Skip to main content

nautilus_architect_ax/websocket/orders/
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//! Orders WebSocket message handler for Ax.
17
18use std::{
19    collections::VecDeque,
20    sync::{
21        Arc,
22        atomic::{AtomicBool, Ordering},
23    },
24};
25
26use ahash::AHashMap;
27use dashmap::DashMap;
28use nautilus_model::identifiers::{ClientOrderId, VenueOrderId};
29use nautilus_network::websocket::{AuthTracker, WebSocketClient};
30use tokio_tungstenite::tungstenite::Message;
31use ustr::Ustr;
32
33use crate::{
34    common::enums::AxOrderRequestType,
35    websocket::{
36        messages::{
37            AxOrdersWsFrame, AxOrdersWsMessage, AxWsCancelOrder, AxWsError, AxWsGetOpenOrders,
38            AxWsOrderEvent, AxWsOrderResponse, AxWsPlaceOrder, OrderMetadata,
39        },
40        parse::parse_order_message,
41    },
42};
43
44/// Simple tracking info for pending WebSocket orders.
45#[derive(Clone, Debug)]
46pub struct WsOrderInfo {
47    /// Client order ID for correlation.
48    pub client_order_id: ClientOrderId,
49    /// Instrument symbol.
50    pub symbol: Ustr,
51    /// Numeric AX client ID.
52    pub cid: u64,
53}
54
55/// Commands sent from the outer client to the inner orders handler.
56#[derive(Debug)]
57pub enum HandlerCommand {
58    /// Set the WebSocket client for this handler.
59    SetClient(WebSocketClient),
60    /// Disconnect the WebSocket connection.
61    Disconnect,
62    /// Mark the current handshake-authenticated session as ready.
63    SessionAuthenticated,
64    /// Place an order.
65    PlaceOrder {
66        /// Request ID for correlation.
67        request_id: i64,
68        /// Order placement message.
69        order: AxWsPlaceOrder,
70        /// Order info for tracking.
71        order_info: WsOrderInfo,
72    },
73    /// Cancel an order.
74    CancelOrder {
75        /// Request ID for correlation.
76        request_id: i64,
77        /// Order ID to cancel.
78        order_id: String,
79    },
80    /// Get open orders.
81    GetOpenOrders {
82        /// Request ID for correlation.
83        request_id: i64,
84    },
85}
86
87/// Orders feed handler that processes WebSocket messages.
88///
89/// Runs in a dedicated Tokio task and owns the WebSocket client exclusively.
90/// Emits raw venue types for downstream consumers to parse into domain events.
91pub(crate) struct AxOrdersWsFeedHandler {
92    signal: Arc<AtomicBool>,
93    inner: Option<WebSocketClient>,
94    cmd_rx: tokio::sync::mpsc::UnboundedReceiver<HandlerCommand>,
95    raw_rx: tokio::sync::mpsc::UnboundedReceiver<Message>,
96    auth_tracker: AuthTracker,
97    pending_orders: AHashMap<i64, WsOrderInfo>,
98    message_queue: VecDeque<AxOrdersWsMessage>,
99    orders_metadata: Arc<DashMap<ClientOrderId, OrderMetadata>>,
100    venue_to_client_order_id: Arc<DashMap<VenueOrderId, ClientOrderId>>,
101    cid_to_client_order_id: Arc<DashMap<u64, ClientOrderId>>,
102    has_authenticated_session: bool,
103    needs_session_restore: bool,
104}
105
106impl AxOrdersWsFeedHandler {
107    /// Creates a new [`AxOrdersWsFeedHandler`] instance.
108    #[must_use]
109    pub(crate) fn new(
110        signal: Arc<AtomicBool>,
111        cmd_rx: tokio::sync::mpsc::UnboundedReceiver<HandlerCommand>,
112        raw_rx: tokio::sync::mpsc::UnboundedReceiver<Message>,
113        auth_tracker: AuthTracker,
114        orders_metadata: Arc<DashMap<ClientOrderId, OrderMetadata>>,
115        venue_to_client_order_id: Arc<DashMap<VenueOrderId, ClientOrderId>>,
116        cid_to_client_order_id: Arc<DashMap<u64, ClientOrderId>>,
117    ) -> Self {
118        Self {
119            signal,
120            inner: None,
121            cmd_rx,
122            raw_rx,
123            auth_tracker,
124            pending_orders: AHashMap::new(),
125            message_queue: VecDeque::new(),
126            orders_metadata,
127            venue_to_client_order_id,
128            cid_to_client_order_id,
129            has_authenticated_session: false,
130            needs_session_restore: false,
131        }
132    }
133
134    fn restore_authenticated_session(&mut self) {
135        if self.has_authenticated_session {
136            log::debug!("Restoring authenticated session after reconnection");
137
138            // The reconnect handshake has already succeeded with the current Bearer header.
139            self.auth_tracker.succeed();
140            self.message_queue
141                .push_back(AxOrdersWsMessage::Authenticated);
142            log::debug!("Authenticated session restored");
143        } else {
144            log::warn!("Cannot restore authentication before the initial session succeeds");
145        }
146    }
147
148    /// Returns the next message from the handler.
149    ///
150    /// This method blocks until a message is available or the handler is stopped.
151    pub(crate) async fn next(&mut self) -> Option<AxOrdersWsMessage> {
152        loop {
153            if self.needs_session_restore && self.message_queue.is_empty() {
154                self.needs_session_restore = false;
155                self.restore_authenticated_session();
156            }
157
158            if let Some(msg) = self.message_queue.pop_front() {
159                return Some(msg);
160            }
161
162            tokio::select! {
163                Some(cmd) = self.cmd_rx.recv() => {
164                    self.handle_command(cmd).await;
165                }
166
167                () = tokio::time::sleep(std::time::Duration::from_millis(100)) => {
168                    if self.signal.load(Ordering::Acquire) {
169                        log::debug!("Stop signal received during idle period");
170                        return None;
171                    }
172                }
173
174                msg = self.raw_rx.recv() => {
175                    let msg = match msg {
176                        Some(msg) => msg,
177                        None => {
178                            log::debug!("WebSocket stream closed");
179                            return None;
180                        }
181                    };
182
183                    if let Message::Ping(data) = &msg {
184                        log::trace!("Received ping frame with {} bytes", data.len());
185
186                        if let Some(client) = &self.inner
187                            && let Err(e) = client.send_pong(data.to_vec()).await
188                        {
189                            log::warn!("Failed to send pong frame: {e}");
190                        }
191                        continue;
192                    }
193
194                    if let Some(messages) = self.parse_raw_message(msg) {
195                        self.message_queue.extend(messages);
196                    }
197
198                    if self.signal.load(Ordering::Acquire) {
199                        log::debug!("Stop signal received");
200                        return None;
201                    }
202                }
203            }
204        }
205    }
206
207    async fn handle_command(&mut self, cmd: HandlerCommand) {
208        match cmd {
209            HandlerCommand::SetClient(client) => {
210                log::debug!("WebSocketClient received by handler");
211                self.inner = Some(client);
212            }
213            HandlerCommand::Disconnect => {
214                log::debug!("Disconnect command received");
215                self.auth_tracker.fail("Disconnected");
216
217                if let Some(inner) = self.inner.take() {
218                    inner.disconnect().await;
219                }
220            }
221            HandlerCommand::SessionAuthenticated => {
222                log::debug!("Session authenticated command received");
223                self.has_authenticated_session = true;
224                self.auth_tracker.succeed();
225                self.message_queue
226                    .push_back(AxOrdersWsMessage::Authenticated);
227            }
228            HandlerCommand::PlaceOrder {
229                request_id,
230                order,
231                order_info,
232            } => {
233                log::debug!(
234                    "PlaceOrder command received: request_id={request_id}, symbol={}",
235                    order.s
236                );
237                self.pending_orders.insert(request_id, order_info.clone());
238
239                if let Err(e) = self.send_json(&order).await {
240                    log::error!("Failed to send place order message: {e}");
241                    self.pending_orders.remove(&request_id);
242                    self.orders_metadata.remove(&order_info.client_order_id);
243                    self.cid_to_client_order_id.remove(&order_info.cid);
244                    self.message_queue
245                        .push_back(AxOrdersWsMessage::Error(AxWsError::new(format!(
246                            "Failed to send place order for {}: {e}",
247                            order_info.client_order_id
248                        ))));
249                }
250            }
251            HandlerCommand::CancelOrder {
252                request_id,
253                order_id,
254            } => {
255                log::debug!(
256                    "CancelOrder command received: request_id={request_id}, order_id={order_id}"
257                );
258                self.send_cancel_order(request_id, &order_id).await;
259            }
260            HandlerCommand::GetOpenOrders { request_id } => {
261                log::debug!("GetOpenOrders command received: request_id={request_id}");
262                self.send_get_open_orders(request_id).await;
263            }
264        }
265    }
266
267    async fn send_cancel_order(&mut self, request_id: i64, order_id: &str) {
268        let msg = AxWsCancelOrder {
269            rid: request_id,
270            t: AxOrderRequestType::CancelOrder,
271            oid: order_id.to_string(),
272        };
273
274        if let Err(e) = self.send_json(&msg).await {
275            log::error!("Failed to send cancel order message: {e}");
276            self.message_queue
277                .push_back(AxOrdersWsMessage::Error(AxWsError::new(format!(
278                    "Failed to send cancel for order {order_id}: {e}"
279                ))));
280        }
281    }
282
283    async fn send_get_open_orders(&mut self, request_id: i64) {
284        let msg = AxWsGetOpenOrders {
285            rid: request_id,
286            t: AxOrderRequestType::GetOpenOrders,
287        };
288
289        if let Err(e) = self.send_json(&msg).await {
290            log::error!("Failed to send get open orders message: {e}");
291            self.message_queue
292                .push_back(AxOrdersWsMessage::Error(AxWsError::new(format!(
293                    "Failed to send get open orders request: {e}"
294                ))));
295        }
296    }
297
298    async fn send_json<T: serde::Serialize>(&self, msg: &T) -> Result<(), String> {
299        let Some(inner) = &self.inner else {
300            return Err("No WebSocket client available".to_string());
301        };
302
303        let payload = serde_json::to_string(msg).map_err(|e| e.to_string())?;
304        log::trace!("Sending WebSocket payload ({} bytes)", payload.len());
305
306        inner
307            .send_text(payload, None)
308            .await
309            .map_err(|e| e.to_string())
310    }
311
312    fn parse_raw_message(&mut self, msg: Message) -> Option<Vec<AxOrdersWsMessage>> {
313        match msg {
314            Message::Text(text) => {
315                if text == nautilus_network::RECONNECTED {
316                    log::info!("Received WebSocket reconnected signal");
317                    self.auth_tracker.fail("Reconnecting");
318                    self.needs_session_restore = true;
319                    return Some(vec![AxOrdersWsMessage::Reconnected]);
320                }
321
322                log::trace!("Raw websocket message: {text}");
323
324                let raw_msg: AxOrdersWsFrame = match parse_order_message(&text) {
325                    Ok(v) => v,
326                    Err(e) => {
327                        log::error!("Failed to parse WebSocket message: {e}: {text}");
328                        return None;
329                    }
330                };
331
332                self.handle_raw_message(raw_msg)
333            }
334            Message::Binary(data) => {
335                log::debug!("Received binary message with {} bytes", data.len());
336                None
337            }
338            Message::Close(_) => {
339                log::debug!("Received close message, waiting for reconnection");
340                None
341            }
342            _ => None,
343        }
344    }
345
346    fn handle_raw_message(&mut self, raw_msg: AxOrdersWsFrame) -> Option<Vec<AxOrdersWsMessage>> {
347        match raw_msg {
348            AxOrdersWsFrame::Error(err) => {
349                log::warn!(
350                    "Order error response: rid={} code={} msg={}",
351                    err.rid,
352                    err.err.code,
353                    err.err.msg
354                );
355
356                if let Some(order_info) = self.pending_orders.remove(&err.rid) {
357                    self.orders_metadata.remove(&order_info.client_order_id);
358                    log::debug!(
359                        "Cleaned up metadata for failed order: {}",
360                        order_info.client_order_id
361                    );
362                }
363
364                Some(vec![AxOrdersWsMessage::Error(err.into())])
365            }
366            AxOrdersWsFrame::Response(resp) => self.handle_response(resp),
367            AxOrdersWsFrame::Event(event) => self.handle_event(*event),
368        }
369    }
370
371    fn handle_response(&mut self, resp: AxWsOrderResponse) -> Option<Vec<AxOrdersWsMessage>> {
372        match resp {
373            AxWsOrderResponse::PlaceOrder(msg) => {
374                log::debug!("Place order response: rid={} oid={}", msg.rid, msg.res.oid);
375                let Some(order_info) = self.pending_orders.remove(&msg.rid) else {
376                    log::warn!("Ignoring unsolicited place order response: rid={}", msg.rid);
377                    return Some(vec![AxOrdersWsMessage::PlaceOrderResponse(msg)]);
378                };
379
380                let venue_order_id = match VenueOrderId::new_checked(&msg.res.oid) {
381                    Ok(venue_order_id) => venue_order_id,
382                    Err(e) => {
383                        log::warn!(
384                            "Invalid venue order ID in place response for {}: {e}",
385                            order_info.client_order_id,
386                        );
387                        return Some(vec![AxOrdersWsMessage::PlaceOrderResponse(msg)]);
388                    }
389                };
390
391                if let Some(mut metadata) =
392                    self.orders_metadata.get_mut(&order_info.client_order_id)
393                {
394                    metadata.venue_order_id = Some(venue_order_id);
395                    self.venue_to_client_order_id
396                        .insert(venue_order_id, order_info.client_order_id);
397                } else {
398                    log::debug!(
399                        "Order tracking already cleared before place response: {}",
400                        order_info.client_order_id,
401                    );
402                }
403
404                Some(vec![AxOrdersWsMessage::PlaceOrderResponse(msg)])
405            }
406            AxWsOrderResponse::CancelOrder(msg) => {
407                log::debug!(
408                    "Cancel order response: rid={} accepted={}",
409                    msg.rid,
410                    msg.res.cxl_rx
411                );
412                Some(vec![AxOrdersWsMessage::CancelOrderResponse(msg)])
413            }
414            AxWsOrderResponse::OpenOrders(msg) => {
415                log::debug!("Open orders response: {} orders", msg.res.orders.len());
416                Some(vec![AxOrdersWsMessage::OpenOrdersResponse(msg)])
417            }
418            AxWsOrderResponse::List(msg) => {
419                let order_count = msg.res.o.as_ref().map_or(0, |o| o.len());
420                log::debug!(
421                    "List subscription response: rid={} li={} orders={}",
422                    msg.rid,
423                    msg.res.li,
424                    order_count
425                );
426                None
427            }
428        }
429    }
430
431    fn handle_event(&self, event: AxWsOrderEvent) -> Option<Vec<AxOrdersWsMessage>> {
432        if matches!(event, AxWsOrderEvent::Heartbeat) {
433            log::trace!("Received heartbeat");
434            return None;
435        }
436        Some(vec![AxOrdersWsMessage::Event(Box::new(event))])
437    }
438}
439
440#[cfg(test)]
441mod tests {
442    use std::sync::{Arc, atomic::AtomicBool};
443
444    use dashmap::DashMap;
445    use nautilus_model::{
446        identifiers::{InstrumentId, StrategyId, TraderId},
447        types::Currency,
448    };
449    use nautilus_network::websocket::AuthTracker;
450    use rstest::rstest;
451    use ustr::Ustr;
452
453    use super::*;
454    use crate::websocket::messages::{
455        AxWsOrderError, AxWsOrderErrorResponse, AxWsPlaceOrderResponse, AxWsPlaceOrderResult,
456    };
457
458    fn test_handler() -> AxOrdersWsFeedHandler {
459        let (_cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel();
460        let (_raw_tx, raw_rx) = tokio::sync::mpsc::unbounded_channel();
461        AxOrdersWsFeedHandler::new(
462            Arc::new(AtomicBool::new(false)),
463            cmd_rx,
464            raw_rx,
465            AuthTracker::default(),
466            Arc::new(DashMap::new()),
467            Arc::new(DashMap::new()),
468            Arc::new(DashMap::new()),
469        )
470    }
471
472    #[rstest]
473    fn test_place_order_response_records_venue_identity() {
474        let mut handler = test_handler();
475        let request_id = 11;
476        let cid = 1011;
477        let client_order_id = ClientOrderId::from("CID-11");
478        let venue_order_id = VenueOrderId::from("OID-11");
479        handler
480            .orders_metadata
481            .insert(client_order_id, test_order_metadata(client_order_id));
482        handler.pending_orders.insert(
483            request_id,
484            WsOrderInfo {
485                client_order_id,
486                symbol: Ustr::from("EURUSD-PERP"),
487                cid,
488            },
489        );
490
491        let response = AxWsOrderResponse::PlaceOrder(AxWsPlaceOrderResponse {
492            rid: request_id,
493            res: AxWsPlaceOrderResult {
494                oid: "OID-11".to_string(),
495            },
496        });
497
498        let messages = handler.handle_response(response).unwrap();
499
500        assert_eq!(messages.len(), 1);
501        assert!(handler.pending_orders.get(&request_id).is_none());
502        assert_eq!(
503            handler
504                .orders_metadata
505                .get(&client_order_id)
506                .and_then(|metadata| metadata.venue_order_id),
507            Some(venue_order_id),
508        );
509        assert_eq!(
510            handler
511                .venue_to_client_order_id
512                .get(&venue_order_id)
513                .map(|client_order_id| *client_order_id),
514            Some(client_order_id),
515        );
516    }
517
518    #[rstest]
519    fn test_late_place_order_response_does_not_restore_cleared_tracking() {
520        let mut handler = test_handler();
521        let request_id = 12;
522        let client_order_id = ClientOrderId::from("CID-12");
523        handler.pending_orders.insert(
524            request_id,
525            WsOrderInfo {
526                client_order_id,
527                symbol: Ustr::from("EURUSD-PERP"),
528                cid: 1012,
529            },
530        );
531
532        let response = AxWsOrderResponse::PlaceOrder(AxWsPlaceOrderResponse {
533            rid: request_id,
534            res: AxWsPlaceOrderResult {
535                oid: "OID-12".to_string(),
536            },
537        });
538
539        let messages = handler.handle_response(response).unwrap();
540
541        assert_eq!(messages.len(), 1);
542        assert!(handler.pending_orders.get(&request_id).is_none());
543        assert!(!handler.orders_metadata.contains_key(&client_order_id));
544        assert!(handler.venue_to_client_order_id.is_empty());
545    }
546
547    #[rstest]
548    fn test_place_order_error_preserves_cid_for_reconciliation() {
549        let mut handler = test_handler();
550        let request_id = 13;
551        let cid = 1013;
552        let client_order_id = ClientOrderId::from("CID-13");
553        handler
554            .orders_metadata
555            .insert(client_order_id, test_order_metadata(client_order_id));
556        handler.cid_to_client_order_id.insert(cid, client_order_id);
557        handler.pending_orders.insert(
558            request_id,
559            WsOrderInfo {
560                client_order_id,
561                symbol: Ustr::from("EURUSD-PERP"),
562                cid,
563            },
564        );
565
566        let messages = handler
567            .handle_raw_message(AxOrdersWsFrame::Error(AxWsOrderErrorResponse {
568                rid: request_id,
569                err: AxWsOrderError {
570                    code: 400,
571                    msg: "invalid order".to_string(),
572                },
573            }))
574            .unwrap();
575
576        assert_eq!(messages.len(), 1);
577        assert!(handler.pending_orders.get(&request_id).is_none());
578        assert!(!handler.orders_metadata.contains_key(&client_order_id));
579        assert_eq!(
580            handler
581                .cid_to_client_order_id
582                .get(&cid)
583                .map(|client_order_id| *client_order_id),
584            Some(client_order_id),
585        );
586    }
587
588    #[rstest]
589    fn test_handle_event_forwards_venue_event() {
590        let handler = test_handler();
591
592        let event = AxWsOrderEvent::Heartbeat;
593        let result = handler.handle_event(event);
594        assert!(result.is_none());
595    }
596
597    #[tokio::test]
598    async fn test_authenticated_session_is_restored_after_reconnect() {
599        let mut handler = test_handler();
600
601        handler
602            .handle_command(HandlerCommand::SessionAuthenticated)
603            .await;
604        let initial = handler.message_queue.pop_front();
605        handler.auth_tracker.fail("Reconnecting");
606        handler.restore_authenticated_session();
607        let restored = handler.message_queue.pop_front();
608
609        assert!(handler.has_authenticated_session);
610        assert!(matches!(initial, Some(AxOrdersWsMessage::Authenticated)));
611        assert!(matches!(restored, Some(AxOrdersWsMessage::Authenticated)));
612        assert!(handler.auth_tracker.is_authenticated());
613    }
614
615    fn test_order_metadata(client_order_id: ClientOrderId) -> OrderMetadata {
616        OrderMetadata {
617            trader_id: TraderId::from("TRADER-001"),
618            strategy_id: StrategyId::from("S-001"),
619            instrument_id: InstrumentId::from("EURUSD-PERP.AX"),
620            client_order_id,
621            venue_order_id: None,
622            ts_init: 0.into(),
623            size_precision: 0,
624            price_precision: 2,
625            quote_currency: Currency::USD(),
626        }
627    }
628}