1use std::sync::{Arc, Mutex};
2use std::time::Instant;
3use tracing::{debug, error, info, instrument, warn};
4
5use serde::Deserialize;
6use serde_json::Value;
7
8use crate::{
9 DataPoint, DataServer, Error, Result, SymbolInfo,
10 chart::options::Range,
11 historical::{HistoricalRequest, HistoricalResult, state::HistoricalState},
12 live::handler::{CommandTx, Handler, HandlerFactory},
13 live::models::TradingViewDataEvent,
14 live::websocket::WebSocketClient,
15 utils::symbol_init,
16};
17
18pub struct HistoricalClient {
20 pub(crate) auth_token: String,
21 pub(crate) server: DataServer,
22}
23
24impl HistoricalClient {
25 pub fn new(auth_token: impl Into<String>, server: DataServer) -> Self {
26 Self {
27 auth_token: auth_token.into(),
28 server,
29 }
30 }
31
32 #[instrument(skip(self), fields(symbol, exchange))]
33 pub async fn retrieve(&self, request: HistoricalRequest) -> Result<HistoricalResult> {
34 let started = Instant::now();
35 let (symbol, exchange) = request.resolve_symbol_exchange()?;
36 debug!(symbol = %symbol, exchange = %exchange, "Historical retrieval started");
37
38 let state = Arc::new(Mutex::new(if let Some(n) = request.num_bars {
39 HistoricalState::with_capacity(n as usize)
40 } else {
41 HistoricalState::new()
42 }));
43
44 let (cmd_tx, _cmd_rx) =
45 tokio::sync::mpsc::channel::<crate::live::handler::command::Command>(16);
46 let factory = HistoricalDataHandlerFactory::new(Arc::clone(&state));
47 let handler = factory.create(cmd_tx);
48
49 let ws = WebSocketClient::builder()
50 .auth_token(&self.auth_token)
51 .server(self.server)
52 .handler(handler)
53 .build()
54 .await?;
55
56 let instrument = format!("{exchange}:{symbol}");
64 let chart_session = format!("cs_{}", crate::utils::gen_id());
65 let symbol_series_id = format!("sds_sym_{}", crate::utils::gen_id());
66 let series_identifier = "sds_1".to_string();
67 let series_id = "s1".to_string();
68
69 ws.create_chart_session(&chart_session).await?;
71 debug!(session = %chart_session, "Chart session created");
72
73 let symbol_init_str = symbol_init().instrument(&instrument).call()?;
75 ws.send(
76 "resolve_symbol",
77 &[
78 Value::from(chart_session.as_str()),
79 Value::from(symbol_series_id.as_str()),
80 Value::from(symbol_init_str),
81 ],
82 )
83 .await?;
84 debug!(instrument = %instrument, "Symbol resolution requested");
85
86 let create_series_args = build_create_series_args(
90 &chart_session,
91 &series_identifier,
92 &series_id,
93 &symbol_series_id,
94 request.interval,
95 request.num_bars,
96 request.range,
97 );
98 ws.send("create_series", &create_series_args).await?;
99 debug!(
100 interval = ?request.interval,
101 bars = request.num_bars,
102 range = ?request.range,
103 "Data series created"
104 );
105
106 let qs = format!("qs_{}", crate::utils::gen_id());
108 ws.send("quote_create_session", &[Value::from(qs.as_str())])
109 .await?;
110 ws.send("quote_set_fields", &[Value::from(qs.as_str())])
111 .await?;
112 ws.send(
113 "quote_add_symbols",
114 &[Value::from(qs.as_str()), Value::from(symbol.as_str())],
115 )
116 .await?;
117
118 Arc::clone(&ws).spawn_reader_task();
119
120 let result = tokio::time::timeout(request.timeout, Self::wait_for_completion(&state)).await;
121
122 let mut state_guard = state.lock().unwrap();
123 let total_bars = state_guard.total_bars;
124 let data = state_guard.finalize();
125 let elapsed = started.elapsed();
126
127 match result {
128 Ok(_) => {
129 if state_guard.errored {
130 let msg = state_guard
131 .error_message
132 .take()
133 .unwrap_or_else(|| "Historical data retrieval failed".to_string());
134 return Err(Error::Internal(msg.into()));
135 }
136 let symbol_info = state_guard
137 .symbol_info
138 .take()
139 .ok_or_else(|| Error::Internal("No symbol info received".into()))?;
140 Ok(HistoricalResult {
141 symbol_info,
142 data,
143 series_info: state_guard.series_info.take(),
144 total_bars_received: total_bars,
145 replay_used: request.with_replay,
146 elapsed,
147 })
148 }
149 Err(_) => Err(Error::Timeout("Historical data retrieval timed out".into())),
150 }
151 }
152
153 pub(crate) async fn wait_for_completion(state: &Arc<Mutex<HistoricalState>>) {
154 let notify = {
155 let guard = state.lock().unwrap();
156 guard.notify.clone()
157 };
158 loop {
159 let notified = notify.notified();
161 tokio::pin!(notified);
162 notified.as_mut().enable();
163
164 {
165 let guard = state.lock().unwrap();
166 if guard.completed || guard.errored {
167 break;
168 }
169 }
170
171 notified.await;
172 }
173 }
174}
175
176pub(crate) fn build_create_series_args(
186 chart_session: &str,
187 series_identifier: &str,
188 series_id: &str,
189 symbol_series_id: &str,
190 interval: crate::Interval,
191 num_bars: Option<u64>,
192 range: Option<Range>,
193) -> Vec<Value> {
194 if let Some(r) = range {
195 vec![
196 Value::from(chart_session),
197 Value::from(series_identifier),
198 Value::from(series_id),
199 Value::from(symbol_series_id),
200 Value::from(interval.to_string()),
201 Value::from(0u64), Value::from(r.to_string()),
203 ]
204 } else {
205 let bar_count = num_bars.unwrap_or(100);
206 vec![
207 Value::from(chart_session),
208 Value::from(series_identifier),
209 Value::from(series_id),
210 Value::from(symbol_series_id),
211 Value::from(interval.to_string()),
212 Value::from(bar_count),
213 ]
214 }
215}
216
217#[derive(Clone)]
225pub struct HistoricalDataHandler {
226 state: Arc<Mutex<HistoricalState>>,
227 #[allow(dead_code)]
228 cmd_tx: CommandTx,
229}
230impl Handler for HistoricalDataHandler {
231 fn handle_events(&self, event: TradingViewDataEvent, message: &[Value]) {
232 match event {
233 TradingViewDataEvent::OnSymbolResolved => {
234 if let Some(sym_info) = message.get(2)
236 && let Ok(info) = SymbolInfo::deserialize(sym_info)
237 {
238 debug!(name = %info.name, "Symbol resolved");
239 self.state.lock().unwrap().record_symbol_info(info);
240 }
241 }
242 TradingViewDataEvent::OnChartData | TradingViewDataEvent::OnChartDataUpdate => {
243 if message.len() < 2 {
244 return;
245 }
246 if let Some(obj) = message[1].as_object() {
247 for (_key, series_val) in obj {
248 if let Some(s_arr) = series_val.get("s").and_then(|v| v.as_array()) {
249 let mut points = Vec::with_capacity(s_arr.len());
250 for v in s_arr {
251 if let Ok(point) = DataPoint::deserialize(v) {
252 points.push(point);
253 }
254 }
255 if !points.is_empty() {
256 let mut state = self.state.lock().unwrap();
257 state.record_points(points, s_arr.len());
258 }
259 }
260 }
261 }
262 }
263 TradingViewDataEvent::OnSeriesCompleted => {
264 info!("Series completed");
265 self.state.lock().unwrap().complete();
266 }
267 TradingViewDataEvent::OnError(tv_error) => {
268 error!(?tv_error, "TradingView protocol error");
269 let mut state = self.state.lock().unwrap();
270 state.fail(format!("TradingView error: {tv_error:?}"));
271 }
272 _ => {}
273 }
274 }
275
276 fn handle_quote_data(&self, _message: &[Value]) {}
277 fn handle_series_data(&self, _event: TradingViewDataEvent, _messages: &[Value]) {}
278
279 fn notify_error(&self, error: Error, _message: &[Value]) {
280 warn!(?error, "Historical handler error");
281 let mut state = self.state.lock().unwrap();
282 if state.record_error() {
283 state.fail(format!("Too many errors: {error:?}"));
284 }
285 }
286}
287
288pub struct HistoricalDataHandlerFactory {
295 state: Arc<Mutex<HistoricalState>>,
296}
297
298impl HistoricalDataHandlerFactory {
299 pub fn new(state: Arc<Mutex<HistoricalState>>) -> Self {
301 Self { state }
302 }
303}
304
305impl HandlerFactory for HistoricalDataHandlerFactory {
306 type Handler = HistoricalDataHandler;
307
308 fn create(&self, command_tx: CommandTx) -> Self::Handler {
309 HistoricalDataHandler {
310 state: Arc::clone(&self.state),
311 cmd_tx: command_tx,
312 }
313 }
314}
315
316#[cfg(test)]
317mod tests {
318 use super::*;
319 use crate::Interval;
320 use crate::chart::options::Range;
321 use crate::error::TradingViewError;
322
323 #[test]
324 fn test_create_series_args_count_mode_default_bars() {
325 let args = build_create_series_args(
326 "cs_test",
327 "sds_1",
328 "s1",
329 "sds_sym_1",
330 Interval::OneDay,
331 None,
332 None,
333 );
334 assert_eq!(
335 args.len(),
336 6,
337 "Count mode without range must have exactly 6 arguments"
338 );
339 assert_eq!(args[0], Value::from("cs_test"));
340 assert_eq!(args[1], Value::from("sds_1"));
341 assert_eq!(args[2], Value::from("s1"));
342 assert_eq!(args[3], Value::from("sds_sym_1"));
343 assert_eq!(args[4], Value::from("1D"));
344 assert_eq!(
345 args[5],
346 Value::from(100u64),
347 "Default bar count must be 100"
348 );
349 }
350
351 #[test]
352 fn test_create_series_args_count_mode_custom_bars() {
353 let args = build_create_series_args(
354 "cs_test",
355 "sds_1",
356 "s1",
357 "sds_sym_1",
358 Interval::FiveMinutes,
359 Some(500),
360 None,
361 );
362 assert_eq!(
363 args.len(),
364 6,
365 "Count mode without range must have exactly 6 arguments"
366 );
367 assert_eq!(args[4], Value::from("5"));
368 assert_eq!(args[5], Value::from(500u64));
369 }
370
371 #[test]
372 fn test_create_series_args_range_mode_from_to() {
373 let range = Range::FromTo(1626220800, 1628640000);
374 let args = build_create_series_args(
375 "cs_test",
376 "sds_1",
377 "s1",
378 "sds_sym_1",
379 Interval::OneDay,
380 Some(500), Some(range),
382 );
383 assert_eq!(args.len(), 7, "Range mode must have exactly 7 arguments");
384 assert_eq!(args[0], Value::from("cs_test"));
385 assert_eq!(args[1], Value::from("sds_1"));
386 assert_eq!(args[2], Value::from("s1"));
387 assert_eq!(args[3], Value::from("sds_sym_1"));
388 assert_eq!(args[4], Value::from("1D"));
389 assert_eq!(
390 args[5],
391 Value::from(0u64),
392 "Bar count in range mode must be 0"
393 );
394 assert_eq!(args[6], Value::from(range.to_string()));
395 }
396
397 #[test]
398 fn test_create_series_args_range_mode_preset() {
399 let range = Range::OneDay;
400 let args = build_create_series_args(
401 "cs_test",
402 "sds_1",
403 "s1",
404 "sds_sym_1",
405 Interval::OneMinute,
406 None,
407 Some(range),
408 );
409 assert_eq!(args.len(), 7, "Range mode must have exactly 7 arguments");
410 assert_eq!(
411 args[5],
412 Value::from(0u64),
413 "Bar count in range mode must be 0"
414 );
415 assert_eq!(args[6], Value::from(range.to_string()));
416 }
417
418 #[tokio::test]
419 async fn test_wait_for_completion_wakes_immediately_on_complete() {
420 let state = Arc::new(Mutex::new(HistoricalState::new()));
421 let state_clone = Arc::clone(&state);
422
423 let wait_handle = tokio::spawn(async move {
424 HistoricalClient::wait_for_completion(&state_clone).await;
425 });
426
427 tokio::task::yield_now().await;
428
429 let start = std::time::Instant::now();
430 state.lock().unwrap().complete();
431
432 let timeout_result =
433 tokio::time::timeout(std::time::Duration::from_millis(50), wait_handle).await;
434 assert!(
435 timeout_result.is_ok(),
436 "wait_for_completion must wake without polling"
437 );
438 assert!(start.elapsed() < std::time::Duration::from_millis(50));
439 assert!(state.lock().unwrap().completed);
440 }
441
442 #[tokio::test]
443 async fn test_wait_for_completion_wakes_immediately_on_fail() {
444 let state = Arc::new(Mutex::new(HistoricalState::new()));
445 let state_clone = Arc::clone(&state);
446
447 let wait_handle = tokio::spawn(async move {
448 HistoricalClient::wait_for_completion(&state_clone).await;
449 });
450
451 tokio::task::yield_now().await;
452
453 let start = std::time::Instant::now();
454 state.lock().unwrap().fail("simulated error".into());
455
456 let timeout_result =
457 tokio::time::timeout(std::time::Duration::from_millis(50), wait_handle).await;
458 assert!(
459 timeout_result.is_ok(),
460 "wait_for_completion must wake immediately on failure"
461 );
462 assert!(start.elapsed() < std::time::Duration::from_millis(50));
463 let guard = state.lock().unwrap();
464 assert!(guard.errored);
465 assert_eq!(guard.error_message.as_deref(), Some("simulated error"));
466 }
467
468 #[tokio::test]
469 async fn test_wait_for_completion_lost_wakeup_safe_pre_completed() {
470 let state = Arc::new(Mutex::new(HistoricalState::new()));
471
472 state.lock().unwrap().complete();
474
475 let result = tokio::time::timeout(
476 std::time::Duration::from_millis(20),
477 HistoricalClient::wait_for_completion(&state),
478 )
479 .await;
480
481 assert!(
482 result.is_ok(),
483 "wait_for_completion must return immediately if already completed"
484 );
485 }
486
487 #[tokio::test]
488 async fn test_wait_for_completion_lost_wakeup_safe_pre_errored() {
489 let state = Arc::new(Mutex::new(HistoricalState::new()));
490
491 state.lock().unwrap().fail("early failure".into());
493
494 let result = tokio::time::timeout(
495 std::time::Duration::from_millis(20),
496 HistoricalClient::wait_for_completion(&state),
497 )
498 .await;
499
500 assert!(
501 result.is_ok(),
502 "wait_for_completion must return immediately if already errored"
503 );
504 }
505
506 #[tokio::test]
507 async fn test_handler_signals_series_completed() {
508 let state = Arc::new(Mutex::new(HistoricalState::new()));
509 let factory = HistoricalDataHandlerFactory::new(Arc::clone(&state));
510 let (cmd_tx, _cmd_rx) = tokio::sync::mpsc::channel(4);
511 let handler = factory.create(cmd_tx);
512
513 let state_clone = Arc::clone(&state);
514 let wait_handle = tokio::spawn(async move {
515 HistoricalClient::wait_for_completion(&state_clone).await;
516 });
517
518 tokio::task::yield_now().await;
519
520 handler.handle_events(TradingViewDataEvent::OnSeriesCompleted, &[]);
521
522 let timeout_result =
523 tokio::time::timeout(std::time::Duration::from_millis(50), wait_handle).await;
524 assert!(
525 timeout_result.is_ok(),
526 "handler must signal completion immediately to waiters"
527 );
528 assert!(state.lock().unwrap().completed);
529 }
530
531 #[tokio::test]
532 async fn test_handler_signals_protocol_error() {
533 let state = Arc::new(Mutex::new(HistoricalState::new()));
534 let factory = HistoricalDataHandlerFactory::new(Arc::clone(&state));
535 let (cmd_tx, _cmd_rx) = tokio::sync::mpsc::channel(4);
536 let handler = factory.create(cmd_tx);
537
538 let state_clone = Arc::clone(&state);
539 let wait_handle = tokio::spawn(async move {
540 HistoricalClient::wait_for_completion(&state_clone).await;
541 });
542
543 tokio::task::yield_now().await;
544
545 handler.handle_events(
546 TradingViewDataEvent::OnError(TradingViewError::SeriesError),
547 &[],
548 );
549
550 let timeout_result =
551 tokio::time::timeout(std::time::Duration::from_millis(50), wait_handle).await;
552 assert!(
553 timeout_result.is_ok(),
554 "handler must signal error immediately to waiters"
555 );
556 assert!(state.lock().unwrap().errored);
557 }
558
559 #[test]
560 fn test_handler_parses_chart_data_without_cloning() {
561 let state = Arc::new(Mutex::new(HistoricalState::new()));
562 let factory = HistoricalDataHandlerFactory::new(Arc::clone(&state));
563 let (cmd_tx, _cmd_rx) = tokio::sync::mpsc::channel(4);
564 let handler = factory.create(cmd_tx);
565
566 let chart_data_payload = serde_json::json!([
567 "session_id",
568 {
569 "s1": {
570 "s": [
571 { "i": 100, "v": [100.0, 105.0, 99.0, 104.0, 1000.0] },
572 { "i": 101, "v": [104.0, 106.0, 103.0, 105.5, 1200.0] }
573 ]
574 }
575 }
576 ]);
577
578 let msg_slice = chart_data_payload.as_array().unwrap();
579 handler.handle_events(TradingViewDataEvent::OnChartData, msg_slice);
580
581 let guard = state.lock().unwrap();
582 assert_eq!(guard.data.len(), 2);
583 assert_eq!(guard.total_bars, 2);
584 assert_eq!(guard.data[0].index, 100);
585 assert_eq!(guard.data[1].index, 101);
586 }
587}