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]
302 pub async fn new(
303 auth_token: Option<&str>,
304 #[builder(default = DataServer::ProData)] server: DataServer,
305 handler: T,
306 #[builder(default = 1024*1024)] buffer_size: usize,
307 #[builder(default)] error_config: ErrorRecoveryConfig,
308 ) -> Result<Arc<Self>> {
309 let (write, read) = Self::connect(server, Some(buffer_size)).await?;
310 Self::init_with_stream(
311 handler,
312 server,
313 auth_token,
314 buffer_size,
315 error_config,
316 write,
317 read,
318 )
319 .await
320 }
321
322 async fn init_with_stream(
325 handler: T,
326 server: DataServer,
327 auth_token: Option<&str>,
328 buffer_size: usize,
329 error_config: ErrorRecoveryConfig,
330 write: SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>,
331 read: SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
332 ) -> Result<Arc<Self>> {
333 let auth_token = Ustr::from(auth_token.unwrap_or("unauthorized_user_token"));
334 let is_closed = Arc::new(AtomicBool::new(false));
335 let auth_token = Arc::new(RwLock::new(auth_token));
336 let read = Arc::new(Mutex::new(read));
337
338 let (write_tx, write_rx) = mpsc::channel(1024);
340 let write_tx = Arc::new(RwLock::new(write_tx));
341 let writer_handle = Arc::new(Mutex::new(None::<JoinHandle<()>>));
342
343 let client = Arc::new(Self {
344 handler,
345 server,
346 read,
347 write_tx: write_tx.clone(),
348 writer_handle: writer_handle.clone(),
349 auth_token,
350 is_closed: is_closed.clone(),
351 buffer_size,
352 cancellation: CancellationToken::new(),
353 error_stats: ErrorStats::default(),
354 error_config,
355 health_metrics: Arc::new(RwLock::new(HealthMetrics::default())),
356 circuit_breaker_open: Arc::new(AtomicBool::new(false)),
357 circuit_breaker_opened_at: Arc::new(RwLock::new(None)),
358 });
359
360 Self::spawn_writer(write, write_rx, writer_handle, is_closed.clone());
362
363 if let Err(e) = client.authenticate_on_connect().await {
367 client.cancellation.cancel();
368 let _ = client.close().await;
369 return Err(e);
370 }
371
372 client.spawn_health_monitor();
374
375 Ok(client)
376 }
377
378 async fn authenticate_on_connect(&self) -> Result<()> {
383 let token = *self.auth_token.read().await;
384 self.set_auth_token(token.as_str()).await
385 }
386
387 pub fn spawn_reader_task(self: Arc<Self>) {
388 tokio::spawn(async move {
389 if let Err(e) = self.subscribe().await {
390 error!("Reader task failed: {}", e);
391 }
392 });
393 }
394
395 fn spawn_writer(
398 mut sink: SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>,
399 mut rx: mpsc::Receiver<Message>,
400 handle_storage: Arc<Mutex<Option<JoinHandle<()>>>>,
401 is_closed: Arc<AtomicBool>,
402 ) {
403 let handle = tokio::spawn(async move {
404 while let Some(msg) = rx.recv().await {
405 if is_closed.load(Ordering::Relaxed) {
406 break;
407 }
408 if sink.send(msg).await.is_err() {
409 is_closed.store(true, Ordering::Relaxed);
410 break;
411 }
412 }
413 let _ = sink.close().await;
415 is_closed.store(true, Ordering::Relaxed);
416 });
417
418 tokio::spawn(async move {
420 let mut guard = handle_storage.lock().await;
421 *guard = Some(handle);
422 });
423 }
424
425 async fn connect(
426 server: DataServer,
427 buffer_size: Option<usize>,
428 ) -> Result<(
429 SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>,
430 SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
431 )> {
432 let url = Url::parse(&format!(
433 "wss://{server}.tradingview.com/socket.io/websocket"
434 ))?;
435
436 let buffer_size = buffer_size.unwrap_or(1024 * 1024);
437
438 let mut request = url.into_client_request()?;
439 request.headers_mut().extend((*WEBSOCKET_HEADERS).clone());
440
441 let conf = WebSocketConfig::default()
443 .read_buffer_size(buffer_size)
444 .write_buffer_size(buffer_size);
445
446 let (socket, response) = connect_async_with_config(request, Some(conf), false).await?;
447
448 info!("WebSocket connected with status: {}", response.status());
449
450 let (write, read) = socket.split();
451
452 Ok((write, read))
453 }
454
455 async fn classify_error_severity(&self, error: &Error, context: &str) -> ErrorSeverity {
457 let consecutive_errors = self.error_stats.get_consecutive_errors();
459 let recent_critical = self.error_stats.get_recent_critical_count(300).await;
460
461 let base_severity = match error {
463 Error::WebSocket(msg) => {
464 if msg.contains("ConnectionClosed") || msg.contains("ConnectionReset") {
465 ErrorSeverity::Critical
466 } else if msg.contains("timeout") || msg.contains("WouldBlock") {
467 ErrorSeverity::Moderate
468 } else if msg.contains("Protocol") {
469 ErrorSeverity::Critical
470 } else {
471 ErrorSeverity::Moderate
472 }
473 }
474
475 Error::TradingView { source } => {
476 use crate::error::TradingViewError;
477 match source {
478 TradingViewError::CriticalError => ErrorSeverity::Fatal,
479 TradingViewError::ProtocolError => ErrorSeverity::Critical,
480 TradingViewError::SymbolError | TradingViewError::SeriesError => {
481 ErrorSeverity::Minor
482 }
483 _ => ErrorSeverity::Trace,
484 }
485 }
486
487 Error::JsonParse(_) => {
488 if consecutive_errors > 3 {
489 ErrorSeverity::Moderate
490 } else {
491 ErrorSeverity::Minor
492 }
493 }
494
495 Error::Internal(msg) => {
496 if msg.contains("connection") || msg.contains("timeout") {
497 ErrorSeverity::Critical
498 } else if context.contains("critical") {
499 ErrorSeverity::Fatal
500 } else {
501 ErrorSeverity::Moderate
502 }
503 }
504
505 _ => ErrorSeverity::Moderate,
506 };
507
508 if recent_critical >= 3 {
510 ErrorSeverity::Fatal
511 } else if consecutive_errors >= self.error_config.max_consecutive_errors {
512 std::cmp::max(base_severity, ErrorSeverity::Critical)
513 } else {
514 base_severity
515 }
516 }
517
518 async fn attempt_error_recovery(&self, severity: ErrorSeverity, error: &Error) -> Result<bool> {
519 if self.is_circuit_breaker_open().await {
521 warn!("Circuit breaker is open, skipping recovery attempt");
522 return Ok(false);
523 }
524
525 match severity {
526 ErrorSeverity::Trace => {
527 trace!("Trace level error, no action needed: {}", error);
528 Ok(true)
529 }
530
531 ErrorSeverity::Minor => {
532 debug!("Minor error occurred, continuing: {}", error);
533 Ok(true)
534 }
535
536 ErrorSeverity::Moderate => {
537 warn!(
538 "Moderate error occurred, attempting soft recovery: {}",
539 error
540 );
541
542 if self
544 .error_stats
545 .should_reset_consecutive_errors(self.error_config.error_reset_interval)
546 .await
547 {
548 self.error_stats.reset_consecutive();
549 info!("Reset consecutive error count after timeout period");
550 }
551
552 if let Err(health_err) = self.perform_health_check().await {
554 warn!("Health check failed during recovery: {}", health_err);
555 return Ok(false);
556 }
557
558 Ok(true)
559 }
560
561 ErrorSeverity::Critical => {
562 error!(
563 "Critical error occurred, attempting reconnection: {}",
564 error
565 );
566 self.error_stats
567 .recovery_attempts
568 .fetch_add(1, Ordering::SeqCst);
569
570 for attempt in 1..=self.error_config.max_recovery_attempts {
572 let delay = self.calculate_backoff_delay(attempt);
573 warn!("Recovery attempt {} after {:?} delay", attempt, delay);
574
575 tokio::time::sleep(delay).await;
576
577 match timeout(self.error_config.connection_timeout, self.reconnect()).await {
578 Ok(Ok(_)) => {
579 info!(
580 "Successfully recovered from critical error (attempt {})",
581 attempt
582 );
583 self.error_stats.reset_consecutive();
584 return Ok(true);
585 }
586 Ok(Err(reconnect_err)) => {
587 error!("Reconnection attempt {} failed: {}", attempt, reconnect_err);
588 }
589 Err(_) => {
590 error!("Reconnection attempt {} timed out", attempt);
591 }
592 }
593 }
594
595 self.open_circuit_breaker().await;
597 Ok(false)
598 }
599
600 ErrorSeverity::Fatal => {
601 error!("Fatal error occurred, terminating connection: {}", error);
602 self.is_closed.store(true, Ordering::Relaxed);
603 self.cancellation.cancel();
604 self.open_circuit_breaker().await;
605 Ok(false)
606 }
607 }
608 }
609
610 fn calculate_backoff_delay(&self, attempt: u32) -> Duration {
612 let delay = self
613 .error_config
614 .backoff_base_delay
615 .mul_f64((2_f64).powi(attempt as i32 - 1));
616
617 std::cmp::min(delay, self.error_config.backoff_max_delay)
618 }
619
620 async fn is_circuit_breaker_open(&self) -> bool {
622 if !self.circuit_breaker_open.load(Ordering::Relaxed) {
623 return false;
624 }
625
626 let opened_at = self.circuit_breaker_opened_at.read().await;
628 if let Some(time) = *opened_at
629 && time.elapsed() > Duration::from_secs(300)
630 {
631 self.circuit_breaker_open.store(false, Ordering::Relaxed);
633 info!("Circuit breaker reset after timeout");
634 return false;
635 }
636
637 true
638 }
639
640 async fn open_circuit_breaker(&self) {
641 self.circuit_breaker_open.store(true, Ordering::Relaxed);
642 let mut opened_at = self.circuit_breaker_opened_at.write().await;
643 *opened_at = Some(Instant::now());
644 error!("Circuit breaker opened due to repeated failures");
645 }
646
647 async fn perform_health_check(&self) -> Result<()> {
649 if self.is_closed() {
650 return Err(Error::Internal(ustr("Connection is closed")));
651 }
652
653 if self
655 .error_stats
656 .is_connection_stale(Duration::from_secs(120))
657 .await
658 {
659 warn!("Connection appears stale, performing ping test");
660 }
661
662 let start = Instant::now();
664 self.try_ping().await?;
665 let ping_duration = start.elapsed();
666
667 let mut metrics = self.health_metrics.write().await;
669 metrics.last_ping_time = Some(start);
670
671 if metrics.avg_response_time.is_zero() {
673 metrics.avg_response_time = ping_duration;
674 } else {
675 metrics.avg_response_time = Duration::from_nanos(
676 (metrics.avg_response_time.as_nanos() as f64 * 0.8
677 + ping_duration.as_nanos() as f64 * 0.2) as u64,
678 );
679 }
680
681 metrics.health = if ping_duration > Duration::from_secs(5) {
683 ConnectionHealth::Degraded
684 } else if self.error_stats.get_consecutive_errors() > 2 {
685 ConnectionHealth::Unstable
686 } else {
687 ConnectionHealth::Healthy
688 };
689
690 debug!(
691 "Health check completed: {:?}, ping: {:?}",
692 metrics.health, ping_duration
693 );
694 Ok(())
695 }
696
697 fn spawn_health_monitor(self: &Arc<Self>) {
699 let client = Arc::clone(self);
700 tokio::spawn(async move {
701 let mut interval = tokio::time::interval(client.error_config.health_check_interval);
702
703 loop {
704 tokio::select! {
705 _ = interval.tick() => {
706 if client.is_closed() {
707 break;
708 }
709
710 if let Err(e) = client.perform_health_check().await {
711 warn!("Scheduled health check failed: {}", e);
712 }
713 }
714 _ = client.cancellation.cancelled() => {
715 debug!("Health monitor task cancelled");
716 break;
717 }
718 }
719 }
720 });
721 }
722
723 async fn notify_error_handlers(&self, error: &Error, context: &str, severity: ErrorSeverity) {
725 if severity >= ErrorSeverity::Critical {
727 let now = SystemTime::now()
728 .duration_since(UNIX_EPOCH)
729 .unwrap_or_default()
730 .as_secs();
731 self.error_stats.record_critical_error(now).await;
732 }
733
734 let health_metrics = self.health_metrics.read().await;
736 let error_context = vec![json!({
737 "error_type": format!("{:?}", error),
738 "context": context,
739 "severity": format!("{:?}", severity),
740 "consecutive_errors": self.error_stats.get_consecutive_errors(),
741 "total_errors": self.error_stats.total_errors.load(Ordering::SeqCst),
742 "recovery_attempts": self.error_stats.recovery_attempts.load(Ordering::SeqCst),
743 "connection_health": format!("{:?}", health_metrics.health),
744 "avg_response_time_ms": health_metrics.avg_response_time.as_millis(),
745 "circuit_breaker_open": self.circuit_breaker_open.load(Ordering::Relaxed),
746 "timestamp": chrono::Utc::now().to_rfc3339(),
747 })];
748
749 self.handler.notify_error(*error, &error_context);
751 }
752
753 fn log_error(&self, error: &Error, context: &str, severity: &ErrorSeverity) {
755 let consecutive = self.error_stats.get_consecutive_errors();
756 let total = self.error_stats.total_errors.load(Ordering::SeqCst);
757 let recovery_attempts = self.error_stats.recovery_attempts.load(Ordering::SeqCst);
758
759 let error_info = format!(
760 "{} (consecutive: {}, total: {}, recovery_attempts: {})",
761 error, consecutive, total, recovery_attempts
762 );
763
764 match severity {
765 ErrorSeverity::Trace => {
766 trace!("Trace error in {}: {}", context, error_info);
767 }
768 ErrorSeverity::Minor => {
769 debug!("Minor error in {}: {}", context, error_info);
770 }
771 ErrorSeverity::Moderate => {
772 warn!("Moderate error in {}: {}", context, error_info);
773 }
774 ErrorSeverity::Critical => {
775 error!("Critical error in {}: {}", context, error_info);
776 }
777 ErrorSeverity::Fatal => {
778 error!("FATAL error in {}: {}", context, error_info);
779 }
780 }
781 }
782
783 pub fn is_closed(&self) -> bool {
784 self.is_closed.load(Ordering::Relaxed)
785 }
786
787 pub async fn reconnect(&self) -> Result<()> {
788 let mut wh = self.writer_handle.lock().await;
790 if let Some(handle) = wh.take() {
791 handle.abort();
792 }
793 drop(wh);
794
795 let (write, read) = Self::connect(self.server, Some(self.buffer_size)).await?;
796
797 let (new_tx, new_rx) = mpsc::channel(1024);
799 Self::spawn_writer(
800 write,
801 new_rx,
802 self.writer_handle.clone(),
803 self.is_closed.clone(),
804 );
805
806 {
808 let mut tx_guard = self.write_tx.write().await;
809 *tx_guard = new_tx;
810 }
811
812 let mut read_guard = self.read.lock().await;
813 *read_guard = read;
814 drop(read_guard);
815
816 self.is_closed.store(false, Ordering::Relaxed);
817 self.authenticate_on_connect().await?;
818 Ok(())
819 }
820
821 #[tracing::instrument(skip(self), level = "debug")]
822 pub async fn send_raw_message(&self, message: &str) -> Result<()> {
823 if self.is_closed.load(Ordering::Relaxed) {
824 return Err(Error::Internal("WebSocket is closed".into()));
825 }
826
827 if self.is_circuit_breaker_open().await {
829 return Err(Error::Internal("Circuit breaker is open".into()));
830 }
831
832 let tx = self.write_tx.read().await;
833 match timeout(
834 Duration::from_secs(10),
835 tx.send(Message::Text(message.into())),
836 )
837 .await
838 {
839 Ok(Ok(_)) => {
840 self.error_stats.update_last_successful_message().await;
841 Ok(())
842 }
843 Ok(Err(e)) => Err(Error::WebSocket(e.to_string().into())),
844 Err(_) => Err(Error::Internal("Send timeout".into())),
845 }
846 }
847
848 #[tracing::instrument(skip(self, p), level = "debug")]
849 pub async fn send(&self, m: &str, p: &[Value]) -> Result<()> {
850 if self.is_closed.load(Ordering::Relaxed) {
851 return Err(Error::Internal("WebSocket is closed".into()));
852 }
853 let tx = self.write_tx.read().await;
854 tx.send(SocketMessageSer::new(m, p).to_message()?)
855 .await
856 .map_err(|e| Error::WebSocket(e.to_string().into()))?;
857 Ok(())
858 }
859
860 #[tracing::instrument(skip(self), level = "debug")]
861 pub async fn ping(&self, ping: &Message) -> Result<()> {
862 let tx = self.write_tx.read().await;
863 tx.send(ping.clone())
864 .await
865 .map_err(|e| Error::WebSocket(e.to_string().into()))?;
866 if ping.is_close() {
867 self.is_closed.store(true, Ordering::Relaxed);
868 tracing::warn!("ping message is close, closing session");
869 }
870 Ok(())
871 }
872
873 pub async fn close(&self) -> Result<()> {
874 self.is_closed.store(true, Ordering::Relaxed);
875 let (dummy_tx, _dummy_rx) = mpsc::channel::<Message>(1);
879 let mut tx_guard = self.write_tx.write().await;
880 let old_tx = std::mem::replace(&mut *tx_guard, dummy_tx);
881 drop(old_tx); drop(tx_guard);
883 Ok(())
884 }
885
886 #[tracing::instrument(skip(self), level = "debug")]
887 pub async fn fast_symbols(&self, quote_session: &str, symbols: &[&str]) -> Result<()> {
888 let mut payloads = payload![quote_session];
889 payloads.extend(symbols.iter().map(|s| Value::from(*s)));
890 self.send("quote_fast_symbols", &payloads).await?;
891 Ok(())
892 }
893
894 #[tracing::instrument(skip(self), level = "debug")]
895 pub async fn create_quote_session(&self, quote_session: &str) -> Result<()> {
896 self.send("quote_create_session", &payload!(quote_session))
897 .await?;
898 Ok(())
899 }
900
901 #[tracing::instrument(skip(self), level = "debug")]
902 pub async fn delete_quote_session(&self, quote_session: &str) -> Result<()> {
903 self.send("quote_delete_session", &payload!(quote_session))
904 .await?;
905 Ok(())
906 }
907
908 #[tracing::instrument(skip(self), level = "debug")]
909 pub async fn set_fields(&self, quote_session: &str) -> Result<()> {
910 let mut quote_fields = payload![quote_session];
911 quote_fields.extend(ALL_QUOTE_FIELDS.iter().copied().map(Value::from));
912 self.send("quote_set_fields", "e_fields).await?;
913 Ok(())
914 }
915
916 #[tracing::instrument(skip(self), level = "debug")]
917 pub async fn add_symbols(&self, quote_session: &str, symbols: &[&str]) -> Result<()> {
918 let mut payloads = payload![quote_session];
919 payloads.extend(symbols.iter().map(|s| Value::from(*s)));
920 self.send("quote_add_symbols", &payloads).await?;
921 info!("Added {} symbols to quote session", symbols.len());
922 Ok(())
923 }
924
925 #[tracing::instrument(skip(self), level = "debug")]
926 pub async fn remove_symbols(&self, quote_session: &str, symbols: &[&str]) -> Result<()> {
927 let mut payloads = payload![quote_session];
928 payloads.extend(symbols.iter().map(|s| Value::from(*s)));
929 self.send("quote_remove_symbols", &payloads).await?;
930 Ok(())
931 }
932
933 #[tracing::instrument(skip(self, auth_token), level = "debug")]
934 pub async fn set_auth_token(&self, auth_token: &str) -> Result<()> {
935 {
936 let mut auth_token_ = self.auth_token.write().await;
937 *auth_token_ = ustr(auth_token);
938 }
939 self.send("set_auth_token", &payload!(auth_token)).await?;
940 Ok(())
941 }
942
943 #[tracing::instrument(skip(self), level = "debug")]
945 pub async fn set_locale(&self, language_code: &str, country: &str) -> Result<()> {
946 self.send("set_locale", &payload!(language_code, country))
947 .await?;
948 Ok(())
949 }
950
951 #[tracing::instrument(skip(self), level = "debug")]
952 pub async fn set_data_quality(&self, data_quality: &str) -> Result<()> {
953 self.send("set_data_quality", &payload!(data_quality))
954 .await?;
955
956 Ok(())
957 }
958
959 #[tracing::instrument(skip(self), level = "debug")]
960 pub async fn set_timezone(&self, chart_session: &str, timezone: Timezone) -> Result<()> {
961 self.send(
962 "switch_timezone",
963 &payload!(chart_session, timezone.to_string()),
964 )
965 .await?;
966
967 Ok(())
968 }
969
970 #[tracing::instrument(skip(self), level = "debug")]
971 pub async fn create_chart_session(&self, session: &str) -> Result<()> {
972 self.send("chart_create_session", &payload!(session, ""))
975 .await?;
976 Ok(())
977 }
978
979 #[tracing::instrument(skip(self), level = "debug")]
996 #[builder]
997 pub async fn create_series(
998 &self,
999 chart_session: &str,
1000 series_identifier: &str, series_id: &str, symbol_series_id: &str, interval: Interval,
1004 bar_count: u64,
1005 range: Option<Range>,
1006 ) -> Result<()> {
1007 if let Some(r) = range {
1010 let args = payload!(
1011 chart_session,
1012 series_identifier,
1013 series_id,
1014 symbol_series_id,
1015 interval.to_string(),
1016 0u64, r.to_string() );
1019 self.send("create_series", &args).await?;
1020 } else {
1021 let args: Vec<Value> = vec![
1022 Value::from(chart_session),
1023 Value::from(series_identifier),
1024 Value::from(series_id),
1025 Value::from(symbol_series_id),
1026 Value::from(interval.to_string()),
1027 Value::from(bar_count),
1028 ];
1029 self.send("create_series", &args).await?;
1030 }
1031
1032 Ok(())
1033 }
1034
1035 #[tracing::instrument(skip(self), level = "debug")]
1039 #[builder]
1040 pub async fn modify_series(
1041 &self,
1042 chart_session: &str,
1043 series_identifier: &str, series_id: &str, symbol_series_id: &str, interval: Interval,
1047 bar_count: u64,
1048 range: Option<Range>,
1049 ) -> Result<()> {
1050 if let Some(r) = range {
1051 let args = payload!(
1052 chart_session,
1053 series_identifier,
1054 series_id,
1055 symbol_series_id,
1056 interval.to_string(),
1057 0u64,
1058 r.to_string()
1059 );
1060 self.send("modify_series", &args).await?;
1061 } else {
1062 let args: Vec<Value> = vec![
1063 Value::from(chart_session),
1064 Value::from(series_identifier),
1065 Value::from(series_id),
1066 Value::from(symbol_series_id),
1067 Value::from(interval.to_string()),
1068 Value::from(bar_count),
1069 ];
1070 self.send("modify_series", &args).await?;
1071 }
1072
1073 Ok(())
1074 }
1075
1076 #[tracing::instrument(skip(self), level = "debug")]
1077 pub async fn remove_series(&self, chart_session: &str, series_identifier: &str) -> Result<()> {
1078 self.send("remove_series", &payload!(chart_session, series_identifier))
1079 .await?;
1080 Ok(())
1081 }
1082
1083 #[tracing::instrument(skip(self), level = "debug")]
1084 #[builder]
1085 pub async fn resolve_symbol(
1086 &self,
1087 session: &str,
1088 symbol_series_id: &str,
1089 instrument: &str,
1090 adjustment: Option<MarketAdjustment>,
1091 currency: Option<Currency>,
1092 session_type: Option<SessionType>,
1093 replay_session: Option<&str>,
1094 ) -> Result<()> {
1095 self.send(
1096 "resolve_symbol",
1097 &payload!(
1098 session,
1099 symbol_series_id,
1100 symbol_init()
1101 .instrument(instrument)
1102 .maybe_adjustment(adjustment)
1103 .maybe_currency(currency)
1104 .maybe_session_type(session_type)
1105 .maybe_replay(replay_session)
1106 .call()?
1107 ),
1108 )
1109 .await?;
1110 Ok(())
1111 }
1112
1113 #[tracing::instrument(skip(self), level = "debug")]
1114 pub async fn delete_chart_session(&self, session: &str) -> Result<()> {
1115 self.send("chart_delete_session", &payload!(session))
1116 .await?;
1117 Ok(())
1118 }
1119
1120 #[tracing::instrument(skip(self), level = "debug")]
1121 pub async fn request_more_data(
1122 &self,
1123 chart_session: &str,
1124 series_id: &str,
1125 num: u64,
1126 ) -> Result<()> {
1127 self.send(
1128 "request_more_data",
1129 &payload!(chart_session, series_id, num),
1130 )
1131 .await?;
1132 Ok(())
1133 }
1134
1135 #[tracing::instrument(skip(self), level = "debug")]
1136 pub async fn request_more_tickmarks(
1137 &self,
1138 chart_session: &str,
1139 series_id: &str,
1140 num: u64,
1141 ) -> Result<()> {
1142 self.send(
1143 "request_more_tickmarks",
1144 &payload!(chart_session, series_id, num),
1145 )
1146 .await?;
1147 Ok(())
1148 }
1149
1150 #[tracing::instrument(skip(self), level = "debug")]
1151 pub async fn create_replay_session(&self, replay_session: &str) -> Result<()> {
1152 self.send("replay_create_session", &payload!(replay_session))
1153 .await?;
1154 Ok(())
1155 }
1156
1157 #[tracing::instrument(skip(self), level = "debug")]
1158 #[builder]
1159 pub async fn add_replay_series(
1160 &self,
1161 replay_session: &str,
1162 request_id: &str,
1163 instrument: &str, adjustment: Option<MarketAdjustment>,
1165 session_type: Option<SessionType>,
1166 currency: Option<Currency>,
1167 interval: Interval,
1168 ) -> Result<()> {
1169 let sym_init = symbol_init()
1170 .instrument(instrument)
1171 .maybe_adjustment(adjustment)
1172 .maybe_currency(currency)
1173 .maybe_session_type(session_type)
1174 .call()?;
1175 let payloads =
1176 build_add_replay_series_payload(replay_session, request_id, sym_init, interval);
1177 self.send("replay_add_series", &payloads).await?;
1178 Ok(())
1179 }
1180
1181 #[tracing::instrument(skip(self), level = "debug")]
1182 pub async fn delete_replay_session(&self, replay_session: &str) -> Result<()> {
1183 self.send("replay_delete_session", &payload!(replay_session))
1184 .await?;
1185 Ok(())
1186 }
1187
1188 #[tracing::instrument(skip(self), level = "debug")]
1189 pub async fn replay_step(
1190 &self,
1191 replay_session: &str,
1192 request_id: &str,
1193 step: u64,
1194 ) -> Result<()> {
1195 let payloads = build_replay_step_payload(replay_session, request_id, step);
1196 self.send("replay_step", &payloads).await?;
1197 Ok(())
1198 }
1199
1200 #[tracing::instrument(skip(self), level = "debug")]
1201 pub async fn replay_start(
1202 &self,
1203 replay_session: &str,
1204 request_id: &str,
1205 interval: u64,
1206 ) -> Result<()> {
1207 let payloads = build_replay_start_payload(replay_session, request_id, interval);
1208 self.send("replay_start", &payloads).await?;
1209 Ok(())
1210 }
1211
1212 #[tracing::instrument(skip(self), level = "debug")]
1213 pub async fn replay_stop(&self, replay_session: &str, request_id: &str) -> Result<()> {
1214 let payloads = build_replay_stop_payload(replay_session, request_id);
1215 self.send("replay_stop", &payloads).await?;
1216 Ok(())
1217 }
1218
1219 #[tracing::instrument(skip(self), level = "debug")]
1220 pub async fn replay_reset(
1221 &self,
1222 replay_session: &str,
1223 request_id: &str,
1224 timestamp: i64,
1225 ) -> Result<()> {
1226 let payloads = build_replay_reset_payload(replay_session, request_id, timestamp);
1227 self.send("replay_reset", &payloads).await?;
1228 Ok(())
1229 }
1230
1231 #[tracing::instrument(skip(self), level = "debug")]
1232 #[builder]
1233 pub async fn create_study(
1234 &self,
1235 chart_session: &str,
1236 study_ids: &[&str; 2],
1237 chart_series_id: &str,
1238 study: StudyConfiguration,
1239 ) -> Result<()> {
1240 let mut payloads: Vec<Value> = vec![
1241 Value::from(chart_session),
1242 Value::from(study_ids[0]),
1243 Value::from(study_ids[1]),
1244 Value::from(chart_series_id),
1245 ];
1246
1247 match study {
1248 StudyConfiguration::Pine(pine_indicator) => {
1249 payloads.push(Value::from(pine_indicator.script_type.to_string()));
1250 payloads.push(pine_indicator.to_study_inputs()?);
1251 }
1252 StudyConfiguration::Builtin(study_name, study_config) => {
1253 payloads.push(Value::from(study_name));
1254 payloads.push(json!(study_config));
1255 }
1256 }
1257
1258 self.send("create_study", &payloads).await?;
1259 Ok(())
1260 }
1261
1262 #[tracing::instrument(skip(self), level = "debug")]
1263 #[builder]
1264 pub async fn modify_study(
1265 &self,
1266 chart_session: &str,
1267 study_ids: &[&str; 2],
1268 study: StudyConfiguration,
1269 ) -> Result<()> {
1270 let inputs = match study {
1271 StudyConfiguration::Pine(pine_indicator) => pine_indicator.to_study_inputs()?,
1272 StudyConfiguration::Builtin(_study_name, study_config) => json!(study_config),
1273 };
1274 let payloads =
1275 build_modify_study_payload(chart_session, study_ids[0], study_ids[1], inputs);
1276 self.send("modify_study", &payloads).await?;
1277 Ok(())
1278 }
1279
1280 #[tracing::instrument(skip(self), level = "debug")]
1281 pub async fn remove_study(&self, chart_session: &str, study_id: &str) -> Result<()> {
1282 self.send("remove_study", &payload!(chart_session, study_id))
1283 .await?;
1284 Ok(())
1285 }
1286
1287 pub async fn delete(&self) -> Result<()> {
1288 if let Err(e) = self.close().await {
1290 error!("Failed to close socket: {:?}", e);
1291 return Err(e);
1292 }
1293
1294 debug!("WebSocket client deleted successfully");
1295 Ok(())
1296 }
1297
1298 pub async fn subscribe(&self) -> Result<()> {
1299 let read = self.read.lock().await;
1300 if let Err(e) = self.event_loop(read).await {
1301 error!("Event loop failed: {}", e);
1302 self.is_closed.store(true, Ordering::Relaxed);
1303 self.cancellation.cancel();
1304 return Err(e);
1305 }
1306 Ok(())
1307 }
1308
1309 pub async fn closed_notifier(&self) {
1310 self.cancellation.cancelled().await;
1311 }
1312
1313 pub async fn try_ping(&self) -> Result<()> {
1315 if self.is_closed() {
1316 return Ok(());
1317 }
1318 self.ping(&Message::Ping(Vec::new().into()))
1319 .await
1320 .map_err(|e| Error::WebSocket(ustr(&format!("{e}"))))?;
1321 Ok(())
1322 }
1323
1324 pub async fn get_connection_stats(&self) -> Value {
1325 let health_metrics = self.health_metrics.read().await;
1326 let is_closed = self.is_closed();
1327
1328 json!({
1329 "consecutive_errors": self.error_stats.get_consecutive_errors(),
1330 "total_errors": self.error_stats.total_errors.load(Ordering::SeqCst),
1331 "recovery_attempts": self.error_stats.recovery_attempts.load(Ordering::SeqCst),
1332 "connection_drops": self.error_stats.connection_drops.load(Ordering::SeqCst),
1333 "health": format!("{:?}", health_metrics.health),
1334 "avg_response_time_ms": health_metrics.avg_response_time.as_millis(),
1335 "circuit_breaker_open": self.circuit_breaker_open.load(Ordering::Relaxed),
1336 "is_closed": is_closed,
1337 })
1338 }
1339}
1340
1341impl<T: Handler> Socket for WebSocketClient<T> {
1342 async fn event_loop(
1343 &self,
1344 mut read: MutexGuard<'_, SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>>,
1345 ) -> Result<()> {
1346 trace!("WebSocket event loop started");
1347
1348 let mut ping_interval = tokio::time::interval(self.error_config.ping_interval);
1349 ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
1350
1351 loop {
1352 if self.is_closed.load(Ordering::Relaxed) {
1353 trace!("WebSocket is closed, ending event loop");
1354 break;
1355 }
1356
1357 tokio::select! {
1358 message_result = timeout(Duration::from_secs(30), read.next()) => {
1360 match message_result {
1361 Ok(Some(Ok(message))) => {
1362 trace!("Received message: {:?}", message);
1363 if let Err(e) = self.handle_raw_messages(message).await {
1364 self.handle_error(e, ustr("handle_raw_messages")).await?;
1365 } else {
1366 if self.error_stats.get_consecutive_errors() > 0 {
1368 self.error_stats.reset_consecutive();
1369 debug!("Reset consecutive errors after successful message processing");
1370 }
1371 self.error_stats.update_last_successful_message().await;
1372 }
1373 }
1374 Ok(Some(Err(e))) => {
1375 error!("Error reading message: {:#?}", e);
1376 self.error_stats.connection_drops.fetch_add(1, Ordering::SeqCst);
1377 self.handle_error(
1378 Error::WebSocket(e.to_string().into()),
1379 ustr("event_loop_read"),
1380 ).await?;
1381
1382 if e.to_string().contains("ConnectionClosed") ||
1384 e.to_string().contains("ConnectionReset") {
1385 break;
1386 }
1387 }
1388 Ok(None) => {
1389 info!("WebSocket stream ended");
1390 self.is_closed.store(true, Ordering::Relaxed);
1391 break;
1392 }
1393 Err(_) => {
1394 warn!("WebSocket read timeout, checking connection health");
1395 if let Err(e) = self.perform_health_check().await {
1396 warn!("Health check failed during timeout: {}", e);
1397 if matches!(self.classify_error_severity(&e, "health_check").await, ErrorSeverity::Fatal) {
1399 break;
1400 }
1401 }
1402 }
1403 }
1404 }
1405
1406 _ = ping_interval.tick() => {
1408 if !self.is_closed()
1409 && let Err(e) = self.try_ping().await
1410 {
1411 warn!("Periodic ping failed: {}", e);
1412 self.handle_error(e, ustr("periodic_ping")).await?;
1413 }
1414 }
1415
1416 _ = self.cancellation.cancelled() => {
1418 info!("Event loop cancelled");
1419 break;
1420 }
1421 }
1422 }
1423
1424 trace!("WebSocket event loop ended");
1425 Ok(())
1426 }
1427
1428 async fn handle_raw_messages(&self, raw: Message) -> Result<()> {
1429 match &raw {
1430 Message::Text(text) => {
1431 trace!("Received text message: {}", text);
1432 self.handle_parsed_messages(parse_packet(text), &raw)
1433 .await?;
1434 }
1435 Message::Close(msg) => {
1436 warn!("Connection closed with code: {:?}", msg);
1437 self.is_closed.store(true, Ordering::Relaxed);
1438 self.cancellation.cancel();
1439 }
1440 Message::Binary(msg) => {
1441 debug!("Received binary message: {:?}", msg);
1442 }
1444 Message::Ping(msg) => {
1445 trace!("Received ping message: {:?}", msg);
1446 }
1447 Message::Pong(msg) => {
1448 trace!("Received pong message: {:?}", msg);
1449 }
1450 Message::Frame(f) => {
1451 debug!("Received frame message: {:?}", f);
1452 }
1453 }
1454 Ok(())
1455 }
1456
1457 async fn handle_parsed_messages(
1458 &self,
1459 messages: Vec<SocketMessage<SocketMessageDe>>,
1460 _raw: &Message,
1461 ) -> Result<()> {
1462 for message in messages {
1463 match message {
1464 SocketMessage::SocketServerInfo(info) => {
1465 trace!("received server info: {:?}", info);
1466 }
1467 SocketMessage::SocketMessage(msg) => {
1468 trace!(
1469 "Processing socket message: method={}, params={:?}",
1470 msg.m, msg.p
1471 );
1472 if let Err(e) = self.handle_message_data(msg).await {
1473 self.handle_error(e, ustr("handle_message_data")).await?;
1474 }
1475 }
1476 SocketMessage::Heartbeat(counter) => {
1477 debug!("handling heartbeat message: {counter}");
1478 let echo = message.heartbeat_echo().unwrap_or_else(|| {
1479 let h = format!("~h~{counter}");
1480 format!("~m~{}~m~{h}", h.len())
1481 });
1482 if let Err(e) = self.send_raw_message(&echo).await {
1483 self.handle_error(e, ustr("heartbeat_echo")).await?;
1484 }
1485 }
1486 SocketMessage::Other(value) => {
1487 trace!("Received other message: {:?}", value);
1488 if let Ok(server_info) = SocketServerInfo::deserialize(&value) {
1489 info!("{}", server_info);
1490 } else {
1491 warn!("Received unrecognized message: {:?}", value);
1492 }
1493 }
1494 SocketMessage::Unknown(s) => {
1495 warn!("unknown message: {:?}", s);
1496 }
1497 }
1498 }
1499 Ok(())
1500 }
1501
1502 #[tracing::instrument(skip(self), level = "trace")]
1503 async fn handle_message_data(&self, message: SocketMessageDe) -> Result<()> {
1504 dispatch_message_data(&self.handler, message);
1505 Ok(())
1506 }
1507
1508 async fn handle_error(&self, error: Error, context: Ustr) -> Result<()> {
1509 let context_str = context.as_str();
1510
1511 let _consecutive_errors = self.error_stats.increment_error();
1513 self.error_stats.update_last_error_time().await;
1514
1515 let severity = self.classify_error_severity(&error, context_str).await;
1517
1518 self.log_error(&error, context_str, &severity);
1520
1521 self.notify_error_handlers(&error, context_str, severity)
1523 .await;
1524
1525 match severity {
1527 ErrorSeverity::Trace | ErrorSeverity::Minor => {
1528 Ok(())
1530 }
1531
1532 ErrorSeverity::Moderate => {
1533 match self.attempt_error_recovery(severity, &error).await {
1535 Ok(true) => Ok(()),
1536 Ok(false) => {
1537 warn!("Moderate error recovery failed, continuing anyway");
1538 Ok(())
1539 }
1540 Err(recovery_err) => {
1541 warn!("Error during moderate recovery: {}", recovery_err);
1542 Ok(()) }
1544 }
1545 }
1546
1547 ErrorSeverity::Critical | ErrorSeverity::Fatal => {
1548 match self.attempt_error_recovery(severity, &error).await {
1550 Ok(recovered) => {
1551 if !recovered {
1552 error!(
1553 "Failed to recover from {} error",
1554 if matches!(severity, ErrorSeverity::Critical) {
1555 "critical"
1556 } else {
1557 "fatal"
1558 }
1559 );
1560 return Err(Error::Internal(ustr("Error recovery failed")));
1561 }
1562 Ok(())
1563 }
1564 Err(recovery_err) => {
1565 error!("Error during recovery attempt: {}", recovery_err);
1566 Err(recovery_err)
1567 }
1568 }
1569 }
1570 }
1571 }
1572}
1573
1574pub fn dispatch_message_data<H: Handler>(handler: &H, message: SocketMessageDe) {
1576 let event = TradingViewDataEvent::from(message.m);
1577 handler.handle_events(event, &message.p);
1578 if event == TradingViewDataEvent::OnQuoteData {
1579 handler.handle_quote_data(&message.p);
1580 }
1581}
1582
1583pub fn build_modify_study_payload(
1587 chart_session: &str,
1588 study_id: &str,
1589 study_sub_id: &str,
1590 inputs: Value,
1591) -> Vec<Value> {
1592 vec![
1593 Value::from(chart_session),
1594 Value::from(study_id),
1595 Value::from(study_sub_id),
1596 inputs,
1597 ]
1598}
1599
1600pub fn build_replay_start_payload(
1604 replay_session: &str,
1605 request_id: &str,
1606 interval_ms: u64,
1607) -> Vec<Value> {
1608 vec![
1609 Value::from(replay_session),
1610 Value::from(request_id),
1611 Value::from(interval_ms),
1612 ]
1613}
1614
1615pub fn build_replay_step_payload(replay_session: &str, request_id: &str, step: u64) -> Vec<Value> {
1619 vec![
1620 Value::from(replay_session),
1621 Value::from(request_id),
1622 Value::from(step),
1623 ]
1624}
1625
1626pub fn build_replay_stop_payload(replay_session: &str, request_id: &str) -> Vec<Value> {
1630 vec![Value::from(replay_session), Value::from(request_id)]
1631}
1632
1633pub fn build_replay_reset_payload(
1637 replay_session: &str,
1638 request_id: &str,
1639 timestamp: i64,
1640) -> Vec<Value> {
1641 vec![
1642 Value::from(replay_session),
1643 Value::from(request_id),
1644 Value::from(timestamp),
1645 ]
1646}
1647
1648pub fn build_add_replay_series_payload(
1652 replay_session: &str,
1653 request_id: &str,
1654 symbol_init: String,
1655 interval: Interval,
1656) -> Vec<Value> {
1657 vec![
1658 Value::from(replay_session),
1659 Value::from(request_id),
1660 Value::from(symbol_init),
1661 Value::from(interval.to_string()),
1662 ]
1663}
1664
1665#[cfg(test)]
1666mod tests {
1667 use super::*;
1668 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1669
1670 struct MockHandler {
1671 event_called: AtomicBool,
1672 quote_called: AtomicBool,
1673 quote_count: AtomicUsize,
1674 }
1675
1676 impl MockHandler {
1677 fn new() -> Self {
1678 Self {
1679 event_called: AtomicBool::new(false),
1680 quote_called: AtomicBool::new(false),
1681 quote_count: AtomicUsize::new(0),
1682 }
1683 }
1684 }
1685
1686 impl Handler for MockHandler {
1687 fn handle_events(&self, _event: TradingViewDataEvent, _message: &[Value]) {
1688 self.event_called.store(true, Ordering::SeqCst);
1689 }
1690 fn handle_quote_data(&self, _message: &[Value]) {
1691 self.quote_called.store(true, Ordering::SeqCst);
1692 self.quote_count.fetch_add(1, Ordering::SeqCst);
1693 }
1694 fn handle_series_data(&self, _event: TradingViewDataEvent, _messages: &[Value]) {}
1695 fn notify_error(&self, _error: Error, _message: &[Value]) {}
1696 }
1697
1698 #[test]
1699 fn test_modify_study_payload_has_4_args() {
1700 let inputs = serde_json::json!({ "text": "//@version=5\nindicator('Test')" });
1701 let payload = build_modify_study_payload("cs_123", "st6", "st1", inputs.clone());
1702
1703 assert_eq!(
1704 payload.len(),
1705 4,
1706 "modify_study payload length must be exactly 4"
1707 );
1708 assert_eq!(payload[0], Value::from("cs_123"));
1709 assert_eq!(payload[1], Value::from("st6"));
1710 assert_eq!(payload[2], Value::from("st1"));
1711 assert_eq!(payload[3], inputs);
1712 }
1713
1714 #[test]
1715 fn test_replay_start_payload_interval_is_numeric_millis() {
1716 let payload = build_replay_start_payload("rs_100", "req_replay_1", 250);
1717
1718 assert_eq!(payload.len(), 3, "replay_start payload length must be 3");
1719 assert_eq!(payload[0], Value::from("rs_100"));
1720 assert_eq!(payload[1], Value::from("req_replay_1"));
1721 assert!(
1722 payload[2].is_number(),
1723 "third argument must be numeric milliseconds"
1724 );
1725 assert_eq!(payload[2].as_u64(), Some(250));
1726 }
1727
1728 #[test]
1729 fn test_replay_other_payloads() {
1730 let step_payload = build_replay_step_payload("rs_100", "req_step_1", 5);
1731 assert_eq!(step_payload.len(), 3);
1732 assert_eq!(step_payload[2], Value::from(5u64));
1733
1734 let stop_payload = build_replay_stop_payload("rs_100", "req_stop_1");
1735 assert_eq!(stop_payload.len(), 2);
1736 assert_eq!(stop_payload[0], Value::from("rs_100"));
1737 assert_eq!(stop_payload[1], Value::from("req_stop_1"));
1738
1739 let reset_payload = build_replay_reset_payload("rs_100", "req_reset_1", 1_700_000_000);
1740 assert_eq!(reset_payload.len(), 3);
1741 assert_eq!(reset_payload[2], Value::from(1_700_000_000i64));
1742
1743 let add_payload = build_add_replay_series_payload(
1744 "rs_100",
1745 "req_add_1",
1746 "={\"symbol\":\"NASDAQ:AAPL\"}".to_string(),
1747 Interval::OneDay,
1748 );
1749 assert_eq!(add_payload.len(), 4);
1750 assert_eq!(add_payload[3], Value::from("1D"));
1751 }
1752
1753 #[test]
1754 fn test_heartbeat_echo_framed_format() {
1755 let msg = SocketMessage::<SocketMessageDe>::Heartbeat(42);
1756 assert_eq!(msg.heartbeat_echo(), Some("~m~5~m~~h~42".to_string()));
1757 }
1758
1759 #[test]
1760 fn test_quote_data_handler_routing() {
1761 let handler = MockHandler::new();
1762 let quote_msg = SocketMessageDe {
1763 m: ustr::ustr("qsd"),
1764 p: vec![
1765 serde_json::json!("quote_session_1"),
1766 serde_json::json!({"s": "ok"}),
1767 ],
1768 t: 0,
1769 t_ms: 0,
1770 };
1771
1772 dispatch_message_data(&handler, quote_msg);
1773 assert!(
1774 handler.event_called.load(Ordering::SeqCst),
1775 "handle_events must be called"
1776 );
1777 assert!(
1778 handler.quote_called.load(Ordering::SeqCst),
1779 "handle_quote_data must be called for qsd"
1780 );
1781 assert_eq!(handler.quote_count.load(Ordering::SeqCst), 1);
1782
1783 let non_quote_msg = SocketMessageDe {
1784 m: ustr::ustr("timescale_update"),
1785 p: vec![serde_json::json!("cs_1")],
1786 t: 0,
1787 t_ms: 0,
1788 };
1789 dispatch_message_data(&handler, non_quote_msg);
1790 assert_eq!(
1791 handler.quote_count.load(Ordering::SeqCst),
1792 1,
1793 "quote_data should not be called for non-qsd"
1794 );
1795 }
1796
1797 #[tokio::test]
1798 async fn test_local_ws_peer_observes_auth_first_and_session_order() {
1799 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1800 let server_addr = listener.local_addr().unwrap();
1801
1802 let server_task = tokio::spawn(async move {
1803 let (tcp_stream, _) = listener.accept().await.unwrap();
1804 let mut ws_stream = tokio_tungstenite::accept_async(tcp_stream).await.unwrap();
1805 let first_msg = ws_stream.next().await.unwrap().unwrap();
1806 let second_msg = ws_stream.next().await.unwrap().unwrap();
1807 (first_msg, second_msg)
1808 });
1809
1810 let client_tcp = tokio::net::TcpStream::connect(server_addr).await.unwrap();
1811 let plain = MaybeTlsStream::Plain(client_tcp);
1812 let (ws_stream, _) =
1813 tokio_tungstenite::client_async(format!("ws://{}", server_addr), plain)
1814 .await
1815 .unwrap();
1816
1817 let (write, read) = ws_stream.split();
1818 let client = WebSocketClient::init_with_stream(
1819 MockHandler::new(),
1820 DataServer::Data,
1821 Some("test_token_secret"),
1822 1024 * 1024,
1823 ErrorRecoveryConfig::default(),
1824 write,
1825 read,
1826 )
1827 .await
1828 .unwrap();
1829
1830 client
1832 .create_quote_session("quote_session_test")
1833 .await
1834 .unwrap();
1835
1836 let (first, second) = server_task.await.unwrap();
1837
1838 if let Message::Text(first_text) = first {
1839 assert!(
1840 first_text.contains("set_auth_token"),
1841 "First message must be set_auth_token: {}",
1842 first_text
1843 );
1844 assert!(
1845 first_text.contains("test_token_secret"),
1846 "First message must contain token: {}",
1847 first_text
1848 );
1849 } else {
1850 panic!("Expected Text message for first frame");
1851 }
1852
1853 if let Message::Text(second_text) = second {
1854 assert!(
1855 second_text.contains("quote_create_session"),
1856 "Second message must be quote_create_session: {}",
1857 second_text
1858 );
1859 assert!(
1860 second_text.contains("quote_session_test"),
1861 "Second message must contain session name: {}",
1862 second_text
1863 );
1864 } else {
1865 panic!("Expected Text message for second frame");
1866 }
1867
1868 let _ = client.close().await;
1869 }
1870
1871 #[tokio::test]
1872 async fn test_local_ws_peer_observes_anonymous_auth() {
1873 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1874 let server_addr = listener.local_addr().unwrap();
1875
1876 let server_task = tokio::spawn(async move {
1877 let (tcp_stream, _) = listener.accept().await.unwrap();
1878 let mut ws_stream = tokio_tungstenite::accept_async(tcp_stream).await.unwrap();
1879 let first_msg = ws_stream.next().await.unwrap().unwrap();
1880 let second_msg = ws_stream.next().await.unwrap().unwrap();
1881 (first_msg, second_msg)
1882 });
1883
1884 let client_tcp = tokio::net::TcpStream::connect(server_addr).await.unwrap();
1885 let plain = MaybeTlsStream::Plain(client_tcp);
1886 let (ws_stream, _) =
1887 tokio_tungstenite::client_async(format!("ws://{}", server_addr), plain)
1888 .await
1889 .unwrap();
1890
1891 let (write, read) = ws_stream.split();
1892 let client = WebSocketClient::init_with_stream(
1893 MockHandler::new(),
1894 DataServer::Data,
1895 None,
1896 1024 * 1024,
1897 ErrorRecoveryConfig::default(),
1898 write,
1899 read,
1900 )
1901 .await
1902 .unwrap();
1903
1904 client.create_quote_session("anon_session").await.unwrap();
1905
1906 let (first, second) = server_task.await.unwrap();
1907 if let Message::Text(first_text) = first {
1908 assert!(
1909 first_text.contains("set_auth_token"),
1910 "First message must be set_auth_token: {}",
1911 first_text
1912 );
1913 assert!(
1914 first_text.contains("unauthorized_user_token"),
1915 "Anonymous auth must contain unauthorized_user_token: {}",
1916 first_text
1917 );
1918 } else {
1919 panic!("Expected Text message for first frame");
1920 }
1921
1922 if let Message::Text(second_text) = second {
1923 assert!(
1924 second_text.contains("quote_create_session"),
1925 "Second message must be quote_create_session: {}",
1926 second_text
1927 );
1928 assert!(
1929 second_text.contains("anon_session"),
1930 "Second message must contain anon_session: {}",
1931 second_text
1932 );
1933 } else {
1934 panic!("Expected Text message for second frame");
1935 }
1936
1937 let _ = client.close().await;
1938 }
1939}