1use std::{
27 collections::VecDeque,
28 fmt::Debug,
29 sync::{
30 Arc,
31 atomic::{AtomicBool, Ordering},
32 },
33};
34
35use nautilus_common::live::dst::time;
36use nautilus_core::{
37 AtomicTime,
38 string::secret::{REDACTED, SecretString},
39};
40pub use nautilus_live::book::snapshot::SnapshotGate;
41use nautilus_model::identifiers::ClientOrderId;
42use nautilus_network::{
43 RECONNECTED,
44 error::SendError,
45 retry::{RetryError, RetryManager, create_websocket_retry_manager},
46 websocket::{AuthTracker, SubscriptionState, TEXT_PING, TEXT_PONG, WebSocketClient},
47};
48use serde_json::{Map, Value};
49use tokio_tungstenite::tungstenite::Message;
50use tokio_util::sync::CancellationToken;
51use ustr::Ustr;
52
53use super::{
54 enums::{OKXSubscriptionEvent, OKXWsChannel, OKXWsOperation},
55 error::OKXWsError,
56 messages::{
57 OKXOrderMsg, OKXSubscription, OKXSubscriptionArg, OKXWebSocketArg, OKXWebSocketError,
58 OKXWsFrame, OKXWsMessage,
59 },
60 subscription::{topic_from_subscription_arg, topic_from_websocket_arg},
61};
62use crate::{
63 common::{
64 consts::{OKX_FIELD_SMSG, OKX_SUCCESS_CODE, should_retry_error_code},
65 enums::{OKXOrderStatus, OKXOrderType},
66 parse::prefer_rpi_response_fields,
67 },
68 websocket::client::OKX_RATE_LIMIT_KEY_SUBSCRIPTION,
69};
70
71pub enum HandlerCommand {
73 SetClient(WebSocketClient),
75 Disconnect,
77 Authenticate { payload: SecretString },
79 Subscribe { args: Vec<OKXSubscriptionArg> },
81 SubscribeBook {
83 subscription: OKXSubscriptionArg,
84 cancel: CancellationToken,
85 gate: SnapshotGate,
86 completion: tokio::sync::oneshot::Sender<Result<(), OKXWsError>>,
87 },
88 Unsubscribe { args: Vec<OKXSubscriptionArg> },
90 Resubscribe {
92 subscription: OKXSubscriptionArg,
93 cancel: CancellationToken,
94 gate: SnapshotGate,
95 completion: tokio::sync::oneshot::Sender<Result<(), OKXWsError>>,
96 },
97 Send {
99 payload: String,
100 rate_limit_keys: Option<Vec<Ustr>>,
101 request_id: Option<String>,
102 client_order_ids: Vec<ClientOrderId>,
103 op: Option<OKXWsOperation>,
104 },
105}
106
107impl Debug for HandlerCommand {
108 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
109 match self {
110 Self::SetClient(_) => f.write_str("SetClient"),
111 Self::Disconnect => f.write_str("Disconnect"),
112 Self::Authenticate { .. } => f
113 .debug_struct(stringify!(Authenticate))
114 .field("payload", &REDACTED)
115 .finish(),
116 Self::Subscribe { args } => f
117 .debug_struct(stringify!(Subscribe))
118 .field("args", args)
119 .finish(),
120 Self::SubscribeBook { subscription, .. } => {
121 f.debug_tuple("SubscribeBook").field(subscription).finish()
122 }
123 Self::Unsubscribe { args } => f
124 .debug_struct(stringify!(Unsubscribe))
125 .field("args", args)
126 .finish(),
127 Self::Resubscribe { subscription, .. } => {
128 f.debug_tuple("Resubscribe").field(subscription).finish()
129 }
130 Self::Send {
131 rate_limit_keys,
132 request_id,
133 client_order_ids,
134 op,
135 ..
136 } => f
137 .debug_struct(stringify!(Send))
138 .field("payload", &REDACTED)
139 .field("rate_limit_keys", rate_limit_keys)
140 .field("request_id", request_id)
141 .field("client_order_ids", client_order_ids)
142 .field("op", op)
143 .finish(),
144 }
145 }
146}
147
148pub(super) struct OKXWsFeedHandler {
149 clock: &'static AtomicTime,
150 signal: Arc<AtomicBool>,
151 inner: Option<WebSocketClient>,
152 cmd_rx: tokio::sync::mpsc::UnboundedReceiver<HandlerCommand>,
153 raw_rx: tokio::sync::mpsc::UnboundedReceiver<Message>,
154 out_tx: tokio::sync::mpsc::UnboundedSender<OKXWsMessage>,
155 auth_tracker: AuthTracker,
156 subscriptions_state: SubscriptionState,
157 retry_manager: RetryManager<OKXWsError>,
158 pending_messages: VecDeque<OKXWsMessage>,
159}
160
161impl OKXWsFeedHandler {
162 pub(super) fn new(
164 signal: Arc<AtomicBool>,
165 cmd_rx: tokio::sync::mpsc::UnboundedReceiver<HandlerCommand>,
166 raw_rx: tokio::sync::mpsc::UnboundedReceiver<Message>,
167 out_tx: tokio::sync::mpsc::UnboundedSender<OKXWsMessage>,
168 auth_tracker: AuthTracker,
169 subscriptions_state: SubscriptionState,
170 clock: &'static AtomicTime,
171 ) -> Self {
172 Self {
173 clock,
174 signal,
175 inner: None,
176 cmd_rx,
177 raw_rx,
178 out_tx,
179 auth_tracker,
180 subscriptions_state,
181 retry_manager: create_websocket_retry_manager(),
182 pending_messages: VecDeque::new(),
183 }
184 }
185
186 pub(super) fn is_stopped(&self) -> bool {
187 self.signal.load(Ordering::Acquire)
188 }
189
190 pub(super) fn send(&self, msg: OKXWsMessage) -> Result<(), ()> {
191 self.out_tx.send(msg).map_err(|_| ())
192 }
193
194 async fn send_with_retry(
195 &self,
196 payload: String,
197 rate_limit_keys: Option<&[Ustr]>,
198 ) -> Result<(), OKXWsError> {
199 self.send_secret_with_retry(payload.into(), rate_limit_keys)
200 .await
201 }
202
203 async fn send_secret_with_retry(
204 &self,
205 payload: SecretString,
206 rate_limit_keys: Option<&[Ustr]>,
207 ) -> Result<(), OKXWsError> {
208 if let Some(client) = &self.inner {
209 let keys_owned: Option<Vec<Ustr>> = rate_limit_keys.map(<[Ustr]>::to_vec);
210 self.retry_manager
211 .invocation(
212 "websocket_send",
213 || {
214 let payload = payload.clone();
215 let keys = keys_owned.clone();
216 async move {
217 client
218 .send_text(payload.expose_secret().to_owned(), keys.as_deref())
219 .await
220 .map_err(OKXWsError::TransportSend)
221 }
222 },
223 should_retry_replay_safe_error,
224 create_okx_retry_error,
225 )
226 .execute()
227 .await
228 } else {
229 Err(OKXWsError::NoActiveClient)
230 }
231 }
232
233 async fn send_on_connection(
234 &self,
235 payload: String,
236 rate_limit_keys: Option<&[Ustr]>,
237 ) -> Result<(), OKXWsError> {
238 let client = self.inner.as_ref().ok_or(OKXWsError::NoActiveClient)?;
239 let connection_epoch = client.connection_epoch();
240 client
241 .send_text_on_connection(payload, rate_limit_keys, connection_epoch)
242 .await
243 .map_err(OKXWsError::TransportSend)
244 }
245
246 pub(super) async fn send_pong(&self) -> anyhow::Result<()> {
247 match self.send_on_connection(TEXT_PONG.to_string(), None).await {
248 Ok(()) => {
249 log::trace!("Sent pong response to OKX text ping");
250 Ok(())
251 }
252 Err(e) => {
253 log::warn!("Failed to send pong: error={e}");
254 Err(anyhow::anyhow!("Failed to send pong: {e}"))
255 }
256 }
257 }
258
259 pub(super) async fn next(&mut self) -> Option<OKXWsMessage> {
260 if let Some(message) = self.pending_messages.pop_front() {
261 return Some(message);
262 }
263
264 let mut poll_raw_next = false;
265
266 loop {
267 if self.signal.load(Ordering::Acquire) {
268 log::debug!("Stop signal received");
269 return None;
270 }
271
272 tokio::select! {
273 biased;
274 Some(cmd) = self.cmd_rx.recv(), if !poll_raw_next => {
275 match cmd {
276 HandlerCommand::SetClient(client) => {
277 log::debug!("Handler received WebSocket client");
278 self.inner = Some(client);
279 }
280 HandlerCommand::Disconnect => {
281 log::debug!("Handler disconnecting WebSocket client");
282 self.inner = None;
283 return None;
284 }
285 HandlerCommand::Authenticate { payload } => {
286 if let Err(e) = self.send_secret_with_retry(
287 payload,
288 Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()),
289 ).await {
290 log::error!(
291 "Failed to send authentication message after retries: error={e}"
292 );
293 }
294 }
295 HandlerCommand::Subscribe { args } => {
296 if let Err(e) = self.handle_subscribe(args).await {
297 log::error!("Failed to handle subscribe command: error={e}");
298 }
299 }
300 HandlerCommand::SubscribeBook { subscription, cancel, gate, completion } => {
301 let result = tokio::select! {
302 biased;
303 () = cancel.cancelled() => continue,
304 result = async {
305 let client = self.inner.as_ref().ok_or(OKXWsError::NoActiveClient)?;
306 self.send_book_subscription(
307 OKXWsOperation::Subscribe,
308 subscription,
309 client.connection_epoch(),
310 ).await
311 } => result,
312 };
313
314 if result.is_ok() {
315 gate.open();
316 }
317 let _ = completion.send(result);
318 }
319 HandlerCommand::Unsubscribe { args } => {
320 if let Err(e) = self.handle_unsubscribe(args).await {
321 log::error!("Failed to handle unsubscribe command: error={e}");
322 }
323 }
324 HandlerCommand::Resubscribe { subscription, cancel, gate, completion } => {
325 let result = tokio::select! {
326 biased;
327 () = cancel.cancelled() => continue,
328 result = async {
329 self.subscriptions_state.mark_failure(&topic_from_subscription_arg(&subscription));
330 let client = self.inner.as_ref().ok_or(OKXWsError::NoActiveClient)?;
331 let epoch = client.connection_epoch();
332 self.send_book_subscription(
333 OKXWsOperation::Unsubscribe,
334 subscription.clone(),
335 epoch,
336 ).await?;
337 self.send_book_subscription(
338 OKXWsOperation::Subscribe,
339 subscription,
340 epoch,
341 ).await
342 } => result,
343 };
344
345 if result.is_ok() {
346 gate.open();
347 }
348 let _ = completion.send(result);
349 }
350 HandlerCommand::Send {
351 payload,
352 rate_limit_keys,
353 request_id,
354 client_order_ids,
355 op,
356 } => {
357 if let Err(e) = self.send_on_connection(
358 payload,
359 rate_limit_keys.as_deref(),
360 ).await {
361 log::error!("Failed to send message: error={e}");
362
363 if let Some(request_id) = request_id {
364 self.pending_messages.push_back(OKXWsMessage::SendFailed {
365 request_id,
366 client_order_ids,
367 op,
368 error: e,
369 });
370 }
371 }
372 }
373 }
374
375 poll_raw_next = true;
376 }
377
378 () = time::sleep(time::Duration::from_millis(100)) => {
379 }
381
382 msg = self.raw_rx.recv() => {
383 let event = match msg {
384 Some(msg) => match Self::parse_raw_message(msg) {
385 Some(event) => event,
386 None => continue,
387 },
388 None => {
389 log::debug!("WebSocket stream closed");
390 return None;
391 }
392 };
393
394 match event {
395 OKXWsFrame::Ping => {
396 if let Err(e) = self.send_pong().await {
397 log::warn!("Failed to send pong response: error={e}");
398 }
399 }
400 OKXWsFrame::Login {
401 code, msg, conn_id, ..
402 } => {
403 if code == OKX_SUCCESS_CODE {
404 self.auth_tracker.succeed();
405 return Some(OKXWsMessage::Authenticated);
406 }
407
408 log::error!("WebSocket authentication failed: error={msg}");
409 self.auth_tracker.fail(msg.clone());
410
411 let error = OKXWebSocketError {
412 code,
413 message: msg,
414 conn_id: Some(conn_id),
415 timestamp: self.clock.get_time_ns().as_u64(),
416 };
417 self.pending_messages.push_back(OKXWsMessage::Error(error));
418 }
419 OKXWsFrame::BookData { arg, action, data } => {
420 return Some(OKXWsMessage::BookData { arg, action, data });
421 }
422 OKXWsFrame::RpiBookData { arg, action, data } => {
423 return Some(OKXWsMessage::RpiBookData { arg, action, data });
424 }
425 OKXWsFrame::OrderResponse {
426 id, op, code, msg, data,
427 } => {
428 return Some(OKXWsMessage::OrderResponse {
429 id, op, code, msg, data,
430 });
431 }
432 OKXWsFrame::Data { arg, data } => {
433 if let Some(output) = self.route_data_message(arg, data) {
434 return Some(output);
435 }
436 }
437 OKXWsFrame::Error { arg, code, msg } => {
438 let arg = arg.or_else(|| subscription_arg_from_error_message(&msg));
439 if let Some(arg) = arg
440 && self.handle_subscription_error(&arg, &code, &msg)
441 {
442 return Some(OKXWsMessage::SubscriptionFailed {
443 channel: arg.channel,
444 inst_id: arg.inst_id,
445 code,
446 msg,
447 });
448 }
449
450 let error = OKXWebSocketError {
451 code,
452 message: msg,
453 conn_id: None,
454 timestamp: self.clock.get_time_ns().as_u64(),
455 };
456 return Some(OKXWsMessage::Error(error));
457 }
458 OKXWsFrame::Reconnected => {
459 self.auth_tracker.invalidate();
460 return Some(OKXWsMessage::Reconnected);
461 }
462 OKXWsFrame::Subscription {
463 event, arg, code, msg,
464 ..
465 } => {
466 let rejected = self
467 .handle_subscription_ack(&event, &arg, code.as_deref(), msg.as_deref());
468
469 if rejected {
470 return Some(OKXWsMessage::SubscriptionFailed {
471 channel: arg.channel,
472 inst_id: arg.inst_id,
473 code: code.unwrap_or_default(),
474 msg: msg.unwrap_or_default(),
475 });
476 }
477 }
478 OKXWsFrame::ChannelConnCount { .. } => {}
479 }
480 }
481
482 () = std::future::ready(()), if poll_raw_next => {
483 poll_raw_next = false;
484 }
485
486 else => {
487 log::debug!("Handler shutting down: stream ended or command channel closed");
488 return None;
489 }
490 }
491 }
492 }
493
494 fn route_data_message(&self, arg: OKXWebSocketArg, mut data: Value) -> Option<OKXWsMessage> {
495 let OKXWebSocketArg {
496 channel, inst_id, ..
497 } = arg;
498
499 match channel {
500 OKXWsChannel::Account => Some(OKXWsMessage::Account(data)),
501 OKXWsChannel::Positions => Some(OKXWsMessage::Positions(data)),
502 OKXWsChannel::Orders => {
503 parse_array_items(data, "orders", false).map(OKXWsMessage::Orders)
504 }
505 OKXWsChannel::SprdOrders => {
506 parse_array_items(data, "spread orders", false).map(OKXWsMessage::SpreadOrders)
507 }
508 OKXWsChannel::OrdersAlgo | OKXWsChannel::AlgoAdvance => {
509 parse_array_items(data, "algo orders", false).map(OKXWsMessage::AlgoOrders)
510 }
511 OKXWsChannel::LiquidationWarning => {
512 parse_array_items(data, "liquidation warnings", false)
513 .map(OKXWsMessage::LiquidationWarnings)
514 }
515 OKXWsChannel::Instruments => {
516 prefer_rpi_response_fields(&mut data);
517 parse_array_items(data, "instruments", true).map(OKXWsMessage::Instruments)
518 }
519 _ => Some(OKXWsMessage::ChannelData {
520 channel,
521 inst_id,
522 data,
523 }),
524 }
525 }
526
527 fn handle_subscription_ack(
528 &self,
529 event: &OKXSubscriptionEvent,
530 arg: &OKXWebSocketArg,
531 code: Option<&str>,
532 msg: Option<&str>,
533 ) -> bool {
534 let topic = topic_from_websocket_arg(arg);
535 let success = code.is_none_or(|c| c == OKX_SUCCESS_CODE);
536
537 match event {
538 OKXSubscriptionEvent::Subscribe => {
539 if success {
540 self.subscriptions_state.confirm_subscribe(&topic);
541 false
542 } else {
543 log::warn!(
544 "Subscription failed: topic={topic:?}, error={msg:?}, code={code:?}"
545 );
546 self.subscriptions_state.mark_failure(&topic);
547 true
548 }
549 }
550 OKXSubscriptionEvent::Unsubscribe => {
551 if success {
552 self.subscriptions_state.confirm_unsubscribe(&topic);
553 } else {
554 log::warn!(
555 "Unsubscription failed - restoring subscription: \
556 topic={topic:?}, error={msg:?}, code={code:?}"
557 );
558 self.subscriptions_state.confirm_unsubscribe(&topic);
559 self.subscriptions_state.mark_subscribe(&topic);
560 self.subscriptions_state.confirm_subscribe(&topic);
561 }
562 false
563 }
564 }
565 }
566
567 fn handle_subscription_error(&self, arg: &OKXWebSocketArg, code: &str, msg: &str) -> bool {
568 let topic = topic_from_websocket_arg(arg);
569 let event = if self
570 .subscriptions_state
571 .pending_unsubscribe_topics()
572 .iter()
573 .any(|pending| pending == &topic)
574 {
575 OKXSubscriptionEvent::Unsubscribe
576 } else if self
577 .subscriptions_state
578 .pending_subscribe_topics()
579 .iter()
580 .any(|pending| pending == &topic)
581 || (arg.channel.is_book()
582 && self
583 .subscriptions_state
584 .all_topics()
585 .iter()
586 .any(|tracked| tracked == &topic))
587 {
588 OKXSubscriptionEvent::Subscribe
589 } else {
590 return false;
591 };
592
593 self.handle_subscription_ack(&event, arg, Some(code), Some(msg))
594 }
595
596 async fn send_book_subscription(
597 &self,
598 op: OKXWsOperation,
599 subscription: OKXSubscriptionArg,
600 connection_epoch: u64,
601 ) -> Result<(), OKXWsError> {
602 let client = self.inner.as_ref().ok_or(OKXWsError::NoActiveClient)?;
603 let message = OKXSubscription {
604 op,
605 args: vec![subscription],
606 };
607 let payload =
608 serde_json::to_string(&message).map_err(|e| OKXWsError::ClientError(e.to_string()))?;
609
610 client
611 .send_text_on_connection(
612 payload,
613 Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()),
614 connection_epoch,
615 )
616 .await
617 .map_err(OKXWsError::TransportSend)
618 }
619
620 async fn handle_subscribe(&self, args: Vec<OKXSubscriptionArg>) -> anyhow::Result<()> {
621 for arg in &args {
622 log::debug!(
623 "Subscribing to channel: channel={:?}, inst_id={:?}",
624 arg.channel,
625 arg.inst_id
626 );
627 }
628
629 let message = OKXSubscription {
630 op: OKXWsOperation::Subscribe,
631 args,
632 };
633
634 let json_txt = serde_json::to_string(&message)
635 .map_err(|e| anyhow::anyhow!("Failed to serialize subscription: {e}"))?;
636
637 self.send_with_retry(json_txt, Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()))
638 .await
639 .map_err(|e| anyhow::anyhow!("Failed to send subscription after retries: {e}"))?;
640 Ok(())
641 }
642
643 async fn handle_unsubscribe(&self, args: Vec<OKXSubscriptionArg>) -> anyhow::Result<()> {
644 for arg in &args {
645 log::debug!(
646 "Unsubscribing from channel: channel={:?}, inst_id={:?}",
647 arg.channel,
648 arg.inst_id
649 );
650 }
651
652 let message = OKXSubscription {
653 op: OKXWsOperation::Unsubscribe,
654 args,
655 };
656
657 let json_txt = serde_json::to_string(&message)
658 .map_err(|e| anyhow::anyhow!("Failed to serialize unsubscription: {e}"))?;
659
660 self.send_with_retry(json_txt, Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()))
661 .await
662 .map_err(|e| anyhow::anyhow!("Failed to send unsubscription after retries: {e}"))?;
663 Ok(())
664 }
665
666 pub(crate) fn parse_raw_message(
667 msg: tokio_tungstenite::tungstenite::Message,
668 ) -> Option<OKXWsFrame> {
669 match msg {
670 tokio_tungstenite::tungstenite::Message::Text(text) => {
671 if text == TEXT_PONG {
672 log::trace!("Received pong from OKX");
673 return None;
674 }
675
676 if text == TEXT_PING {
677 log::trace!("Received ping from OKX (text)");
678 return Some(OKXWsFrame::Ping);
679 }
680
681 if text == RECONNECTED {
682 log::debug!("Received WebSocket reconnection signal");
683 return Some(OKXWsFrame::Reconnected);
684 }
685 log::trace!("Received WebSocket message: {text}");
686
687 match serde_json::from_str(&text) {
688 Ok(ws_event) => match &ws_event {
689 OKXWsFrame::Error { code, msg, .. } => {
690 if should_retry_error_code(code) {
691 log::warn!("WebSocket error: {code} - {msg}");
692 } else {
693 log::error!("WebSocket error: {code} - {msg}");
694 }
695 Some(ws_event)
696 }
697 OKXWsFrame::Login {
698 event,
699 code,
700 msg,
701 conn_id,
702 } => {
703 if code == OKX_SUCCESS_CODE {
704 log::debug!("WebSocket authenticated: conn_id={conn_id}");
705 } else {
706 log::error!(
707 "WebSocket authentication failed: \
708 event={event}, code={code}, error={msg}"
709 );
710 }
711 Some(ws_event)
712 }
713 OKXWsFrame::Subscription {
714 event,
715 arg,
716 conn_id,
717 ..
718 } => {
719 let channel_str = serde_json::to_string(&arg.channel)
720 .expect("Invalid OKX websocket channel")
721 .trim_matches('"')
722 .to_string();
723 log::debug!("{event}d: channel={channel_str}, conn_id={conn_id}");
724 Some(ws_event)
725 }
726 OKXWsFrame::ChannelConnCount {
727 channel,
728 conn_count,
729 conn_id,
730 ..
731 } => {
732 let channel_str = serde_json::to_string(channel)
733 .expect("Invalid OKX websocket channel")
734 .trim_matches('"')
735 .to_string();
736 log::debug!(
737 "Channel connection status: \
738 channel={channel_str}, connections={conn_count}, conn_id={conn_id}",
739 );
740 None
741 }
742 OKXWsFrame::Ping => {
743 log::trace!("Ignoring ping event parsed from text payload");
744 None
745 }
746 OKXWsFrame::Data { .. }
747 | OKXWsFrame::BookData { .. }
748 | OKXWsFrame::RpiBookData { .. } => Some(ws_event),
749 OKXWsFrame::OrderResponse {
750 id, op, code, data, ..
751 } => {
752 if code == OKX_SUCCESS_CODE {
753 log::debug!(
754 "Order operation successful: id={id:?}, op={op}, code={code}"
755 );
756
757 if let Some(order_data) = data.first() {
758 let success_msg = order_data
759 .get(OKX_FIELD_SMSG)
760 .and_then(|s| s.as_str())
761 .unwrap_or("Order operation successful");
762 log::debug!("Order success details: {success_msg}");
763 }
764 }
765 Some(ws_event)
766 }
767 OKXWsFrame::Reconnected => {
768 log::warn!("Unexpected Reconnected event from deserialization");
769 None
770 }
771 },
772 Err(e) => {
773 log::error!("Failed to parse message: {e}: {text}");
774 None
775 }
776 }
777 }
778 Message::Ping(_payload) => {
779 log::trace!("Received binary ping frame from OKX");
780 Some(OKXWsFrame::Ping)
781 }
782 Message::Pong(payload) => {
783 log::trace!("Received pong frame from OKX ({} bytes)", payload.len());
784 None
785 }
786 Message::Binary(msg) => {
787 log::debug!("Raw binary frame ({} bytes)", msg.len());
788 log::trace!("Raw binary: {msg:?}");
789 None
790 }
791 Message::Close(_) => {
792 log::debug!("Received close message");
793 None
794 }
795 msg => {
796 log::warn!("Unexpected message: {msg}");
797 None
798 }
799 }
800 }
801}
802
803fn subscription_arg_from_error_message(msg: &str) -> Option<OKXWebSocketArg> {
804 let descriptor = msg
805 .strip_prefix("Wrong URL or channel:")?
806 .split_whitespace()
807 .next()?;
808 let mut fields = descriptor.split(',');
809 let channel = fields.next()?;
810 let mut arg = Map::new();
811 arg.insert("channel".to_string(), Value::String(channel.to_string()));
812
813 for field in fields {
814 let (key, value) = field.split_once(':')?;
815 if !matches!(key, "instId" | "sprdId" | "instType" | "instFamily") {
816 return None;
817 }
818 arg.insert(key.to_string(), Value::String(value.to_string()));
819 }
820
821 serde_json::from_value(Value::Object(arg)).ok()
822}
823
824pub fn is_post_only_auto_cancel(msg: &OKXOrderMsg) -> bool {
826 use crate::common::{consts::OKX_POST_ONLY_CANCEL_SOURCE, enums::OKXOrderStatus};
827
828 if msg.state != OKXOrderStatus::Canceled {
829 return false;
830 }
831
832 let cancel_source_matches = matches!(
833 msg.cancel_source.as_deref(),
834 Some(source) if source == OKX_POST_ONLY_CANCEL_SOURCE
835 );
836
837 let reason_matches = matches!(
838 msg.cancel_source_reason.as_deref(),
839 Some(reason) if reason.contains("POST_ONLY")
840 );
841
842 if !(cancel_source_matches || reason_matches) {
843 return false;
844 }
845
846 msg.acc_fill_sz
847 .as_ref()
848 .is_none_or(|filled| filled == "0" || filled.is_empty())
849}
850
851pub fn is_unfilled_rpi_cancel(msg: &OKXOrderMsg) -> bool {
853 msg.ord_type == OKXOrderType::Rpi
854 && msg.state == OKXOrderStatus::Canceled
855 && msg
856 .acc_fill_sz
857 .as_ref()
858 .is_none_or(|filled| filled == "0" || filled.is_empty())
859}
860
861fn parse_array_items<T: serde::de::DeserializeOwned>(
863 data: Value,
864 label: &str,
865 warn_on_parse_error: bool,
866) -> Option<Vec<T>> {
867 let Value::Array(items) = data else {
868 if warn_on_parse_error {
869 log::warn!("Expected {label} payload to be a JSON array");
870 } else {
871 log::error!("Expected {label} payload to be a JSON array");
872 }
873 return None;
874 };
875
876 let mut parsed = Vec::with_capacity(items.len());
877 for (idx, item) in items.into_iter().enumerate() {
878 match serde_json::from_value::<T>(item) {
879 Ok(value) => parsed.push(value),
880 Err(e) => {
881 if warn_on_parse_error {
882 log::warn!("Failed to parse {label} item at index {idx}: {e}");
883 } else {
884 log::error!("Failed to parse {label} item at index {idx}: {e}");
885 }
886 }
887 }
888 }
889
890 if parsed.is_empty() {
891 None
892 } else {
893 Some(parsed)
894 }
895}
896
897fn should_retry_replay_safe_error(error: &OKXWsError) -> bool {
898 match error {
899 OKXWsError::OkxError { error_code, .. } => should_retry_error_code(error_code),
900 OKXWsError::TransportSend(SendError::Timeout | SendError::ConnectionChanged)
901 | OKXWsError::TungsteniteError(_)
902 | OKXWsError::OperationTimeout { .. } => true,
903 OKXWsError::AuthenticationError(_)
904 | OKXWsError::JsonError(_)
905 | OKXWsError::ParsingError(_)
906 | OKXWsError::ClientError(_)
907 | OKXWsError::NoActiveClient
908 | OKXWsError::HandlerUnavailable(_)
909 | OKXWsError::TransportSend(
910 SendError::InvalidInput(_)
911 | SendError::Closed
912 | SendError::WriteTimeout
913 | SendError::BrokenPipe(_),
914 )
915 | OKXWsError::SendFailed(_) => false,
916 }
917}
918
919fn create_okx_retry_error(error: RetryError) -> OKXWsError {
920 match error {
921 RetryError::OperationTimeout { timeout_ms } => OKXWsError::OperationTimeout { timeout_ms },
922 RetryError::InvalidConfiguration { message } => OKXWsError::ClientError(message),
923 RetryError::Canceled => {
924 OKXWsError::SendFailed("Adapter disconnecting or shutting down".to_string())
925 }
926 error @ RetryError::ElapsedBudgetExceeded { .. } => {
927 OKXWsError::SendFailed(error.to_string())
928 }
929 }
930}
931
932#[cfg(test)]
933mod tests {
934 use std::{
935 num::NonZeroU32,
936 sync::{Arc, atomic::AtomicBool},
937 time::Duration,
938 };
939
940 use futures_util::{SinkExt, StreamExt};
941 use nautilus_core::{collections::AtomicMap, time::get_atomic_clock_realtime};
942 use nautilus_model::identifiers::InstrumentId;
943 use nautilus_network::{
944 ratelimiter::{RateLimiter, quota::Quota},
945 websocket::{AuthTracker, SubscriptionState, WebSocketConfig, channel_message_handler},
946 };
947 use rstest::rstest;
948 use serde_json::json;
949
950 use super::*;
951 use crate::{
952 book::{BookChannelScope, BookSequenceOutcome, sync::BookSyncTracker},
953 common::{
954 consts::OKX_WS_TOPIC_DELIMITER,
955 enums::{OKXBookAction, OKXBookChannel, OKXRpiPermission},
956 testing::load_test_json,
957 },
958 };
959
960 fn create_handler() -> OKXWsFeedHandler {
961 let signal = Arc::new(AtomicBool::new(false));
962 let (_cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
963 let (_raw_tx, raw_rx) = tokio::sync::mpsc::unbounded_channel();
964 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
965
966 OKXWsFeedHandler::new(
967 signal,
968 cmd_rx,
969 raw_rx,
970 out_tx,
971 AuthTracker::new(),
972 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
973 get_atomic_clock_realtime(),
974 )
975 }
976
977 #[rstest]
978 fn test_command_debug_redacts_payloads() {
979 let payload = "authentication-secret";
980 let authenticate = HandlerCommand::Authenticate {
981 payload: SecretString::from(payload.to_string()),
982 };
983 let send = HandlerCommand::Send {
984 payload: payload.to_string(),
985 rate_limit_keys: None,
986 request_id: None,
987 client_order_ids: Vec::new(),
988 op: None,
989 };
990
991 let debug = format!("{authenticate:?} {send:?}");
992
993 assert!(debug.contains(REDACTED));
994 assert!(!debug.contains(payload));
995 }
996
997 #[tokio::test]
998 async fn test_next_polls_raw_after_one_ready_command() {
999 let signal = Arc::new(AtomicBool::new(false));
1000 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
1001 let (raw_tx, raw_rx) = tokio::sync::mpsc::unbounded_channel();
1002 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
1003 let mut handler = OKXWsFeedHandler::new(
1004 signal,
1005 cmd_rx,
1006 raw_rx,
1007 out_tx,
1008 AuthTracker::new(),
1009 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
1010 get_atomic_clock_realtime(),
1011 );
1012
1013 for _ in 0..3 {
1014 cmd_tx
1015 .send(HandlerCommand::Subscribe { args: Vec::new() })
1016 .unwrap();
1017 }
1018 raw_tx
1019 .send(Message::Text(RECONNECTED.to_string().into()))
1020 .unwrap();
1021
1022 let message = handler.next().await;
1023
1024 assert!(matches!(message, Some(OKXWsMessage::Reconnected)));
1025 assert_eq!(handler.cmd_rx.len(), 2);
1026 }
1027
1028 #[rstest]
1029 #[case::sent(false)]
1030 #[case::canceled(true)]
1031 #[tokio::test]
1032 async fn initial_book_completion_waits_for_send(#[case] canceled: bool) {
1033 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1034 let url = format!("ws://{}/", listener.local_addr().unwrap());
1035 let (wire_tx, mut wire_rx) = tokio::sync::mpsc::unbounded_channel();
1036
1037 let server = tokio::spawn(async move {
1038 let (socket, _) = listener.accept().await.unwrap();
1039 let mut socket = tokio_tungstenite::accept_async(socket).await.unwrap();
1040 while let Some(Ok(frame)) = socket.next().await {
1041 wire_tx.send(frame).unwrap();
1042 }
1043 });
1044
1045 let (message_handler, raw_rx) = channel_message_handler();
1046 let client = WebSocketClient::builder()
1047 .config(WebSocketConfig::builder().url(url).build().unwrap())
1048 .message_handler(message_handler)
1049 .default_quota(Quota::with_period(Duration::from_secs(10)).unwrap())
1050 .connect()
1051 .await
1052 .unwrap();
1053 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
1054 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
1055
1056 let mut handler = OKXWsFeedHandler::new(
1057 Arc::new(AtomicBool::new(false)),
1058 cmd_rx,
1059 raw_rx,
1060 out_tx,
1061 AuthTracker::new(),
1062 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
1063 get_atomic_clock_realtime(),
1064 );
1065 handler.inner = Some(client);
1066 cmd_tx
1067 .send(HandlerCommand::Subscribe { args: Vec::new() })
1068 .unwrap();
1069 let cancel = CancellationToken::new();
1070 let gate = SnapshotGate::default();
1071 gate.lock().close();
1072 let (completion, mut result) = tokio::sync::oneshot::channel();
1073 cmd_tx
1074 .send(HandlerCommand::SubscribeBook {
1075 subscription: OKXSubscriptionArg {
1076 channel: OKXWsChannel::Books,
1077 inst_id: Some(Ustr::from("BTC-USDT")),
1078 inst_type: None,
1079 inst_family: None,
1080 },
1081 cancel: cancel.clone(),
1082 gate: gate.clone(),
1083 completion,
1084 })
1085 .unwrap();
1086
1087 let running = tokio::spawn(async move { handler.next().await });
1088 let first = tokio::time::timeout(Duration::from_secs(3), wire_rx.recv())
1089 .await
1090 .unwrap()
1091 .unwrap();
1092 let first: Value = serde_json::from_str(first.to_text().unwrap()).unwrap();
1093 assert_eq!(first, serde_json::json!({"op": "subscribe", "args": []}));
1094 assert!(matches!(
1095 result.try_recv(),
1096 Err(tokio::sync::oneshot::error::TryRecvError::Empty)
1097 ));
1098
1099 if canceled {
1100 cancel.cancel();
1101 }
1102
1103 tokio::time::pause();
1104 tokio::time::advance(Duration::from_secs(10)).await;
1105 tokio::time::resume();
1106 let result = tokio::time::timeout(Duration::from_secs(3), result)
1107 .await
1108 .unwrap();
1109
1110 if canceled {
1111 assert!(gate.lock().is_closed());
1112 assert!(result.is_err());
1113 assert!(
1114 tokio::time::timeout(Duration::from_millis(100), wire_rx.recv())
1115 .await
1116 .is_err()
1117 );
1118 } else {
1119 result.unwrap().unwrap();
1120 assert!(!gate.lock().is_closed());
1121 let frame = tokio::time::timeout(Duration::from_secs(3), wire_rx.recv())
1122 .await
1123 .unwrap()
1124 .unwrap();
1125 let frame: Value = serde_json::from_str(frame.to_text().unwrap()).unwrap();
1126 assert_eq!(
1127 frame,
1128 serde_json::json!({"op": "subscribe", "args": [{"channel": "books", "instId": "BTC-USDT"}]})
1129 );
1130 }
1131
1132 running.abort();
1133 server.abort();
1134 }
1135
1136 #[tokio::test]
1137 async fn book_subscription_rejects_changed_connection() {
1138 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1139 let url = format!("ws://{}/", listener.local_addr().unwrap());
1140 let (wire_tx, mut wire_rx) = tokio::sync::mpsc::unbounded_channel();
1141
1142 let server = tokio::spawn(async move {
1143 let (socket, _) = listener.accept().await.unwrap();
1144 let mut socket = tokio_tungstenite::accept_async(socket).await.unwrap();
1145 while let Some(Ok(frame)) = socket.next().await {
1146 wire_tx.send(frame).unwrap();
1147 }
1148 });
1149 let (message_handler, raw_rx) = channel_message_handler();
1150 let client = WebSocketClient::builder()
1151 .config(WebSocketConfig::builder().url(url).build().unwrap())
1152 .message_handler(message_handler)
1153 .connect()
1154 .await
1155 .unwrap();
1156 let wrong_epoch = client.connection_epoch() + 1;
1157 let (_cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
1158 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
1159 let mut handler = OKXWsFeedHandler::new(
1160 Arc::new(AtomicBool::new(false)),
1161 cmd_rx,
1162 raw_rx,
1163 out_tx,
1164 AuthTracker::new(),
1165 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
1166 get_atomic_clock_realtime(),
1167 );
1168 handler.inner = Some(client);
1169 let result = handler
1170 .send_book_subscription(
1171 OKXWsOperation::Subscribe,
1172 OKXSubscriptionArg {
1173 channel: OKXWsChannel::Books,
1174 inst_id: Some(Ustr::from("BTC-USDT")),
1175 inst_type: None,
1176 inst_family: None,
1177 },
1178 wrong_epoch,
1179 )
1180 .await;
1181
1182 assert!(matches!(
1183 result,
1184 Err(OKXWsError::TransportSend(SendError::ConnectionChanged))
1185 ));
1186 assert!(matches!(
1187 wire_rx.try_recv(),
1188 Err(tokio::sync::mpsc::error::TryRecvError::Empty)
1189 ));
1190 handler.inner.as_ref().unwrap().disconnect().await;
1191 server.abort();
1192 }
1193
1194 #[tokio::test]
1195 async fn recovery_cancel_between_sends_keeps_snapshot_gate_closed() {
1196 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1197 let url = format!("ws://{}/", listener.local_addr().unwrap());
1198 let (received, first_frame) = tokio::sync::oneshot::channel();
1199
1200 let server = tokio::spawn(async move {
1201 let (socket, _) = listener.accept().await.unwrap();
1202 let mut socket = tokio_tungstenite::accept_async(socket).await.unwrap();
1203 let frame = socket.next().await.unwrap().unwrap();
1204 received.send(frame).unwrap();
1205
1206 while socket.next().await.is_some() {}
1207 });
1208
1209 let (message_handler, raw_rx) = channel_message_handler();
1210 let client = WebSocketClient::builder()
1211 .config(WebSocketConfig::builder().url(url).build().unwrap())
1212 .message_handler(message_handler)
1213 .default_quota(Quota::per_hour(NonZeroU32::new(1).unwrap()))
1214 .connect()
1215 .await
1216 .unwrap();
1217 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
1218 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
1219
1220 let mut handler = OKXWsFeedHandler::new(
1221 Arc::new(AtomicBool::new(false)),
1222 cmd_rx,
1223 raw_rx,
1224 out_tx,
1225 AuthTracker::new(),
1226 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
1227 get_atomic_clock_realtime(),
1228 );
1229 handler.inner = Some(client);
1230 let tracker = BookSyncTracker::default();
1231 let instrument_id = InstrumentId::from("BTC-USDT.OKX");
1232 tracker.record_subscription(instrument_id, time::Instant::now(), SnapshotGate::default());
1233 let recovery = tracker.claim_recovery(instrument_id).unwrap();
1234 assert!(recovery.begin_replacement());
1235 let cancel = recovery.cancellation.child_token();
1236 let (completion, result) = tokio::sync::oneshot::channel();
1237 cmd_tx
1238 .send(HandlerCommand::Resubscribe {
1239 subscription: OKXSubscriptionArg {
1240 channel: OKXWsChannel::Books,
1241 inst_id: Some(Ustr::from("BTC-USDT")),
1242 inst_type: None,
1243 inst_family: None,
1244 },
1245 cancel: cancel.clone(),
1246 gate: recovery.gate.clone(),
1247 completion,
1248 })
1249 .unwrap();
1250
1251 let running = tokio::spawn(async move { handler.next().await });
1252 let frame = tokio::time::timeout(Duration::from_secs(3), first_frame)
1253 .await
1254 .unwrap()
1255 .unwrap();
1256 let request: Value = serde_json::from_str(frame.to_text().unwrap()).unwrap();
1257 assert_eq!(request["op"], "unsubscribe");
1258 assert_eq!(
1259 tracker.validate_sequence(
1260 instrument_id,
1261 true,
1262 &[(Some(-1), 42)],
1263 Duration::ZERO,
1264 time::Instant::now()
1265 ),
1266 BookSequenceOutcome::Suppress
1267 );
1268 cancel.cancel();
1269 assert!(
1270 tokio::time::timeout(Duration::from_secs(1), result)
1271 .await
1272 .unwrap()
1273 .is_err()
1274 );
1275
1276 assert!(recovery.gate.lock().is_closed());
1277 assert_eq!(
1278 tracker.validate_sequence(
1279 instrument_id,
1280 true,
1281 &[(Some(-1), 43)],
1282 Duration::ZERO,
1283 time::Instant::now()
1284 ),
1285 BookSequenceOutcome::Suppress
1286 );
1287 cmd_tx.send(HandlerCommand::Disconnect).unwrap();
1288 tokio::time::timeout(Duration::from_secs(3), running)
1289 .await
1290 .unwrap()
1291 .unwrap();
1292 server.abort();
1293 }
1294
1295 #[rstest]
1296 #[case::disabled_deadline(Duration::ZERO)]
1297 #[case::snapshot_deadline(Duration::from_secs(3))]
1298 #[tokio::test]
1299 async fn reconnect_replay_survives_inflight_recovery(#[case] snapshot_timeout: Duration) {
1300 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1301 let url = format!("ws://{}/", listener.local_addr().unwrap());
1302 let (wire_tx, mut wire_rx) = tokio::sync::mpsc::unbounded_channel();
1303
1304 let server = tokio::spawn(async move {
1305 let (socket, _) = listener.accept().await.unwrap();
1306 let mut socket = tokio_tungstenite::accept_async(socket).await.unwrap();
1307
1308 while let Some(Ok(Message::Text(text))) = socket.next().await {
1309 let request: Value = serde_json::from_str(&text).unwrap();
1310 wire_tx
1311 .send(request["op"].as_str().unwrap().to_owned())
1312 .unwrap();
1313
1314 if request["op"] == "subscribe" {
1315 for (fixture, sequence, previous) in [
1316 ("ws_books_snapshot.json", 100, -1),
1317 ("ws_books_update.json", 101, 100),
1318 ] {
1319 let mut frame: Value =
1320 serde_json::from_str(&load_test_json(fixture)).unwrap();
1321 frame["data"][0]["seqId"] = json!(sequence);
1322 frame["data"][0]["prevSeqId"] = json!(previous);
1323 socket
1324 .send(Message::Text(frame.to_string().into()))
1325 .await
1326 .unwrap();
1327 }
1328 }
1329 }
1330 });
1331
1332 let limiter = Arc::new(RateLimiter::new_with_quota(
1334 Some(
1335 Quota::with_period(Duration::from_secs(10))
1336 .unwrap()
1337 .allow_burst(NonZeroU32::new(2).unwrap()),
1338 ),
1339 Vec::new(),
1340 ));
1341 let (message_handler, raw_rx) = channel_message_handler();
1342 let client = WebSocketClient::builder()
1343 .config(WebSocketConfig::builder().url(url).build().unwrap())
1344 .message_handler(message_handler)
1345 .rate_limiter(Arc::clone(&limiter))
1346 .connect()
1347 .await
1348 .unwrap();
1349 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
1350 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
1351
1352 let mut handler = OKXWsFeedHandler::new(
1353 Arc::new(AtomicBool::new(false)),
1354 cmd_rx,
1355 raw_rx,
1356 out_tx,
1357 AuthTracker::new(),
1358 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
1359 get_atomic_clock_realtime(),
1360 );
1361 handler.inner = Some(client);
1362 let (messages_tx, mut messages_rx) = tokio::sync::mpsc::unbounded_channel();
1363
1364 let running = tokio::spawn(async move {
1365 while let Some(message) = handler.next().await {
1366 messages_tx.send(message).unwrap();
1367 }
1368 });
1369
1370 let instrument_id = InstrumentId::from("BTC-USDT.OKX");
1371 let channels = AtomicMap::new();
1372 channels.insert(instrument_id, OKXBookChannel::Book);
1373 let tracker = BookSyncTracker::default();
1374 tracker.record_subscription(instrument_id, time::Instant::now(), SnapshotGate::default());
1375 let recovery = tracker.claim_recovery(instrument_id).unwrap();
1376 assert!(recovery.begin_replacement());
1377
1378 let arg = OKXSubscriptionArg {
1379 channel: OKXWsChannel::Books,
1380 inst_id: Some(Ustr::from("BTC-USDT")),
1381 inst_type: None,
1382 inst_family: None,
1383 };
1384
1385 cmd_tx
1386 .send(HandlerCommand::Subscribe {
1387 args: vec![arg.clone()],
1388 })
1389 .unwrap();
1390
1391 assert_eq!(
1392 tokio::time::timeout(Duration::from_secs(3), wire_rx.recv())
1393 .await
1394 .unwrap()
1395 .unwrap(),
1396 "subscribe"
1397 );
1398
1399 for action in [OKXBookAction::Snapshot, OKXBookAction::Update] {
1400 let message = tokio::time::timeout(Duration::from_secs(3), messages_rx.recv())
1401 .await
1402 .unwrap()
1403 .unwrap();
1404 assert!(
1405 matches!(message, OKXWsMessage::BookData { action: received, .. } if received == action)
1406 );
1407 }
1408
1409 let (completion, result) = tokio::sync::oneshot::channel();
1410 cmd_tx
1411 .send(HandlerCommand::Resubscribe {
1412 subscription: arg.clone(),
1413 cancel: recovery.cancellation.child_token(),
1414 gate: recovery.gate.clone(),
1415 completion,
1416 })
1417 .unwrap();
1418
1419 assert_eq!(
1420 tokio::time::timeout(Duration::from_secs(3), wire_rx.recv())
1421 .await
1422 .unwrap()
1423 .unwrap(),
1424 "unsubscribe"
1425 );
1426
1427 tokio::time::pause();
1429 tracker.reset_sequences(&channels, BookChannelScope::Public);
1430
1431 if !snapshot_timeout.is_zero() {
1432 tracker.seed_pending_snapshots(
1433 &channels,
1434 BookChannelScope::Public,
1435 snapshot_timeout,
1436 time::Instant::now(),
1437 );
1438 }
1439
1440 tokio::time::advance(Duration::from_secs(10)).await;
1441 tokio::time::resume();
1442 let _ = tokio::time::timeout(Duration::from_secs(3), result)
1443 .await
1444 .expect("recovery command completes after reconnect");
1445
1446 let resumed = tokio::time::timeout(Duration::from_secs(3), async {
1447 for (action, sequence) in [(OKXBookAction::Snapshot, 100), (OKXBookAction::Update, 101)]
1448 {
1449 let Some(OKXWsMessage::BookData {
1450 action: received,
1451 data,
1452 ..
1453 }) = messages_rx.recv().await
1454 else {
1455 panic!("expected book data after reconnect");
1456 };
1457
1458 assert_eq!(received, action);
1459 assert_eq!(data[0].seq_id, sequence);
1460 assert_eq!(
1461 tracker.validate_sequence(
1462 instrument_id,
1463 received == OKXBookAction::Snapshot,
1464 &[(data[0].prev_seq_id, data[0].seq_id)],
1465 snapshot_timeout,
1466 time::Instant::now(),
1467 ),
1468 BookSequenceOutcome::Accept,
1469 );
1470 }
1471 })
1472 .await;
1473
1474 if resumed.is_err() {
1476 cmd_tx
1477 .send(HandlerCommand::Subscribe { args: vec![arg] })
1478 .unwrap();
1479
1480 for action in [OKXBookAction::Snapshot, OKXBookAction::Update] {
1481 let message = tokio::time::timeout(Duration::from_secs(3), messages_rx.recv())
1482 .await
1483 .unwrap()
1484 .unwrap();
1485 assert!(
1486 matches!(message, OKXWsMessage::BookData { action: received, .. } if received == action)
1487 );
1488 }
1489 }
1490
1491 cmd_tx.send(HandlerCommand::Disconnect).unwrap();
1492 tokio::time::timeout(Duration::from_secs(3), running)
1493 .await
1494 .unwrap()
1495 .unwrap();
1496 server.abort();
1497
1498 assert!(
1499 resumed.is_ok(),
1500 "reconnect must preserve an in-flight recovery and resume book output"
1501 );
1502 }
1503
1504 #[rstest]
1505 fn test_should_retry_typed_transport_and_timeout_errors() {
1506 assert!(should_retry_replay_safe_error(&OKXWsError::TransportSend(
1507 SendError::Timeout
1508 )));
1509 assert!(should_retry_replay_safe_error(&OKXWsError::TransportSend(
1510 SendError::ConnectionChanged
1511 )));
1512 assert!(!should_retry_replay_safe_error(&OKXWsError::TransportSend(
1513 SendError::WriteTimeout
1514 )));
1515 assert!(!should_retry_replay_safe_error(&OKXWsError::TransportSend(
1516 SendError::BrokenPipe("connection reset".to_string())
1517 )));
1518 assert!(should_retry_replay_safe_error(
1519 &OKXWsError::OperationTimeout { timeout_ms: 1_000 }
1520 ));
1521 assert!(!should_retry_replay_safe_error(&OKXWsError::NoActiveClient));
1522 assert!(!should_retry_replay_safe_error(
1523 &OKXWsError::HandlerUnavailable("closed".to_string())
1524 ));
1525 }
1526
1527 #[rstest]
1528 fn test_retryability_uses_websocket_error_type_not_message() {
1529 let message = "connection reset".to_string();
1530 let temporary = OKXWsError::OkxError {
1531 error_code: "50011".to_string(),
1532 message: message.clone(),
1533 };
1534 let permanent = OKXWsError::ClientError(message.clone());
1535 let ambiguous = OKXWsError::SendFailed(message);
1536
1537 assert!(should_retry_replay_safe_error(&temporary));
1538 assert!(!should_retry_replay_safe_error(&permanent));
1539 assert!(!should_retry_replay_safe_error(&ambiguous));
1540 }
1541
1542 #[rstest]
1543 fn test_subscription_error_restores_failed_unsubscribe() {
1544 let handler = create_handler();
1545 let arg = OKXWebSocketArg {
1546 channel: OKXWsChannel::Books,
1547 inst_id: Some(Ustr::from("BTC-USD")),
1548 inst_type: None,
1549 inst_family: None,
1550 bar: None,
1551 };
1552 let topic = topic_from_websocket_arg(&arg);
1553 handler.subscriptions_state.mark_subscribe(&topic);
1554 handler.subscriptions_state.confirm_subscribe(&topic);
1555 handler.subscriptions_state.mark_unsubscribe(&topic);
1556
1557 let rejected_subscription =
1558 handler.handle_subscription_error(&arg, "60019", "Unsubscription failed");
1559
1560 assert!(!rejected_subscription);
1561 assert_eq!(handler.subscriptions_state.all_topics(), vec![topic]);
1562 assert!(
1563 handler
1564 .subscriptions_state
1565 .pending_subscribe_topics()
1566 .is_empty()
1567 );
1568 assert!(
1569 handler
1570 .subscriptions_state
1571 .pending_unsubscribe_topics()
1572 .is_empty()
1573 );
1574 }
1575
1576 #[rstest]
1577 fn test_subscription_arg_from_error_message_matches_mainnet_shape() {
1578 let msg = "Wrong URL or channel:books,instId:BTC-USDT-SWAP doesn't exist. Please use the \
1579 correct URL, channel and parameters referring to API document.";
1580
1581 let arg = subscription_arg_from_error_message(msg).unwrap();
1582
1583 assert_eq!(arg.channel, OKXWsChannel::Books);
1584 assert_eq!(arg.inst_id, Some(Ustr::from("BTC-USDT-SWAP")));
1585 assert_eq!(arg.inst_type, None);
1586 assert_eq!(arg.inst_family, None);
1587 assert_eq!(arg.bar, None);
1588 }
1589
1590 #[rstest]
1591 fn test_subscription_error_ignores_non_pending_topic() {
1592 let handler = create_handler();
1593 let arg = OKXWebSocketArg {
1594 channel: OKXWsChannel::Books,
1595 inst_id: Some(Ustr::from("BTC-USDT-SWAP")),
1596 inst_type: None,
1597 inst_family: None,
1598 bar: None,
1599 };
1600
1601 let rejected_subscription =
1602 handler.handle_subscription_error(&arg, "60018", "Subscription failed");
1603
1604 assert!(!rejected_subscription);
1605 assert!(handler.subscriptions_state.all_topics().is_empty());
1606 assert!(
1607 handler
1608 .subscriptions_state
1609 .pending_subscribe_topics()
1610 .is_empty()
1611 );
1612 assert!(
1613 handler
1614 .subscriptions_state
1615 .pending_unsubscribe_topics()
1616 .is_empty()
1617 );
1618 }
1619
1620 #[derive(serde::Deserialize, Debug, PartialEq)]
1621 struct ParseArrayItem {
1622 value: i64,
1623 }
1624
1625 #[rstest]
1626 fn test_parse_array_items_keeps_good_items_when_one_fails() {
1627 let data = json!([
1628 {"value": 1},
1629 {"value": "not a number"},
1630 {"value": 3},
1631 ]);
1632
1633 let parsed: Vec<ParseArrayItem> =
1634 parse_array_items(data, "test", false).expect("non-empty");
1635 assert_eq!(
1636 parsed,
1637 vec![ParseArrayItem { value: 1 }, ParseArrayItem { value: 3 }],
1638 );
1639 }
1640
1641 #[rstest]
1642 fn test_parse_array_items_returns_none_when_payload_not_array() {
1643 let data = json!({"not": "an array"});
1644 let parsed: Option<Vec<ParseArrayItem>> = parse_array_items(data, "test", false);
1645 assert!(parsed.is_none());
1646 }
1647
1648 #[rstest]
1649 fn test_parse_array_items_returns_none_when_all_items_fail() {
1650 let data = json!([{"value": "bad"}]);
1651 let parsed: Option<Vec<ParseArrayItem>> = parse_array_items(data, "test", false);
1652 assert!(parsed.is_none());
1653 }
1654
1655 #[rstest]
1656 fn test_route_instruments_keeps_valid_items_when_one_item_fails() {
1657 let handler = create_handler();
1658 let mut frame: Value =
1659 serde_json::from_str(&load_test_json("ws_instruments.json")).expect("valid fixture");
1660 let data = frame
1661 .get_mut("data")
1662 .and_then(Value::as_array_mut)
1663 .expect("data array");
1664 let mut invalid_item = data[0].clone();
1665 invalid_item["tickSz"] = json!(7);
1666 data.insert(0, invalid_item);
1667
1668 let arg: OKXWebSocketArg = serde_json::from_value(frame["arg"].clone()).expect("valid arg");
1669 let msg = handler
1670 .route_data_message(arg, frame["data"].clone())
1671 .expect("instruments message");
1672
1673 match msg {
1674 OKXWsMessage::Instruments(instruments) => {
1675 assert_eq!(instruments.len(), 1);
1676 assert_eq!(instruments[0].inst_id, "BTC-USDT-SWAP");
1677 }
1678 other => panic!("Expected Instruments, was {other:?}"),
1679 }
1680 }
1681
1682 #[rstest]
1683 fn test_route_instruments_prefers_rpi_over_legacy_alias() {
1684 let handler = create_handler();
1685 let mut frame: Value =
1686 serde_json::from_str(&load_test_json("ws_instruments.json")).expect("valid fixture");
1687 let instrument = &mut frame["data"][0];
1688 instrument["rpi"] = json!("2");
1689 instrument["elp"] = json!("1");
1690
1691 let arg: OKXWebSocketArg = serde_json::from_value(frame["arg"].clone()).expect("valid arg");
1692 let msg = handler
1693 .route_data_message(arg, frame["data"].clone())
1694 .expect("instruments message");
1695
1696 match msg {
1697 OKXWsMessage::Instruments(instruments) => {
1698 assert_eq!(instruments.len(), 1);
1699 assert_eq!(instruments[0].rpi, Some(OKXRpiPermission::Permitted));
1700 }
1701 other => panic!("Expected Instruments, was {other:?}"),
1702 }
1703 }
1704
1705 #[rstest]
1706 fn test_route_liquidation_warnings() {
1707 let handler = create_handler();
1708 let frame: Value = serde_json::from_str(&load_test_json("ws_liquidation_warning.json"))
1709 .expect("valid fixture");
1710
1711 let arg: OKXWebSocketArg = serde_json::from_value(frame["arg"].clone()).expect("valid arg");
1712 let msg = handler
1713 .route_data_message(arg, frame["data"].clone())
1714 .expect("liquidation warning message");
1715
1716 match msg {
1717 OKXWsMessage::LiquidationWarnings(warnings) => {
1718 assert_eq!(warnings.len(), 1);
1719 assert_eq!(warnings[0].inst_id, "BTC-USDT-SWAP");
1720 assert_eq!(warnings[0].mgn_ratio, "0.62");
1721 }
1722 other => panic!("Expected LiquidationWarnings, was {other:?}"),
1723 }
1724 }
1725}