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};
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
71/// Commands sent from the outer client to the inner message handler.
72pub enum HandlerCommand {
73    /// Set the `WebSocketClient` for the handler to use.
74    SetClient(WebSocketClient),
75    /// Disconnect the WebSocket connection.
76    Disconnect,
77    /// Send authentication payload to the WebSocket.
78    Authenticate { payload: SecretString },
79    /// Subscribe to the given channels.
80    Subscribe { args: Vec<OKXSubscriptionArg> },
81    /// Subscribes to a book and reports completion of the transport send.
82    SubscribeBook {
83        subscription: OKXSubscriptionArg,
84        cancel: CancellationToken,
85        gate: SnapshotGate,
86        completion: tokio::sync::oneshot::Sender<Result<(), OKXWsError>>,
87    },
88    /// Unsubscribe from the given channels.
89    Unsubscribe { args: Vec<OKXSubscriptionArg> },
90    /// Replaces a subscription without removing reconnect intent, unless canceled.
91    Resubscribe {
92        subscription: OKXSubscriptionArg,
93        cancel: CancellationToken,
94        gate: SnapshotGate,
95        completion: tokio::sync::oneshot::Sender<Result<(), OKXWsError>>,
96    },
97    /// Send a pre-serialized payload (used for order operations).
98    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    /// Creates a new [`OKXWsFeedHandler`] instance.
163    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                    // Wake the loop to poll the stop signal while both channels are idle
380                }
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
824/// Returns `true` when an OKX WebSocket order message represents a post-only auto-cancel.
825pub 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
851/// Returns `true` when an RPI order update is canceled without any fill.
852pub 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
861// Per-item deserialization so one malformed entry does not drop the batch.
862fn 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        // Replay and unsubscribe consume the burst; the replacement subscribe cannot send
1333        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        // Paused time keeps the subscribe blocked until the reset completes
1428        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        // A successful manual subscribe isolates a failure to automatic replay
1475        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}