Skip to main content

nautilus_okx/websocket/
handler.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! WebSocket message handler for OKX.
17//!
18//! The handler is a thin I/O boundary between the network layer and the client. It owns the
19//! `WebSocketClient`, deserializes raw venue messages into `OKXWsMessage` events, and handles
20//! subscription management, authentication, and retry logic.
21//!
22//! All domain parsing (venue types to Nautilus types) occurs outside the handler:
23//! - Data parsing in `PyOKXWebSocketClient` (uses an instruments cache)
24//! - Execution parsing in `execution.rs` (uses the system Cache)
25
26use 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
69/// Commands sent from the outer client to the inner message handler.
70pub enum HandlerCommand {
71    /// Set the `WebSocketClient` for the handler to use.
72    SetClient(WebSocketClient),
73    /// Disconnect the WebSocket connection.
74    Disconnect,
75    /// Send authentication payload to the WebSocket.
76    Authenticate { payload: SecretString },
77    /// Subscribe to the given channels.
78    Subscribe { args: Vec<OKXSubscriptionArg> },
79    /// Unsubscribe from the given channels.
80    Unsubscribe { args: Vec<OKXSubscriptionArg> },
81    /// Send a pre-serialized payload (used for order operations).
82    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    /// Creates a new [`OKXWsFeedHandler`] instance.
141    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                    // Wake the loop to poll the stop signal while both channels are idle
313                }
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
727/// Returns `true` when an OKX WebSocket order message represents a post-only auto-cancel.
728pub 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
754/// Returns `true` when an RPI order update is canceled without any fill.
755pub 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
764// Per-item deserialization so one malformed entry does not drop the batch.
765fn 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}