1use 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#[derive(Clone, Debug)]
46pub struct WsOrderInfo {
47 pub client_order_id: ClientOrderId,
49 pub symbol: Ustr,
51 pub cid: u64,
53}
54
55#[derive(Debug)]
57pub enum HandlerCommand {
58 SetClient(WebSocketClient),
60 Disconnect,
62 SessionAuthenticated,
64 PlaceOrder {
66 request_id: i64,
68 order: AxWsPlaceOrder,
70 order_info: WsOrderInfo,
72 },
73 CancelOrder {
75 request_id: i64,
77 order_id: String,
79 },
80 GetOpenOrders {
82 request_id: i64,
84 },
85}
86
87pub(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 #[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 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 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}