Skip to main content

tradingview/
utils.rs

1use crate::{
2    Result, UserCookies,
3    live::models::{SocketMessage, SocketMessageDe},
4    models::{MarketAdjustment, SessionType},
5};
6use bon::builder;
7use iso_currency::Currency;
8use rand::{Rng, distr::Alphanumeric};
9use reqwest::{
10    Response,
11    header::{ACCEPT, COOKIE, HeaderMap, HeaderValue, ORIGIN, REFERER},
12};
13use serde::Serialize;
14use std::{collections::HashMap, sync::LazyLock};
15use tokio_tungstenite::tungstenite::protocol::Message;
16use tracing::{debug, error, warn};
17use ustr::Ustr;
18
19// ---------------------------------------------------------------------------
20// Shared HTTP client — built once, reused for all requests.
21// Enables connection pooling, DNS caching, TLS session reuse, and HTTP/2.
22// ---------------------------------------------------------------------------
23static SHARED_CLIENT: LazyLock<reqwest::Client> = LazyLock::new(|| {
24    let mut headers = HeaderMap::new();
25    headers.insert(ACCEPT, HeaderValue::from_static("application/json"));
26    headers.insert(
27        ORIGIN,
28        HeaderValue::from_static("https://www.tradingview.com"),
29    );
30    headers.insert(
31        REFERER,
32        HeaderValue::from_static("https://www.tradingview.com/"),
33    );
34
35    let mut builder = reqwest::Client::builder()
36        .default_headers(headers)
37        .https_only(true)
38        .user_agent(crate::UA);
39
40    #[cfg(feature = "rustls-tls")]
41    {
42        builder = builder.use_rustls_tls();
43    }
44    #[cfg(feature = "native-tls")]
45    {
46        builder = builder.use_native_tls();
47    }
48
49    builder.build().expect("Failed to build shared HTTP client")
50});
51
52#[macro_export]
53macro_rules! payload {
54    ($($payload:expr),*) => {
55        {
56        let payload_vec = vec![$(serde_json::Value::from($payload)),*];
57        payload_vec
58        }
59    };
60}
61
62/// Returns a clone of the shared `reqwest::Client`.
63///
64/// The client is built once at first use and reused for all subsequent calls.
65/// `reqwest::Client` is cheap to clone (it wraps an `Arc` internally), so
66/// cloning enables connection pooling, DNS caching, and TLS session reuse.
67///
68/// For authenticated requests, use [`http_client`] and add the cookie
69/// per-request via `.header(COOKIE, ...)`.
70pub fn http_client() -> reqwest::Client {
71    SHARED_CLIENT.clone()
72}
73
74/// Build a `reqwest::Client` with optional authentication cookies baked into
75/// its default headers.
76///
77/// **Deprecated for hot paths.**  Prefer [`http_client`] + per-request
78/// `.header(COOKIE, cookie_str)` to benefit from connection pooling.
79///
80/// This function is retained for backward compatibility.  When no cookie is
81/// supplied it clones the shared client (zero-cost).  When a cookie IS
82/// supplied it builds a fresh client (suboptimal — prefer per-request cookies).
83#[deprecated(
84    since = "0.2.0",
85    note = "Use http_client() and add cookies per-request via .header(COOKIE, ...) for connection pooling"
86)]
87pub fn build_request(cookie: Option<&str>) -> Result<reqwest::Client> {
88    if cookie.is_some() {
89        warn!(
90            "build_request() with cookies bypasses connection pooling; use http_client() + .header(COOKIE, ...) instead"
91        );
92        // Legacy path: build a dedicated client with cookies in default headers.
93        // This is suboptimal but maintains backward compatibility.
94        let mut headers = HeaderMap::new();
95        headers.insert(ACCEPT, HeaderValue::from_static("application/json"));
96        headers.insert(
97            ORIGIN,
98            HeaderValue::from_static("https://www.tradingview.com"),
99        );
100        headers.insert(
101            REFERER,
102            HeaderValue::from_static("https://www.tradingview.com/"),
103        );
104        if let Some(cookie) = cookie {
105            headers.insert(COOKIE, HeaderValue::from_str(cookie)?);
106        }
107
108        let mut builder = reqwest::Client::builder()
109            .default_headers(headers)
110            .https_only(true)
111            .user_agent(crate::UA);
112
113        #[cfg(feature = "rustls-tls")]
114        {
115            builder = builder.use_rustls_tls();
116        }
117        #[cfg(feature = "native-tls")]
118        {
119            builder = builder.use_native_tls();
120        }
121
122        return Ok(builder.build()?);
123    }
124
125    // No cookies: return the shared client (connection pooling enabled).
126    Ok(SHARED_CLIENT.clone())
127}
128
129pub fn gen_session_id(session_type: &str) -> String {
130    session_type.to_owned() + "_" + &gen_id()
131}
132
133#[inline]
134pub fn gen_id() -> String {
135    let mut rng = rand::rng();
136    let buf: [u8; 12] = std::array::from_fn(|_| rng.sample(Alphanumeric));
137    // SAFETY: `Alphanumeric` samples only ASCII bytes (0-9, A-Z, a-z),
138    // which are always valid UTF-8.
139    let s = core::str::from_utf8(&buf).expect("Alphanumeric produces only ASCII");
140    s.to_owned()
141}
142
143/// Extract properly-framed `~m~<len>~m~~h~<counter>` heartbeat echoes from
144/// a raw TradingView protocol text frame.
145///
146/// Each returned string is a complete, properly framed heartbeat echo ready
147/// to be sent back to the server. The framing length is **computed from the
148/// actual `~h~<counter>` payload size**, not blindly echoed from the server
149/// frame. This guarantees correctness even if the server sends a malformed
150/// length.
151pub fn extract_heartbeat_echoes(raw: &str) -> Vec<String> {
152    parse_packet(raw)
153        .into_iter()
154        .filter_map(|msg| msg.heartbeat_echo())
155        .collect()
156}
157
158fn classify_json_value(value: serde_json::Value) -> SocketMessage<SocketMessageDe> {
159    match value {
160        serde_json::Value::Object(mut map) => {
161            if matches!(map.get("m"), Some(serde_json::Value::String(_)))
162                && matches!(map.get("p"), Some(serde_json::Value::Array(_)))
163            {
164                let m_str = match map.remove("m").unwrap() {
165                    serde_json::Value::String(s) => s,
166                    _ => unreachable!(),
167                };
168                let p_vec = match map.remove("p").unwrap() {
169                    serde_json::Value::Array(a) => a,
170                    _ => unreachable!(),
171                };
172                let t = map.get("t").and_then(|v| v.as_u64()).unwrap_or(0);
173                let t_ms = map.get("t_ms").and_then(|v| v.as_u64()).unwrap_or(0);
174                return SocketMessage::SocketMessage(SocketMessageDe {
175                    m: Ustr::from(&m_str),
176                    p: p_vec,
177                    t,
178                    t_ms,
179                });
180            }
181
182            if map.contains_key("session_id") && map.contains_key("timestamp") {
183                let obj = serde_json::Value::Object(map);
184                if let Ok(info) =
185                    serde_json::from_value::<crate::live::models::SocketServerInfo>(obj.clone())
186                {
187                    return SocketMessage::SocketServerInfo(info);
188                }
189                return SocketMessage::Other(obj);
190            }
191
192            SocketMessage::Other(serde_json::Value::Object(map))
193        }
194        other => SocketMessage::Other(other),
195    }
196}
197
198pub fn parse_packet(message: &str) -> Vec<SocketMessage<SocketMessageDe>> {
199    if message.is_empty() {
200        return vec![];
201    }
202
203    let bytes = message.as_bytes();
204    let mut pos = 0;
205    let mut packets = Vec::new();
206
207    while pos < bytes.len() {
208        if bytes[pos..].starts_with(b"~m~") {
209            let header_start = pos + 3;
210            let mut len_end = header_start;
211            while len_end < bytes.len() && bytes[len_end].is_ascii_digit() {
212                len_end += 1;
213            }
214
215            if len_end > header_start && bytes[len_end..].starts_with(b"~m~") {
216                let len_str = &message[header_start..len_end];
217                if let Ok(payload_len) = len_str.parse::<usize>() {
218                    let payload_start = len_end + 3;
219                    let slice = &message[payload_start..];
220                    let mut utf16_count = 0;
221                    let mut actual_bytes = slice.len();
222                    let mut found = false;
223
224                    for (byte_offset, ch) in slice.char_indices() {
225                        if utf16_count >= payload_len {
226                            actual_bytes = byte_offset;
227                            found = true;
228                            break;
229                        }
230                        utf16_count += ch.len_utf16();
231                    }
232
233                    if !found && utf16_count <= payload_len {
234                        actual_bytes = slice.len();
235                    }
236
237                    if actual_bytes > 0 {
238                        let payload = &slice[..actual_bytes];
239                        if payload.starts_with("~h~") {
240                            let mut hb_len = 3;
241                            while hb_len < payload.len()
242                                && payload.as_bytes()[hb_len].is_ascii_digit()
243                            {
244                                hb_len += 1;
245                            }
246                            if hb_len > 3 {
247                                if let Ok(counter) = payload[3..hb_len].parse::<u64>() {
248                                    packets.push(SocketMessage::Heartbeat(counter));
249                                }
250                                pos = payload_start + hb_len;
251                                continue;
252                            }
253                        }
254
255                        match serde_json::from_str::<serde_json::Value>(payload) {
256                            Ok(val) => {
257                                packets.push(classify_json_value(val));
258                            }
259                            Err(err) => {
260                                if err.is_syntax() {
261                                    error!("error parsing packet, invalid JSON: {}", err);
262                                } else {
263                                    error!("error parsing packet: {}", err);
264                                }
265                                packets.push(SocketMessage::Unknown(payload.to_string()));
266                            }
267                        }
268                    }
269
270                    pos = payload_start + actual_bytes;
271                    continue;
272                }
273            }
274
275            pos += 3;
276        } else if bytes[pos..].starts_with(b"~h~") {
277            let counter_start = pos + 3;
278            let mut counter_end = counter_start;
279            while counter_end < bytes.len() && bytes[counter_end].is_ascii_digit() {
280                counter_end += 1;
281            }
282
283            if counter_end > counter_start {
284                if let Ok(counter) = message[counter_start..counter_end].parse::<u64>() {
285                    packets.push(SocketMessage::Heartbeat(counter));
286                }
287                pos = counter_end;
288            } else {
289                pos += 3;
290            }
291        } else {
292            pos += 1;
293        }
294    }
295
296    packets
297}
298
299pub fn format_packet<T: Serialize>(packet: T) -> Result<Message> {
300    let json_string = serde_json::to_string(&packet)?;
301    let utf16_len = json_string.encode_utf16().count();
302    let formatted_message = format!("~m~{}~m~{}", utf16_len, json_string);
303    debug!("Formatted packet: {}", formatted_message);
304    Ok(Message::Text(formatted_message.into()))
305}
306
307#[builder]
308pub fn symbol_init(
309    instrument: &str, // The instrument symbol, e.g., "HOSE:FPT"
310    adjustment: Option<MarketAdjustment>,
311    currency: Option<Currency>,
312    session_type: Option<SessionType>,
313    replay: Option<&str>,
314) -> Result<String> {
315    let mut symbol_init: HashMap<Ustr, Ustr> = HashMap::new();
316    if let Some(s) = replay {
317        symbol_init.insert(Ustr::from("replay"), Ustr::from(s));
318    }
319    if let Some(a) = adjustment {
320        symbol_init.insert(Ustr::from("adjustment"), Ustr::from(&a.to_string()));
321    }
322    symbol_init.insert(Ustr::from("symbol"), Ustr::from(instrument));
323    if let Some(c) = currency {
324        symbol_init.insert(Ustr::from("currency-id"), Ustr::from(c.code()));
325    }
326    if let Some(s) = session_type {
327        symbol_init.insert(Ustr::from("session"), Ustr::from(&s.to_string()));
328    }
329    let symbol_init_json = serde_json::to_value(&symbol_init)?;
330    Ok(format!("={symbol_init_json}"))
331}
332
333pub async fn get(
334    client: Option<&UserCookies>,
335    url: &str,
336    queries: &[(&str, &str)],
337) -> Result<Response> {
338    let mut req = SHARED_CLIENT.get(url);
339    if let Some(c) = client {
340        let cookie = format!(
341            "sessionid={}; sessionid_sign={}; device_t={};",
342            c.session, c.session_signature, c.device_token
343        );
344        req = req.header(COOKIE, &cookie);
345    }
346    let response = req.query(queries).send().await?;
347    Ok(response)
348}
349
350#[cfg(test)]
351mod tests {
352    use serde_json::{Value, json};
353
354    use crate::{
355        live::models,
356        models::{MarketAdjustment, SessionType},
357        utils::*,
358    };
359
360    // ──────────────────────────────────────────────────────────────────
361    // parse_packet — basic smoke tests
362    // ──────────────────────────────────────────────────────────────────
363
364    #[test]
365    fn parse_packet_from_file() {
366        let current_dir = std::env::current_dir().unwrap().display().to_string();
367        let messages =
368            std::fs::read_to_string(format!("{current_dir}/tests/data/socket_messages.txt"))
369                .unwrap();
370        let result = parse_packet(messages.as_str());
371        assert_eq!(result.len(), 42);
372    }
373
374    #[test]
375    fn parse_packet_empty_returns_empty() {
376        assert!(parse_packet("").is_empty());
377    }
378
379    #[test]
380    fn parse_packet_only_heartbeats_returns_empty() {
381        assert!(parse_packet("~h~~h~~h~").is_empty());
382    }
383
384    #[test]
385    fn parse_packet_single_valid_message() {
386        // {"m":"test","p":["hello"]} = 25 bytes
387        let result = parse_packet(r#"~m~25~m~{"m":"test","p":["hello"]}"#);
388        assert_eq!(result.len(), 1);
389    }
390
391    #[test]
392    fn parse_packet_multiple_messages() {
393        let input = concat!(
394            r#"~m~23~m~{"m":"test1","p":["a"]}"#,
395            r#"~m~23~m~{"m":"test2","p":["b"]}"#,
396            r#"~m~23~m~{"m":"test3","p":["c"]}"#,
397        );
398        assert_eq!(parse_packet(input).len(), 3);
399    }
400
401    #[test]
402    fn parse_packet_skips_interleaved_heartbeats() {
403        let input = "~h~~m~15~m~{\"key\":\"value\"}~h~~m~5~m~12345";
404        assert_eq!(parse_packet(input).len(), 2);
405    }
406
407    // ──────────────────────────────────────────────────────────────────
408    // parse_packet — type verification (deserialization correctness)
409    // ──────────────────────────────────────────────────────────────────
410
411    #[test]
412    fn parse_packet_deserializes_socket_message_de() {
413        // A well-formed SocketMessageDe with m, p, t, t_ms fields.
414        let payload = serde_json::json!({
415            "m": "timescale_update",
416            "p": [{"sds_5": {"s": [{"i": 0, "v": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0]}]}}],
417            "t": 1685633880_u64,
418            "t_ms": 1685633880000_u64,
419        });
420        let payload_str = payload.to_string();
421        let packet = format!("~m~{}~m~{}", payload_str.len(), payload_str);
422        let result = parse_packet(&packet);
423        assert_eq!(result.len(), 1);
424
425        match &result[0] {
426            SocketMessage::SocketMessage(de) => {
427                assert_eq!(de.m.as_str(), "timescale_update");
428                assert_eq!(de.p.len(), 1);
429                assert_eq!(de.t, 1685633880);
430                assert_eq!(de.t_ms, 1685633880000);
431            }
432            other => panic!("expected SocketMessageDe, got {other:?}"),
433        }
434    }
435
436    #[test]
437    fn parse_packet_deserializes_socket_server_info() {
438        // SocketServerInfo is matched by the untagged enum before SocketMessageDe.
439        // `#[serde(rename_all = "camelCase")]` applies to most fields;
440        // `session_id`, `studies_metadata_hash`, and `auth_scheme_vsn` have
441        // explicit `#[serde(rename)]` overrides.
442        let info = serde_json::json!({
443            "session_id": "cs_abc123",
444            "timestamp": 1685633880_i64,
445            "timestampMs": 1685633880000_i64,
446            "release": "v24.10",
447            "studies_metadata_hash": "hash123",
448            "auth_scheme_vsn": 2_i64,
449            "protocol": "json",
450            "via": "direct",
451            "javastudies": ["study1", "study2"],
452        });
453        let payload_str = info.to_string();
454        let packet = format!("~m~{}~m~{}", payload_str.len(), payload_str);
455        let result = parse_packet(&packet);
456        assert_eq!(result.len(), 1);
457
458        match &result[0] {
459            SocketMessage::SocketServerInfo(si) => {
460                assert_eq!(si.session_id.as_str(), "cs_abc123");
461                assert_eq!(si.timestamp, 1685633880);
462                assert_eq!(si.release.as_str(), "v24.10");
463                assert_eq!(si.sjavastudies.len(), 2);
464            }
465            other => panic!("expected SocketServerInfo, got {other:?}"),
466        }
467    }
468
469    #[test]
470    fn parse_packet_other_variant_for_unknown_json_structure() {
471        // Valid JSON that doesn't match SocketServerInfo or SocketMessageDe
472        // should fall into the Other(Value) variant.
473        let payload = serde_json::json!({"unexpected_field": "strange", "count": 42});
474        let payload_str = payload.to_string();
475        let packet = format!("~m~{}~m~{}", payload_str.len(), payload_str);
476        let result = parse_packet(&packet);
477        assert_eq!(result.len(), 1);
478
479        match &result[0] {
480            SocketMessage::Other(v) => {
481                assert_eq!(v["unexpected_field"], "strange");
482                assert_eq!(v["count"], 42);
483            }
484            other => panic!("expected Other(Value), got {other:?}"),
485        }
486    }
487
488    #[test]
489    fn parse_packet_unknown_variant_for_invalid_json() {
490        // Non-JSON text should produce an Unknown variant.
491        let input = "~m~11~m~not_a_json!";
492        let result = parse_packet(input);
493        assert_eq!(result.len(), 1);
494
495        match &result[0] {
496            SocketMessage::Unknown(s) => {
497                assert_eq!(s.as_str(), "not_a_json!");
498            }
499            other => panic!("expected Unknown, got {other:?}"),
500        }
501    }
502
503    // ──────────────────────────────────────────────────────────────────
504    // parse_packet — edge case: non-UTF8 payload
505    // ──────────────────────────────────────────────────────────────────
506
507    #[test]
508    fn parse_packet_non_utf8_payload_becomes_unknown() {
509        // Build a packet with non-UTF8 bytes.  The byte sequence 0xFF is
510        // never valid UTF-8, so the parser falls back to lossy conversion.
511        let non_utf8_payload = vec![0xFF, 0xFE, 0xFD, b'a', b'b', b'c'];
512        let payload_len = non_utf8_payload.len();
513        let mut packet = format!("~m~{}~m~", payload_len);
514        packet.push_str(
515            // SAFETY: we're constructing a string-like frame; the payload
516            // bytes are appended directly for test purposes.
517            core::str::from_utf8(&non_utf8_payload).unwrap_or(""),
518        );
519        // For a fully non-UTF8 payload, we construct raw bytes manually.
520        let raw = [b"~m~6~m~" as &[u8], &[0xFF, 0xFE, 0xFD, b'a', b'b', b'c']].concat();
521        let lossy = String::from_utf8_lossy(&raw);
522        let result = parse_packet(&lossy);
523        // Should produce an Unknown variant with the lossy text.
524        assert_eq!(result.len(), 1);
525        match &result[0] {
526            SocketMessage::Unknown(_) => { /* expected */ }
527            other => panic!("expected Unknown for non-UTF8 payload, got {other:?}"),
528        }
529    }
530
531    // ──────────────────────────────────────────────────────────────────
532    // parse_packet — edge case: length field variants
533    // ──────────────────────────────────────────────────────────────────
534
535    #[test]
536    fn parse_packet_length_zero_is_skipped() {
537        let result = parse_packet("~m~0~m~");
538        assert!(result.is_empty());
539    }
540
541    #[test]
542    fn parse_packet_leading_zeros_in_length() {
543        // "~m~005~m~hello" — length 5, payload "hello"
544        let result = parse_packet("~m~005~m~hello");
545        assert_eq!(result.len(), 1);
546    }
547
548    #[test]
549    fn parse_packet_negative_like_length_is_skipped() {
550        // "~m~-1~m~xxx" — '-' is not a digit, so not a valid length header.
551        // Protocol invariant: malformed frames never panic.
552        let result = parse_packet("~m~-1~m~xxx");
553        assert!(result.is_empty());
554    }
555
556    #[test]
557    fn parse_packet_truncated_before_length_delimiter_does_not_panic() {
558        let _ = parse_packet("~m~999");
559    }
560
561    #[test]
562    fn parse_packet_length_exceeds_remaining_bytes_clamped() {
563        // Length declares 999 bytes but only "short" is available.
564        // The parser clamps at the end of input and tries to parse "short".
565        let result = parse_packet("~m~999~m~short");
566        assert!(result.len() <= 1);
567    }
568
569    #[test]
570    fn parse_packet_payload_contains_tilde_m_delimiter_substring() {
571        // If the payload itself contains "~m~", it must be included in the
572        // payload, not treated as a new delimiter (since we use length-based
573        // extraction, not delimiter scanning).
574        let json_payload = serde_json::json!({"note": "look for ~m~ in payload"});
575        let payload_str = json_payload.to_string();
576        let packet = format!("~m~{}~m~{}", payload_str.len(), payload_str);
577        let result = parse_packet(&packet);
578        assert_eq!(result.len(), 1);
579        match &result[0] {
580            SocketMessage::Other(v) => {
581                assert_eq!(v["note"], "look for ~m~ in payload");
582            }
583            other => panic!("expected a parsed message, got {other:?}"),
584        }
585    }
586
587    #[test]
588    fn parse_packet_random_garbage_never_panics() {
589        let garbage = [
590            "",
591            "~~~",
592            "~m~",
593            "~m~abc~m~",
594            "~m~-1~m~",
595            "not a packet at all",
596            "~m~5~m~hello~m~3~m~bye",
597            "\x00\x01\x02\x03",
598            "~m~999999999999999999999999~m~", // huge length
599            "~h~~h~~m~~m~~h~",
600            "~m~5~m~",
601            "~m~~m~5~m~hello",
602        ];
603        for input in &garbage {
604            let _result = parse_packet(input);
605            // No panic => pass
606        }
607    }
608
609    // ──────────────────────────────────────────────────────────────────
610    // parse_packet — heartbeat / ping edge cases
611    // ──────────────────────────────────────────────────────────────────
612
613    #[test]
614    fn parse_packet_ping_digits_before_frame() {
615        // ~h~9999999999~m~5~m~hello => heartbeat + frame
616        let result = parse_packet("~h~9999999999~m~5~m~hello");
617        assert_eq!(result.len(), 2);
618        assert_eq!(result[0], SocketMessage::Heartbeat(9999999999));
619        assert_eq!(result[1], SocketMessage::Unknown("hello".to_string()));
620    }
621
622    #[test]
623    fn parse_packet_consecutive_pings_only() {
624        // 10 repetitions of "~h~9999999999" — 10 typed Heartbeat packets.
625        let input = "~h~9999999999".repeat(10);
626        let result = parse_packet(&input);
627        assert_eq!(result.len(), 10);
628        for msg in result {
629            assert_eq!(msg, SocketMessage::Heartbeat(9999999999));
630        }
631    }
632
633    #[test]
634    fn parse_packet_mixed_pings_and_frames() {
635        let input = concat!(
636            "~h~",                                     // bare heartbeat without counter (skipped)
637            "~m~23~m~{\"m\":\"test1\",\"p\":[\"a\"]}", // valid frame
638            "~h~9999999999",                           // ping digits
639            "~m~23~m~{\"m\":\"test2\",\"p\":[\"b\"]}", // valid frame
640            "~h~88888888",                             // more ping digits
641        );
642        let result = parse_packet(input);
643        assert_eq!(result.len(), 4);
644        assert!(matches!(&result[0], SocketMessage::SocketMessage(_)));
645        assert_eq!(result[1], SocketMessage::Heartbeat(9999999999));
646        assert!(matches!(&result[2], SocketMessage::SocketMessage(_)));
647        assert_eq!(result[3], SocketMessage::Heartbeat(88888888));
648    }
649
650    #[test]
651    fn parse_packet_ping_digits_resembling_frame_prefix() {
652        let result = parse_packet("~h~10~m~hello");
653        assert_eq!(result.len(), 1);
654        assert_eq!(result[0], SocketMessage::Heartbeat(10));
655    }
656
657    #[test]
658    fn parse_packet_large_ping_no_panic() {
659        let ping = format!("~h~{}", "9".repeat(100));
660        let _result = parse_packet(&ping);
661    }
662
663    // ──────────────────────────────────────────────────────────────────
664    // format_packet → parse_packet round-trip
665    // ──────────────────────────────────────────────────────────────────
666
667    #[test]
668    fn roundtrip_format_then_parse_basic() {
669        let msg = serde_json::json!({"m": "test_method", "p": [{"key": "value"}]});
670        let formatted = super::format_packet(&msg).expect("format succeeds");
671        let text = match &formatted {
672            tokio_tungstenite::tungstenite::protocol::Message::Text(t) => t.as_str(),
673            _ => panic!("expected text message"),
674        };
675        let parsed = parse_packet(text);
676        assert!(!parsed.is_empty());
677    }
678
679    #[test]
680    fn roundtrip_format_then_parse_socket_message_ser() {
681        let msg = models::SocketMessageSer::new(
682            "qsd",
683            serde_json::json!([{
684                "n": "AAPL",
685                "v": {"bid": 150.25, "ask": 150.30, "lp": 150.28}
686            }]),
687        );
688        let formatted = msg.to_message().expect("format succeeds");
689        let text = match &formatted {
690            tokio_tungstenite::tungstenite::protocol::Message::Text(t) => t.as_str(),
691            _ => panic!("expected text message"),
692        };
693        let parsed = parse_packet(text);
694        assert_eq!(parsed.len(), 1);
695        match &parsed[0] {
696            SocketMessage::SocketMessage(de) => {
697                assert_eq!(de.m.as_str(), "qsd");
698                assert_eq!(de.p[0]["n"], "AAPL");
699            }
700            other => panic!(
701                "expected SocketMessage(SocketMessageDe) for SocketMessageSer round-trip, got {other:?}"
702            ),
703        }
704    }
705
706    #[test]
707    fn roundtrip_multiple_formatted_packets() {
708        let msgs: Vec<models::SocketMessageSer> = (0..5)
709            .map(|i| {
710                models::SocketMessageSer::new(
711                    format!("method_{i}"),
712                    serde_json::json!([{"index": i}]),
713                )
714            })
715            .collect();
716
717        // Concatenate formatted packets.
718        let mut combined = String::new();
719        for msg in &msgs {
720            let fmt = msg.to_message().expect("format succeeds");
721            if let tokio_tungstenite::tungstenite::protocol::Message::Text(t) = &fmt {
722                combined.push_str(t.as_str());
723            }
724        }
725
726        let parsed = parse_packet(&combined);
727        assert_eq!(parsed.len(), msgs.len());
728        for (i, p) in parsed.iter().enumerate() {
729            match p {
730                SocketMessage::SocketMessage(de) => {
731                    assert_eq!(de.m.as_str(), format!("method_{i}"));
732                    assert_eq!(de.p[0]["index"], i);
733                }
734                other => {
735                    panic!("expected SocketMessage(SocketMessageDe) at index {i}, got {other:?}")
736                }
737            }
738        }
739    }
740
741    #[test]
742    fn roundtrip_with_unicode_payload() {
743        // Payload containing Unicode characters.
744        let msg = serde_json::json!({
745            "m": "study_data",
746            "p": [{"name": "📈 Moving Average", "currency": "€"}]
747        });
748        let formatted = super::format_packet(&msg).expect("format succeeds");
749        let text = match &formatted {
750            tokio_tungstenite::tungstenite::protocol::Message::Text(t) => t.as_str(),
751            _ => panic!("expected text message"),
752        };
753        let parsed = parse_packet(text);
754        assert_eq!(parsed.len(), 1);
755    }
756
757    #[test]
758    fn roundtrip_full_socket_message_de() {
759        // A complete SocketMessageDe includes m, p, t, and t_ms.
760        // The untagged enum should deserialize this as SocketMessage(SocketMessageDe).
761        let payload = serde_json::json!({
762            "m": "timescale_update",
763            "p": [{"sds_5": {"s": [{"i": 0, "v": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0]}]}}],
764            "t": 1685633880_u64,
765            "t_ms": 1685633880000_u64,
766        });
767        let formatted = super::format_packet(&payload).expect("format succeeds");
768        let text = match &formatted {
769            tokio_tungstenite::tungstenite::protocol::Message::Text(t) => t.as_str(),
770            _ => panic!("expected text message"),
771        };
772        let parsed = parse_packet(text);
773        assert_eq!(parsed.len(), 1);
774        match &parsed[0] {
775            SocketMessage::SocketMessage(de) => {
776                assert_eq!(de.m.as_str(), "timescale_update");
777                assert_eq!(de.t, 1685633880);
778                assert_eq!(de.t_ms, 1685633880000);
779            }
780            other => panic!("expected SocketMessage(SocketMessageDe), got {other:?}"),
781        }
782    }
783
784    // ──────────────────────────────────────────────────────────────────
785    // gen_session_id / gen_id
786    // ──────────────────────────────────────────────────────────────────
787
788    #[test]
789    fn gen_session_id_produces_correct_format() {
790        let session_type = "qc";
791        let session_id = gen_session_id(session_type);
792        // 2 (session_type) + 1 (_) + 12 (random alphanumeric chars)
793        assert_eq!(session_id.len(), 15);
794        assert!(session_id.starts_with(session_type));
795        assert!(session_id.as_bytes()[2] == b'_');
796    }
797
798    #[test]
799    fn gen_id_is_unique_across_many_calls() {
800        let ids: Vec<String> = (0..100).map(|_| gen_id()).collect();
801        let unique: std::collections::HashSet<_> = ids.iter().collect();
802        assert_eq!(unique.len(), 100, "gen_id should produce unique values");
803    }
804
805    #[test]
806    fn gen_id_produces_only_alphanumeric() {
807        for _ in 0..50 {
808            let id = gen_id();
809            assert!(
810                id.chars().all(|c| c.is_ascii_alphanumeric()),
811                "gen_id produced non-alphanumeric: {id}"
812            );
813        }
814    }
815
816    #[test]
817    fn two_heartbeats_produce_exactly_two_typed_heartbeats_and_echoes() {
818        let raw = "~m~5~m~~h~42~m~25~m~{\"m\":\"du\",\"p\":[\"cs_xxx\"]}~m~5~m~~h~43";
819        let parsed = parse_packet(raw);
820        assert_eq!(parsed.len(), 3);
821        assert_eq!(parsed[0], SocketMessage::Heartbeat(42));
822        assert!(matches!(parsed[1], SocketMessage::SocketMessage(_)));
823        assert_eq!(parsed[2], SocketMessage::Heartbeat(43));
824
825        let echoes = extract_heartbeat_echoes(raw);
826        assert_eq!(echoes, vec!["~m~5~m~~h~42", "~m~5~m~~h~43"]);
827    }
828
829    #[test]
830    fn embedded_tilde_m_and_tilde_h_survive() {
831        let payload = serde_json::json!({
832            "m": "quote",
833            "p": ["embedded ~m~5~m~ and ~h~ text"]
834        });
835        let formatted = super::format_packet(&payload).expect("format succeeds");
836        let text = match &formatted {
837            tokio_tungstenite::tungstenite::protocol::Message::Text(t) => t.as_str(),
838            _ => panic!("expected text"),
839        };
840        let parsed = parse_packet(text);
841        assert_eq!(parsed.len(), 1);
842        match &parsed[0] {
843            SocketMessage::SocketMessage(de) => {
844                assert_eq!(de.m.as_str(), "quote");
845                assert_eq!(de.p[0], "embedded ~m~5~m~ and ~h~ text");
846            }
847            other => panic!("expected SocketMessage, got {other:?}"),
848        }
849    }
850
851    #[test]
852    fn unicode_byte_lengths_roundtrip() {
853        let payload = serde_json::json!({
854            "m": "study_data",
855            "p": [{"name": "📈 Moving Average", "currency": "€"}]
856        });
857        let formatted = super::format_packet(&payload).expect("format succeeds");
858        let text = match &formatted {
859            tokio_tungstenite::tungstenite::protocol::Message::Text(t) => t.as_str(),
860            _ => panic!("expected text"),
861        };
862        let parsed = parse_packet(text);
863        assert_eq!(parsed.len(), 1);
864        match &parsed[0] {
865            SocketMessage::SocketMessage(de) => {
866                assert_eq!(de.m.as_str(), "study_data");
867                assert_eq!(de.p[0]["name"], "📈 Moving Average");
868                assert_eq!(de.p[0]["currency"], "€");
869            }
870            other => panic!("expected SocketMessage, got {other:?}"),
871        }
872    }
873
874    #[test]
875    fn m_p_maps_to_socket_message_de() {
876        let raw = r#"~m~23~m~{"m":"test","p":["a"]}"#;
877        let parsed = parse_packet(raw);
878        assert_eq!(parsed.len(), 1);
879        match &parsed[0] {
880            SocketMessage::SocketMessage(de) => {
881                assert_eq!(de.m.as_str(), "test");
882                assert_eq!(de.p[0], "a");
883                assert_eq!(de.t, 0);
884                assert_eq!(de.t_ms, 0);
885            }
886            other => panic!("expected SocketMessage(SocketMessageDe), got {other:?}"),
887        }
888    }
889
890    // ──────────────────────────────────────────────────────────────────
891    // symbol_init
892    // ──────────────────────────────────────────────────────────────────
893
894    #[test]
895    fn symbol_init_minimal() {
896        let test1 = symbol_init().instrument("NSE:NIFTY").call();
897        assert!(test1.is_ok());
898        assert_eq!(test1.unwrap(), r#"={"symbol":"NSE:NIFTY"}"#.to_string());
899    }
900
901    #[test]
902    fn symbol_init_all_fields() {
903        let result = symbol_init()
904            .instrument("HOSE:FPT")
905            .adjustment(MarketAdjustment::Dividends)
906            .currency(Currency::USD)
907            .session_type(SessionType::Extended)
908            .replay("aaaaaaaaaaaa")
909            .call();
910        assert!(result.is_ok());
911        let json_str = result.unwrap().replace('=', "");
912        let parsed: Value = serde_json::from_str(&json_str).unwrap();
913        let expected = json!({
914            "adjustment": "dividends",
915            "currency-id": "USD",
916            "replay": "aaaaaaaaaaaa",
917            "session": "extended",
918            "symbol": "HOSE:FPT"
919        });
920        assert_eq!(parsed, expected);
921    }
922
923    // ──────────────────────────────────────────────────────────────────
924    // Shared HTTP client tests
925    // ──────────────────────────────────────────────────────────────────
926
927    #[test]
928    fn http_client_is_reusable() {
929        let c1 = http_client();
930        let c2 = http_client();
931        let c3 = http_client();
932        drop(c1);
933        drop(c2);
934        drop(c3);
935    }
936
937    #[test]
938    fn http_client_supports_many_clones() {
939        let clients: Vec<_> = (0..100).map(|_| http_client()).collect();
940        assert_eq!(clients.len(), 100);
941    }
942
943    #[test]
944    #[allow(deprecated)]
945    fn build_request_no_cookie_returns_usable_client() {
946        let client = build_request(None).expect("build_request without cookie");
947        drop(client);
948    }
949
950    #[test]
951    fn cookie_format_is_correct() {
952        let cookies = UserCookies {
953            session: "abc123".into(),
954            session_signature: "sig456".into(),
955            device_token: "dev789".into(),
956            ..Default::default()
957        };
958        let cookie = format!(
959            "sessionid={}; sessionid_sign={}; device_t={};",
960            cookies.session, cookies.session_signature, cookies.device_token
961        );
962        assert_eq!(
963            cookie,
964            "sessionid=abc123; sessionid_sign=sig456; device_t=dev789;"
965        );
966    }
967
968    #[test]
969    fn deprecated_build_request_with_cookie_still_works() {
970        #[allow(deprecated)]
971        let client =
972            build_request(Some("sessionid=test; sessionid_sign=sig;")).expect("with cookie");
973        drop(client);
974    }
975
976    // ──────────────────────────────────────────────────────────────────
977    // extract_heartbeat_echoes
978    // ──────────────────────────────────────────────────────────────────
979
980    #[test]
981    fn extract_heartbeat_single() {
982        let raw = "~m~25~m~{\"m\":\"du\",\"p\":[\"cs_xxx\"]}~m~5~m~~h~42";
983        let echoes = super::extract_heartbeat_echoes(raw);
984        assert_eq!(echoes, vec!["~m~5~m~~h~42"]);
985    }
986
987    #[test]
988    fn extract_heartbeat_multiple() {
989        let raw = "~m~5~m~~h~42~m~25~m~{\"m\":\"du\",\"p\":[\"cs_xxx\"]}~m~5~m~~h~43";
990        let echoes = super::extract_heartbeat_echoes(raw);
991        assert_eq!(echoes, vec!["~m~5~m~~h~42", "~m~5~m~~h~43"]);
992    }
993
994    #[test]
995    fn extract_heartbeat_none() {
996        let raw = "~m~25~m~{\"m\":\"du\",\"p\":[\"cs_xxx\"]}";
997        let echoes = super::extract_heartbeat_echoes(raw);
998        assert!(echoes.is_empty());
999    }
1000
1001    #[test]
1002    fn extract_heartbeat_interleaved() {
1003        // "~h~99" = 5 bytes → ~m~5~m~, "~h~100" = 6 bytes → ~m~6~m~
1004        let raw = "~m~5~m~~h~99~m~25~m~{\"m\":\"qsd\",\"p\":[\"qs_xxx\"]}~m~6~m~~h~100";
1005        let echoes = super::extract_heartbeat_echoes(raw);
1006        assert_eq!(echoes, vec!["~m~5~m~~h~99", "~m~6~m~~h~100"]);
1007    }
1008
1009    #[test]
1010    fn extract_heartbeat_standalone() {
1011        let raw = "~m~5~m~~h~42";
1012        let echoes = super::extract_heartbeat_echoes(raw);
1013        assert_eq!(echoes, vec!["~m~5~m~~h~42"]);
1014    }
1015
1016    #[test]
1017    fn extract_heartbeat_empty_string() {
1018        let echoes = super::extract_heartbeat_echoes("");
1019        assert!(echoes.is_empty());
1020    }
1021
1022    #[test]
1023    fn extract_heartbeat_consecutive() {
1024        let raw = "~m~5~m~~h~10~m~5~m~~h~20";
1025        let echoes = super::extract_heartbeat_echoes(raw);
1026        assert_eq!(echoes, vec!["~m~5~m~~h~10", "~m~5~m~~h~20"]);
1027    }
1028
1029    #[test]
1030    fn extract_heartbeat_properly_framed() {
1031        // Server sends a standard heartbeat: ~m~5~m~~h~1
1032        // "~h~1" = 4 bytes → correct framing should be ~m~4~m~
1033        let raw = "~m~5~m~~h~1";
1034        let echoes = super::extract_heartbeat_echoes(raw);
1035        assert_eq!(echoes, vec!["~m~4~m~~h~1"]);
1036    }
1037
1038    #[test]
1039    fn extract_heartbeat_variable_length_counter() {
1040        // Counter "42": "~h~42" = 5 bytes → ~m~5~m~
1041        let raw = "~m~5~m~~h~42";
1042        let echoes = super::extract_heartbeat_echoes(raw);
1043        assert_eq!(echoes, vec!["~m~5~m~~h~42"]);
1044
1045        // Counter "123": "~h~123" = 6 bytes → ~m~6~m~
1046        let raw = "~m~6~m~~h~123";
1047        let echoes = super::extract_heartbeat_echoes(raw);
1048        assert_eq!(echoes, vec!["~m~6~m~~h~123"]);
1049
1050        // Counter "9999": "~h~9999" = 7 bytes → ~m~7~m~
1051        let raw = "~m~7~m~~h~9999";
1052        let echoes = super::extract_heartbeat_echoes(raw);
1053        assert_eq!(echoes, vec!["~m~7~m~~h~9999"]);
1054    }
1055
1056    #[test]
1057    fn extract_heartbeat_self_heals_bad_length() {
1058        // Server sends malformed frame: claims length 9 but actual "~h~42" = 5
1059        // We compute correct length from payload, not trusting server's length.
1060        let raw = "~m~9~m~~h~42";
1061        let echoes = super::extract_heartbeat_echoes(raw);
1062        assert_eq!(echoes, vec!["~m~5~m~~h~42"]);
1063    }
1064
1065    #[test]
1066    fn extract_heartbeat_matches_exact_not_partial() {
1067        // A JSON message containing "~h~" as data, not a real heartbeat.
1068        // Should NOT be extracted as a heartbeat.
1069        let raw = "~m~30~m~{\"m\":\"set\",\"p\":[\"~h~\"]}";
1070        let echoes = super::extract_heartbeat_echoes(raw);
1071        assert!(echoes.is_empty());
1072    }
1073
1074    #[test]
1075    fn test_parse_packet_utf16_code_units_vietnamese() {
1076        let json_val = serde_json::json!({
1077            "m": "symbol_resolved",
1078            "p": [
1079                "sds_sym_1",
1080                {
1081                    "name": "HOSE:FPT",
1082                    "local_description": "CÔNG TY CỔ PHẦN FPT"
1083                }
1084            ]
1085        });
1086        let json_str = json_val.to_string();
1087        let utf16_len = json_str.encode_utf16().count();
1088        let packet = format!("~m~{}~m~{}", utf16_len, json_str);
1089
1090        let result = parse_packet(&packet);
1091        assert_eq!(result.len(), 1);
1092        match &result[0] {
1093            SocketMessage::SocketMessage(de) => {
1094                assert_eq!(de.m, "symbol_resolved");
1095                assert_eq!(de.p.len(), 2);
1096            }
1097            other => panic!("expected SocketMessage, got {other:?}"),
1098        }
1099    }
1100}