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};
40use nautilus_model::identifiers::ClientOrderId;
41use nautilus_network::{
42 RECONNECTED,
43 error::SendError,
44 retry::{RetryError, RetryManager, create_websocket_retry_manager},
45 websocket::{AuthTracker, SubscriptionState, TEXT_PING, TEXT_PONG, WebSocketClient},
46};
47use serde_json::{Map, Value};
48use tokio_tungstenite::tungstenite::Message;
49use ustr::Ustr;
50
51use super::{
52 enums::{OKXSubscriptionEvent, OKXWsChannel, OKXWsOperation},
53 error::OKXWsError,
54 messages::{
55 OKXOrderMsg, OKXSubscription, OKXSubscriptionArg, OKXWebSocketArg, OKXWebSocketError,
56 OKXWsFrame, OKXWsMessage,
57 },
58 subscription::topic_from_websocket_arg,
59};
60use crate::{
61 common::{
62 consts::{OKX_FIELD_SMSG, OKX_SUCCESS_CODE, should_retry_error_code},
63 enums::{OKXOrderStatus, OKXOrderType},
64 parse::prefer_rpi_response_fields,
65 },
66 websocket::client::OKX_RATE_LIMIT_KEY_SUBSCRIPTION,
67};
68
69pub enum HandlerCommand {
71 SetClient(WebSocketClient),
73 Disconnect,
75 Authenticate { payload: SecretString },
77 Subscribe { args: Vec<OKXSubscriptionArg> },
79 Unsubscribe { args: Vec<OKXSubscriptionArg> },
81 Send {
83 payload: String,
84 rate_limit_keys: Option<Vec<Ustr>>,
85 request_id: Option<String>,
86 client_order_ids: Vec<ClientOrderId>,
87 op: Option<OKXWsOperation>,
88 },
89}
90
91impl Debug for HandlerCommand {
92 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
93 match self {
94 Self::SetClient(_) => f.write_str("SetClient"),
95 Self::Disconnect => f.write_str("Disconnect"),
96 Self::Authenticate { .. } => f
97 .debug_struct(stringify!(Authenticate))
98 .field("payload", &REDACTED)
99 .finish(),
100 Self::Subscribe { args } => f
101 .debug_struct(stringify!(Subscribe))
102 .field("args", args)
103 .finish(),
104 Self::Unsubscribe { args } => f
105 .debug_struct(stringify!(Unsubscribe))
106 .field("args", args)
107 .finish(),
108 Self::Send {
109 rate_limit_keys,
110 request_id,
111 client_order_ids,
112 op,
113 ..
114 } => f
115 .debug_struct(stringify!(Send))
116 .field("payload", &REDACTED)
117 .field("rate_limit_keys", rate_limit_keys)
118 .field("request_id", request_id)
119 .field("client_order_ids", client_order_ids)
120 .field("op", op)
121 .finish(),
122 }
123 }
124}
125
126pub(super) struct OKXWsFeedHandler {
127 clock: &'static AtomicTime,
128 signal: Arc<AtomicBool>,
129 inner: Option<WebSocketClient>,
130 cmd_rx: tokio::sync::mpsc::UnboundedReceiver<HandlerCommand>,
131 raw_rx: tokio::sync::mpsc::UnboundedReceiver<Message>,
132 out_tx: tokio::sync::mpsc::UnboundedSender<OKXWsMessage>,
133 auth_tracker: AuthTracker,
134 subscriptions_state: SubscriptionState,
135 retry_manager: RetryManager<OKXWsError>,
136 pending_messages: VecDeque<OKXWsMessage>,
137}
138
139impl OKXWsFeedHandler {
140 pub(super) fn new(
142 signal: Arc<AtomicBool>,
143 cmd_rx: tokio::sync::mpsc::UnboundedReceiver<HandlerCommand>,
144 raw_rx: tokio::sync::mpsc::UnboundedReceiver<Message>,
145 out_tx: tokio::sync::mpsc::UnboundedSender<OKXWsMessage>,
146 auth_tracker: AuthTracker,
147 subscriptions_state: SubscriptionState,
148 clock: &'static AtomicTime,
149 ) -> Self {
150 Self {
151 clock,
152 signal,
153 inner: None,
154 cmd_rx,
155 raw_rx,
156 out_tx,
157 auth_tracker,
158 subscriptions_state,
159 retry_manager: create_websocket_retry_manager(),
160 pending_messages: VecDeque::new(),
161 }
162 }
163
164 pub(super) fn is_stopped(&self) -> bool {
165 self.signal.load(Ordering::Acquire)
166 }
167
168 pub(super) fn send(&self, msg: OKXWsMessage) -> Result<(), ()> {
169 self.out_tx.send(msg).map_err(|_| ())
170 }
171
172 async fn send_with_retry(
173 &self,
174 payload: String,
175 rate_limit_keys: Option<&[Ustr]>,
176 ) -> Result<(), OKXWsError> {
177 self.send_secret_with_retry(payload.into(), rate_limit_keys)
178 .await
179 }
180
181 async fn send_secret_with_retry(
182 &self,
183 payload: SecretString,
184 rate_limit_keys: Option<&[Ustr]>,
185 ) -> Result<(), OKXWsError> {
186 if let Some(client) = &self.inner {
187 let keys_owned: Option<Vec<Ustr>> = rate_limit_keys.map(<[Ustr]>::to_vec);
188 self.retry_manager
189 .invocation(
190 "websocket_send",
191 || {
192 let payload = payload.clone();
193 let keys = keys_owned.clone();
194 async move {
195 client
196 .send_text(payload.expose_secret().to_owned(), keys.as_deref())
197 .await
198 .map_err(OKXWsError::TransportSend)
199 }
200 },
201 should_retry_replay_safe_error,
202 create_okx_retry_error,
203 )
204 .execute()
205 .await
206 } else {
207 Err(OKXWsError::NoActiveClient)
208 }
209 }
210
211 async fn send_on_connection(
212 &self,
213 payload: String,
214 rate_limit_keys: Option<&[Ustr]>,
215 ) -> Result<(), OKXWsError> {
216 let client = self.inner.as_ref().ok_or(OKXWsError::NoActiveClient)?;
217 let connection_epoch = client.connection_epoch();
218 client
219 .send_text_on_connection(payload, rate_limit_keys, connection_epoch)
220 .await
221 .map_err(OKXWsError::TransportSend)
222 }
223
224 pub(super) async fn send_pong(&self) -> anyhow::Result<()> {
225 match self.send_on_connection(TEXT_PONG.to_string(), None).await {
226 Ok(()) => {
227 log::trace!("Sent pong response to OKX text ping");
228 Ok(())
229 }
230 Err(e) => {
231 log::warn!("Failed to send pong: error={e}");
232 Err(anyhow::anyhow!("Failed to send pong: {e}"))
233 }
234 }
235 }
236
237 pub(super) async fn next(&mut self) -> Option<OKXWsMessage> {
238 if let Some(message) = self.pending_messages.pop_front() {
239 return Some(message);
240 }
241
242 let mut poll_raw_next = false;
243
244 loop {
245 if self.signal.load(Ordering::Acquire) {
246 log::debug!("Stop signal received");
247 return None;
248 }
249
250 tokio::select! {
251 biased;
252 Some(cmd) = self.cmd_rx.recv(), if !poll_raw_next => {
253 match cmd {
254 HandlerCommand::SetClient(client) => {
255 log::debug!("Handler received WebSocket client");
256 self.inner = Some(client);
257 }
258 HandlerCommand::Disconnect => {
259 log::debug!("Handler disconnecting WebSocket client");
260 self.inner = None;
261 return None;
262 }
263 HandlerCommand::Authenticate { payload } => {
264 if let Err(e) = self.send_secret_with_retry(
265 payload,
266 Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()),
267 ).await {
268 log::error!(
269 "Failed to send authentication message after retries: error={e}"
270 );
271 }
272 }
273 HandlerCommand::Subscribe { args } => {
274 if let Err(e) = self.handle_subscribe(args).await {
275 log::error!("Failed to handle subscribe command: error={e}");
276 }
277 }
278 HandlerCommand::Unsubscribe { args } => {
279 if let Err(e) = self.handle_unsubscribe(args).await {
280 log::error!("Failed to handle unsubscribe command: error={e}");
281 }
282 }
283 HandlerCommand::Send {
284 payload,
285 rate_limit_keys,
286 request_id,
287 client_order_ids,
288 op,
289 } => {
290 if let Err(e) = self.send_on_connection(
291 payload,
292 rate_limit_keys.as_deref(),
293 ).await {
294 log::error!("Failed to send message: error={e}");
295
296 if let Some(request_id) = request_id {
297 self.pending_messages.push_back(OKXWsMessage::SendFailed {
298 request_id,
299 client_order_ids,
300 op,
301 error: e,
302 });
303 }
304 }
305 }
306 }
307
308 poll_raw_next = true;
309 }
310
311 () = time::sleep(time::Duration::from_millis(100)) => {
312 }
314
315 msg = self.raw_rx.recv() => {
316 let event = match msg {
317 Some(msg) => match Self::parse_raw_message(msg) {
318 Some(event) => event,
319 None => continue,
320 },
321 None => {
322 log::debug!("WebSocket stream closed");
323 return None;
324 }
325 };
326
327 match event {
328 OKXWsFrame::Ping => {
329 if let Err(e) = self.send_pong().await {
330 log::warn!("Failed to send pong response: error={e}");
331 }
332 }
333 OKXWsFrame::Login {
334 code, msg, conn_id, ..
335 } => {
336 if code == OKX_SUCCESS_CODE {
337 self.auth_tracker.succeed();
338 return Some(OKXWsMessage::Authenticated);
339 }
340
341 log::error!("WebSocket authentication failed: error={msg}");
342 self.auth_tracker.fail(msg.clone());
343
344 let error = OKXWebSocketError {
345 code,
346 message: msg,
347 conn_id: Some(conn_id),
348 timestamp: self.clock.get_time_ns().as_u64(),
349 };
350 self.pending_messages.push_back(OKXWsMessage::Error(error));
351 }
352 OKXWsFrame::BookData { arg, action, data } => {
353 return Some(OKXWsMessage::BookData { arg, action, data });
354 }
355 OKXWsFrame::RpiBookData { arg, action, data } => {
356 return Some(OKXWsMessage::RpiBookData { arg, action, data });
357 }
358 OKXWsFrame::OrderResponse {
359 id, op, code, msg, data,
360 } => {
361 return Some(OKXWsMessage::OrderResponse {
362 id, op, code, msg, data,
363 });
364 }
365 OKXWsFrame::Data { arg, data } => {
366 if let Some(output) = self.route_data_message(arg, data) {
367 return Some(output);
368 }
369 }
370 OKXWsFrame::Error { arg, code, msg } => {
371 let arg = arg.or_else(|| subscription_arg_from_error_message(&msg));
372 if let Some(arg) = arg
373 && self.handle_subscription_error(&arg, &code, &msg)
374 {
375 return Some(OKXWsMessage::SubscriptionFailed {
376 channel: arg.channel,
377 inst_id: arg.inst_id,
378 code,
379 msg,
380 });
381 }
382
383 let error = OKXWebSocketError {
384 code,
385 message: msg,
386 conn_id: None,
387 timestamp: self.clock.get_time_ns().as_u64(),
388 };
389 return Some(OKXWsMessage::Error(error));
390 }
391 OKXWsFrame::Reconnected => {
392 self.auth_tracker.invalidate();
393 return Some(OKXWsMessage::Reconnected);
394 }
395 OKXWsFrame::Subscription {
396 event, arg, code, msg,
397 ..
398 } => {
399 let rejected = self
400 .handle_subscription_ack(&event, &arg, code.as_deref(), msg.as_deref());
401
402 if rejected {
403 return Some(OKXWsMessage::SubscriptionFailed {
404 channel: arg.channel,
405 inst_id: arg.inst_id,
406 code: code.unwrap_or_default(),
407 msg: msg.unwrap_or_default(),
408 });
409 }
410 }
411 OKXWsFrame::ChannelConnCount { .. } => {}
412 }
413 }
414
415 () = std::future::ready(()), if poll_raw_next => {
416 poll_raw_next = false;
417 }
418
419 else => {
420 log::debug!("Handler shutting down: stream ended or command channel closed");
421 return None;
422 }
423 }
424 }
425 }
426
427 fn route_data_message(&self, arg: OKXWebSocketArg, mut data: Value) -> Option<OKXWsMessage> {
428 let OKXWebSocketArg {
429 channel, inst_id, ..
430 } = arg;
431
432 match channel {
433 OKXWsChannel::Account => Some(OKXWsMessage::Account(data)),
434 OKXWsChannel::Positions => Some(OKXWsMessage::Positions(data)),
435 OKXWsChannel::Orders => {
436 parse_array_items(data, "orders", false).map(OKXWsMessage::Orders)
437 }
438 OKXWsChannel::SprdOrders => {
439 parse_array_items(data, "spread orders", false).map(OKXWsMessage::SpreadOrders)
440 }
441 OKXWsChannel::OrdersAlgo | OKXWsChannel::AlgoAdvance => {
442 parse_array_items(data, "algo orders", false).map(OKXWsMessage::AlgoOrders)
443 }
444 OKXWsChannel::LiquidationWarning => {
445 parse_array_items(data, "liquidation warnings", false)
446 .map(OKXWsMessage::LiquidationWarnings)
447 }
448 OKXWsChannel::Instruments => {
449 prefer_rpi_response_fields(&mut data);
450 parse_array_items(data, "instruments", true).map(OKXWsMessage::Instruments)
451 }
452 _ => Some(OKXWsMessage::ChannelData {
453 channel,
454 inst_id,
455 data,
456 }),
457 }
458 }
459
460 fn handle_subscription_ack(
461 &self,
462 event: &OKXSubscriptionEvent,
463 arg: &OKXWebSocketArg,
464 code: Option<&str>,
465 msg: Option<&str>,
466 ) -> bool {
467 let topic = topic_from_websocket_arg(arg);
468 let success = code.is_none_or(|c| c == OKX_SUCCESS_CODE);
469
470 match event {
471 OKXSubscriptionEvent::Subscribe => {
472 if success {
473 self.subscriptions_state.confirm_subscribe(&topic);
474 false
475 } else {
476 log::warn!(
477 "Subscription failed: topic={topic:?}, error={msg:?}, code={code:?}"
478 );
479 self.subscriptions_state.mark_failure(&topic);
480 true
481 }
482 }
483 OKXSubscriptionEvent::Unsubscribe => {
484 if success {
485 self.subscriptions_state.confirm_unsubscribe(&topic);
486 } else {
487 log::warn!(
488 "Unsubscription failed - restoring subscription: \
489 topic={topic:?}, error={msg:?}, code={code:?}"
490 );
491 self.subscriptions_state.confirm_unsubscribe(&topic);
492 self.subscriptions_state.mark_subscribe(&topic);
493 self.subscriptions_state.confirm_subscribe(&topic);
494 }
495 false
496 }
497 }
498 }
499
500 fn handle_subscription_error(&self, arg: &OKXWebSocketArg, code: &str, msg: &str) -> bool {
501 let topic = topic_from_websocket_arg(arg);
502 let event = if self
503 .subscriptions_state
504 .pending_unsubscribe_topics()
505 .iter()
506 .any(|pending| pending == &topic)
507 {
508 OKXSubscriptionEvent::Unsubscribe
509 } else if self
510 .subscriptions_state
511 .pending_subscribe_topics()
512 .iter()
513 .any(|pending| pending == &topic)
514 {
515 OKXSubscriptionEvent::Subscribe
516 } else {
517 return false;
518 };
519
520 self.handle_subscription_ack(&event, arg, Some(code), Some(msg))
521 }
522
523 async fn handle_subscribe(&self, args: Vec<OKXSubscriptionArg>) -> anyhow::Result<()> {
524 for arg in &args {
525 log::debug!(
526 "Subscribing to channel: channel={:?}, inst_id={:?}",
527 arg.channel,
528 arg.inst_id
529 );
530 }
531
532 let message = OKXSubscription {
533 op: OKXWsOperation::Subscribe,
534 args,
535 };
536
537 let json_txt = serde_json::to_string(&message)
538 .map_err(|e| anyhow::anyhow!("Failed to serialize subscription: {e}"))?;
539
540 self.send_with_retry(json_txt, Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()))
541 .await
542 .map_err(|e| anyhow::anyhow!("Failed to send subscription after retries: {e}"))?;
543 Ok(())
544 }
545
546 async fn handle_unsubscribe(&self, args: Vec<OKXSubscriptionArg>) -> anyhow::Result<()> {
547 for arg in &args {
548 log::debug!(
549 "Unsubscribing from channel: channel={:?}, inst_id={:?}",
550 arg.channel,
551 arg.inst_id
552 );
553 }
554
555 let message = OKXSubscription {
556 op: OKXWsOperation::Unsubscribe,
557 args,
558 };
559
560 let json_txt = serde_json::to_string(&message)
561 .map_err(|e| anyhow::anyhow!("Failed to serialize unsubscription: {e}"))?;
562
563 self.send_with_retry(json_txt, Some(OKX_RATE_LIMIT_KEY_SUBSCRIPTION.as_slice()))
564 .await
565 .map_err(|e| anyhow::anyhow!("Failed to send unsubscription after retries: {e}"))?;
566 Ok(())
567 }
568
569 pub(crate) fn parse_raw_message(
570 msg: tokio_tungstenite::tungstenite::Message,
571 ) -> Option<OKXWsFrame> {
572 match msg {
573 tokio_tungstenite::tungstenite::Message::Text(text) => {
574 if text == TEXT_PONG {
575 log::trace!("Received pong from OKX");
576 return None;
577 }
578
579 if text == TEXT_PING {
580 log::trace!("Received ping from OKX (text)");
581 return Some(OKXWsFrame::Ping);
582 }
583
584 if text == RECONNECTED {
585 log::debug!("Received WebSocket reconnection signal");
586 return Some(OKXWsFrame::Reconnected);
587 }
588 log::trace!("Received WebSocket message: {text}");
589
590 match serde_json::from_str(&text) {
591 Ok(ws_event) => match &ws_event {
592 OKXWsFrame::Error { code, msg, .. } => {
593 if should_retry_error_code(code) {
594 log::warn!("WebSocket error: {code} - {msg}");
595 } else {
596 log::error!("WebSocket error: {code} - {msg}");
597 }
598 Some(ws_event)
599 }
600 OKXWsFrame::Login {
601 event,
602 code,
603 msg,
604 conn_id,
605 } => {
606 if code == OKX_SUCCESS_CODE {
607 log::debug!("WebSocket authenticated: conn_id={conn_id}");
608 } else {
609 log::error!(
610 "WebSocket authentication failed: \
611 event={event}, code={code}, error={msg}"
612 );
613 }
614 Some(ws_event)
615 }
616 OKXWsFrame::Subscription {
617 event,
618 arg,
619 conn_id,
620 ..
621 } => {
622 let channel_str = serde_json::to_string(&arg.channel)
623 .expect("Invalid OKX websocket channel")
624 .trim_matches('"')
625 .to_string();
626 log::debug!("{event}d: channel={channel_str}, conn_id={conn_id}");
627 Some(ws_event)
628 }
629 OKXWsFrame::ChannelConnCount {
630 channel,
631 conn_count,
632 conn_id,
633 ..
634 } => {
635 let channel_str = serde_json::to_string(channel)
636 .expect("Invalid OKX websocket channel")
637 .trim_matches('"')
638 .to_string();
639 log::debug!(
640 "Channel connection status: \
641 channel={channel_str}, connections={conn_count}, conn_id={conn_id}",
642 );
643 None
644 }
645 OKXWsFrame::Ping => {
646 log::trace!("Ignoring ping event parsed from text payload");
647 None
648 }
649 OKXWsFrame::Data { .. }
650 | OKXWsFrame::BookData { .. }
651 | OKXWsFrame::RpiBookData { .. } => Some(ws_event),
652 OKXWsFrame::OrderResponse {
653 id, op, code, data, ..
654 } => {
655 if code == OKX_SUCCESS_CODE {
656 log::debug!(
657 "Order operation successful: id={id:?}, op={op}, code={code}"
658 );
659
660 if let Some(order_data) = data.first() {
661 let success_msg = order_data
662 .get(OKX_FIELD_SMSG)
663 .and_then(|s| s.as_str())
664 .unwrap_or("Order operation successful");
665 log::debug!("Order success details: {success_msg}");
666 }
667 }
668 Some(ws_event)
669 }
670 OKXWsFrame::Reconnected => {
671 log::warn!("Unexpected Reconnected event from deserialization");
672 None
673 }
674 },
675 Err(e) => {
676 log::error!("Failed to parse message: {e}: {text}");
677 None
678 }
679 }
680 }
681 Message::Ping(_payload) => {
682 log::trace!("Received binary ping frame from OKX");
683 Some(OKXWsFrame::Ping)
684 }
685 Message::Pong(payload) => {
686 log::trace!("Received pong frame from OKX ({} bytes)", payload.len());
687 None
688 }
689 Message::Binary(msg) => {
690 log::debug!("Raw binary frame ({} bytes)", msg.len());
691 log::trace!("Raw binary: {msg:?}");
692 None
693 }
694 Message::Close(_) => {
695 log::debug!("Received close message");
696 None
697 }
698 msg => {
699 log::warn!("Unexpected message: {msg}");
700 None
701 }
702 }
703 }
704}
705
706fn subscription_arg_from_error_message(msg: &str) -> Option<OKXWebSocketArg> {
707 let descriptor = msg
708 .strip_prefix("Wrong URL or channel:")?
709 .split_whitespace()
710 .next()?;
711 let mut fields = descriptor.split(',');
712 let channel = fields.next()?;
713 let mut arg = Map::new();
714 arg.insert("channel".to_string(), Value::String(channel.to_string()));
715
716 for field in fields {
717 let (key, value) = field.split_once(':')?;
718 if !matches!(key, "instId" | "sprdId" | "instType" | "instFamily") {
719 return None;
720 }
721 arg.insert(key.to_string(), Value::String(value.to_string()));
722 }
723
724 serde_json::from_value(Value::Object(arg)).ok()
725}
726
727pub fn is_post_only_auto_cancel(msg: &OKXOrderMsg) -> bool {
729 use crate::common::{consts::OKX_POST_ONLY_CANCEL_SOURCE, enums::OKXOrderStatus};
730
731 if msg.state != OKXOrderStatus::Canceled {
732 return false;
733 }
734
735 let cancel_source_matches = matches!(
736 msg.cancel_source.as_deref(),
737 Some(source) if source == OKX_POST_ONLY_CANCEL_SOURCE
738 );
739
740 let reason_matches = matches!(
741 msg.cancel_source_reason.as_deref(),
742 Some(reason) if reason.contains("POST_ONLY")
743 );
744
745 if !(cancel_source_matches || reason_matches) {
746 return false;
747 }
748
749 msg.acc_fill_sz
750 .as_ref()
751 .is_none_or(|filled| filled == "0" || filled.is_empty())
752}
753
754pub fn is_unfilled_rpi_cancel(msg: &OKXOrderMsg) -> bool {
756 msg.ord_type == OKXOrderType::Rpi
757 && msg.state == OKXOrderStatus::Canceled
758 && msg
759 .acc_fill_sz
760 .as_ref()
761 .is_none_or(|filled| filled == "0" || filled.is_empty())
762}
763
764fn parse_array_items<T: serde::de::DeserializeOwned>(
766 data: Value,
767 label: &str,
768 warn_on_parse_error: bool,
769) -> Option<Vec<T>> {
770 let Value::Array(items) = data else {
771 if warn_on_parse_error {
772 log::warn!("Expected {label} payload to be a JSON array");
773 } else {
774 log::error!("Expected {label} payload to be a JSON array");
775 }
776 return None;
777 };
778
779 let mut parsed = Vec::with_capacity(items.len());
780 for (idx, item) in items.into_iter().enumerate() {
781 match serde_json::from_value::<T>(item) {
782 Ok(value) => parsed.push(value),
783 Err(e) => {
784 if warn_on_parse_error {
785 log::warn!("Failed to parse {label} item at index {idx}: {e}");
786 } else {
787 log::error!("Failed to parse {label} item at index {idx}: {e}");
788 }
789 }
790 }
791 }
792
793 if parsed.is_empty() {
794 None
795 } else {
796 Some(parsed)
797 }
798}
799
800fn should_retry_replay_safe_error(error: &OKXWsError) -> bool {
801 match error {
802 OKXWsError::OkxError { error_code, .. } => should_retry_error_code(error_code),
803 OKXWsError::TransportSend(SendError::Timeout | SendError::ConnectionChanged)
804 | OKXWsError::TungsteniteError(_)
805 | OKXWsError::OperationTimeout { .. } => true,
806 OKXWsError::AuthenticationError(_)
807 | OKXWsError::JsonError(_)
808 | OKXWsError::ParsingError(_)
809 | OKXWsError::ClientError(_)
810 | OKXWsError::NoActiveClient
811 | OKXWsError::HandlerUnavailable(_)
812 | OKXWsError::TransportSend(
813 SendError::InvalidInput(_)
814 | SendError::Closed
815 | SendError::WriteTimeout
816 | SendError::BrokenPipe(_),
817 )
818 | OKXWsError::SendFailed(_) => false,
819 }
820}
821
822fn create_okx_retry_error(error: RetryError) -> OKXWsError {
823 match error {
824 RetryError::OperationTimeout { timeout_ms } => OKXWsError::OperationTimeout { timeout_ms },
825 RetryError::InvalidConfiguration { message } => OKXWsError::ClientError(message),
826 RetryError::Canceled => {
827 OKXWsError::SendFailed("Adapter disconnecting or shutting down".to_string())
828 }
829 error @ RetryError::ElapsedBudgetExceeded { .. } => {
830 OKXWsError::SendFailed(error.to_string())
831 }
832 }
833}
834
835#[cfg(test)]
836mod tests {
837 use std::sync::{Arc, atomic::AtomicBool};
838
839 use nautilus_core::time::get_atomic_clock_realtime;
840 use nautilus_network::websocket::{AuthTracker, SubscriptionState};
841 use rstest::rstest;
842 use serde_json::json;
843
844 use super::*;
845 use crate::common::{
846 consts::OKX_WS_TOPIC_DELIMITER, enums::OKXRpiPermission, testing::load_test_json,
847 };
848
849 fn create_handler() -> OKXWsFeedHandler {
850 let signal = Arc::new(AtomicBool::new(false));
851 let (_cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
852 let (_raw_tx, raw_rx) = tokio::sync::mpsc::unbounded_channel();
853 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
854
855 OKXWsFeedHandler::new(
856 signal,
857 cmd_rx,
858 raw_rx,
859 out_tx,
860 AuthTracker::new(),
861 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
862 get_atomic_clock_realtime(),
863 )
864 }
865
866 #[rstest]
867 fn test_command_debug_redacts_payloads() {
868 let payload = "authentication-secret";
869 let authenticate = HandlerCommand::Authenticate {
870 payload: SecretString::from(payload.to_string()),
871 };
872 let send = HandlerCommand::Send {
873 payload: payload.to_string(),
874 rate_limit_keys: None,
875 request_id: None,
876 client_order_ids: Vec::new(),
877 op: None,
878 };
879
880 let debug = format!("{authenticate:?} {send:?}");
881
882 assert!(debug.contains(REDACTED));
883 assert!(!debug.contains(payload));
884 }
885
886 #[tokio::test]
887 async fn test_next_polls_raw_after_one_ready_command() {
888 let signal = Arc::new(AtomicBool::new(false));
889 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
890 let (raw_tx, raw_rx) = tokio::sync::mpsc::unbounded_channel();
891 let (out_tx, _out_rx) = tokio::sync::mpsc::unbounded_channel();
892 let mut handler = OKXWsFeedHandler::new(
893 signal,
894 cmd_rx,
895 raw_rx,
896 out_tx,
897 AuthTracker::new(),
898 SubscriptionState::new(OKX_WS_TOPIC_DELIMITER),
899 get_atomic_clock_realtime(),
900 );
901
902 for _ in 0..3 {
903 cmd_tx
904 .send(HandlerCommand::Subscribe { args: Vec::new() })
905 .unwrap();
906 }
907 raw_tx
908 .send(Message::Text(RECONNECTED.to_string().into()))
909 .unwrap();
910
911 let message = handler.next().await;
912
913 assert!(matches!(message, Some(OKXWsMessage::Reconnected)));
914 assert_eq!(handler.cmd_rx.len(), 2);
915 }
916
917 #[rstest]
918 fn test_should_retry_typed_transport_and_timeout_errors() {
919 assert!(should_retry_replay_safe_error(&OKXWsError::TransportSend(
920 SendError::Timeout
921 )));
922 assert!(should_retry_replay_safe_error(&OKXWsError::TransportSend(
923 SendError::ConnectionChanged
924 )));
925 assert!(!should_retry_replay_safe_error(&OKXWsError::TransportSend(
926 SendError::WriteTimeout
927 )));
928 assert!(!should_retry_replay_safe_error(&OKXWsError::TransportSend(
929 SendError::BrokenPipe("connection reset".to_string())
930 )));
931 assert!(should_retry_replay_safe_error(
932 &OKXWsError::OperationTimeout { timeout_ms: 1_000 }
933 ));
934 assert!(!should_retry_replay_safe_error(&OKXWsError::NoActiveClient));
935 assert!(!should_retry_replay_safe_error(
936 &OKXWsError::HandlerUnavailable("closed".to_string())
937 ));
938 }
939
940 #[rstest]
941 fn test_retryability_uses_websocket_error_type_not_message() {
942 let message = "connection reset".to_string();
943 let temporary = OKXWsError::OkxError {
944 error_code: "50011".to_string(),
945 message: message.clone(),
946 };
947 let permanent = OKXWsError::ClientError(message.clone());
948 let ambiguous = OKXWsError::SendFailed(message);
949
950 assert!(should_retry_replay_safe_error(&temporary));
951 assert!(!should_retry_replay_safe_error(&permanent));
952 assert!(!should_retry_replay_safe_error(&ambiguous));
953 }
954
955 #[rstest]
956 fn test_subscription_error_restores_failed_unsubscribe() {
957 let handler = create_handler();
958 let arg = OKXWebSocketArg {
959 channel: OKXWsChannel::Books,
960 inst_id: Some(Ustr::from("BTC-USD")),
961 inst_type: None,
962 inst_family: None,
963 bar: None,
964 };
965 let topic = topic_from_websocket_arg(&arg);
966 handler.subscriptions_state.mark_subscribe(&topic);
967 handler.subscriptions_state.confirm_subscribe(&topic);
968 handler.subscriptions_state.mark_unsubscribe(&topic);
969
970 let rejected_subscription =
971 handler.handle_subscription_error(&arg, "60019", "Unsubscription failed");
972
973 assert!(!rejected_subscription);
974 assert_eq!(handler.subscriptions_state.all_topics(), vec![topic]);
975 assert!(
976 handler
977 .subscriptions_state
978 .pending_subscribe_topics()
979 .is_empty()
980 );
981 assert!(
982 handler
983 .subscriptions_state
984 .pending_unsubscribe_topics()
985 .is_empty()
986 );
987 }
988
989 #[rstest]
990 fn test_subscription_arg_from_error_message_matches_mainnet_shape() {
991 let msg = "Wrong URL or channel:books,instId:BTC-USDT-SWAP doesn't exist. Please use the \
992 correct URL, channel and parameters referring to API document.";
993
994 let arg = subscription_arg_from_error_message(msg).unwrap();
995
996 assert_eq!(arg.channel, OKXWsChannel::Books);
997 assert_eq!(arg.inst_id, Some(Ustr::from("BTC-USDT-SWAP")));
998 assert_eq!(arg.inst_type, None);
999 assert_eq!(arg.inst_family, None);
1000 assert_eq!(arg.bar, None);
1001 }
1002
1003 #[rstest]
1004 fn test_subscription_error_ignores_non_pending_topic() {
1005 let handler = create_handler();
1006 let arg = OKXWebSocketArg {
1007 channel: OKXWsChannel::Books,
1008 inst_id: Some(Ustr::from("BTC-USDT-SWAP")),
1009 inst_type: None,
1010 inst_family: None,
1011 bar: None,
1012 };
1013
1014 let rejected_subscription =
1015 handler.handle_subscription_error(&arg, "60018", "Subscription failed");
1016
1017 assert!(!rejected_subscription);
1018 assert!(handler.subscriptions_state.all_topics().is_empty());
1019 assert!(
1020 handler
1021 .subscriptions_state
1022 .pending_subscribe_topics()
1023 .is_empty()
1024 );
1025 assert!(
1026 handler
1027 .subscriptions_state
1028 .pending_unsubscribe_topics()
1029 .is_empty()
1030 );
1031 }
1032
1033 #[derive(serde::Deserialize, Debug, PartialEq)]
1034 struct ParseArrayItem {
1035 value: i64,
1036 }
1037
1038 #[rstest]
1039 fn test_parse_array_items_keeps_good_items_when_one_fails() {
1040 let data = json!([
1041 {"value": 1},
1042 {"value": "not a number"},
1043 {"value": 3},
1044 ]);
1045
1046 let parsed: Vec<ParseArrayItem> =
1047 parse_array_items(data, "test", false).expect("non-empty");
1048 assert_eq!(
1049 parsed,
1050 vec![ParseArrayItem { value: 1 }, ParseArrayItem { value: 3 }],
1051 );
1052 }
1053
1054 #[rstest]
1055 fn test_parse_array_items_returns_none_when_payload_not_array() {
1056 let data = json!({"not": "an array"});
1057 let parsed: Option<Vec<ParseArrayItem>> = parse_array_items(data, "test", false);
1058 assert!(parsed.is_none());
1059 }
1060
1061 #[rstest]
1062 fn test_parse_array_items_returns_none_when_all_items_fail() {
1063 let data = json!([{"value": "bad"}]);
1064 let parsed: Option<Vec<ParseArrayItem>> = parse_array_items(data, "test", false);
1065 assert!(parsed.is_none());
1066 }
1067
1068 #[rstest]
1069 fn test_route_instruments_keeps_valid_items_when_one_item_fails() {
1070 let handler = create_handler();
1071 let mut frame: Value =
1072 serde_json::from_str(&load_test_json("ws_instruments.json")).expect("valid fixture");
1073 let data = frame
1074 .get_mut("data")
1075 .and_then(Value::as_array_mut)
1076 .expect("data array");
1077 let mut invalid_item = data[0].clone();
1078 invalid_item["tickSz"] = json!(7);
1079 data.insert(0, invalid_item);
1080
1081 let arg: OKXWebSocketArg = serde_json::from_value(frame["arg"].clone()).expect("valid arg");
1082 let msg = handler
1083 .route_data_message(arg, frame["data"].clone())
1084 .expect("instruments message");
1085
1086 match msg {
1087 OKXWsMessage::Instruments(instruments) => {
1088 assert_eq!(instruments.len(), 1);
1089 assert_eq!(instruments[0].inst_id, "BTC-USDT-SWAP");
1090 }
1091 other => panic!("Expected Instruments, was {other:?}"),
1092 }
1093 }
1094
1095 #[rstest]
1096 fn test_route_instruments_prefers_rpi_over_legacy_alias() {
1097 let handler = create_handler();
1098 let mut frame: Value =
1099 serde_json::from_str(&load_test_json("ws_instruments.json")).expect("valid fixture");
1100 let instrument = &mut frame["data"][0];
1101 instrument["rpi"] = json!("2");
1102 instrument["elp"] = json!("1");
1103
1104 let arg: OKXWebSocketArg = serde_json::from_value(frame["arg"].clone()).expect("valid arg");
1105 let msg = handler
1106 .route_data_message(arg, frame["data"].clone())
1107 .expect("instruments message");
1108
1109 match msg {
1110 OKXWsMessage::Instruments(instruments) => {
1111 assert_eq!(instruments.len(), 1);
1112 assert_eq!(instruments[0].rpi, Some(OKXRpiPermission::Permitted));
1113 }
1114 other => panic!("Expected Instruments, was {other:?}"),
1115 }
1116 }
1117
1118 #[rstest]
1119 fn test_route_liquidation_warnings() {
1120 let handler = create_handler();
1121 let frame: Value = serde_json::from_str(&load_test_json("ws_liquidation_warning.json"))
1122 .expect("valid fixture");
1123
1124 let arg: OKXWebSocketArg = serde_json::from_value(frame["arg"].clone()).expect("valid arg");
1125 let msg = handler
1126 .route_data_message(arg, frame["data"].clone())
1127 .expect("liquidation warning message");
1128
1129 match msg {
1130 OKXWsMessage::LiquidationWarnings(warnings) => {
1131 assert_eq!(warnings.len(), 1);
1132 assert_eq!(warnings[0].inst_id, "BTC-USDT-SWAP");
1133 assert_eq!(warnings[0].mgn_ratio, "0.62");
1134 }
1135 other => panic!("Expected LiquidationWarnings, was {other:?}"),
1136 }
1137 }
1138}