1use std::{
19 fmt::Debug,
20 num::NonZeroU32,
21 sync::{
22 Arc,
23 atomic::{AtomicBool, AtomicI64, AtomicU8, Ordering},
24 },
25 time::Duration,
26};
27
28use arc_swap::ArcSwap;
29use dashmap::{DashMap, mapref::entry::Entry};
30use nautilus_common::cache::InstrumentLookupError;
31use nautilus_core::{
32 AtomicMap,
33 consts::NAUTILUS_USER_AGENT,
34 nanos::UnixNanos,
35 time::{AtomicTime, get_atomic_clock_realtime},
36};
37use nautilus_live::{
38 SocketControl,
39 task::{SharedTaskSlot, TaskJoinOutcome},
40};
41use nautilus_model::{
42 enums::{OrderSide, TimeInForce},
43 identifiers::{AccountId, ClientOrderId, InstrumentId, StrategyId, TraderId, VenueOrderId},
44 instruments::{Instrument, InstrumentAny},
45 types::{Price, Quantity},
46};
47use nautilus_network::{
48 http::USER_AGENT,
49 mode::ConnectionMode,
50 websocket::{
51 AuthTracker, InitialConnectRetryPolicy, PingHandler, ReconnectHeaders, TransportBackend,
52 WebSocketClient, WebSocketConfig, channel_message_handler,
53 },
54};
55use parking_lot::Mutex;
56use tokio_util::sync::CancellationToken;
57use ustr::Ustr;
58
59use super::handler::{AxOrdersWsFeedHandler, HandlerCommand, WsOrderInfo};
60use crate::{
61 common::{
62 consts::AX_NAUTILUS_TAG,
63 enums::{AxOrderRequestType, AxOrderSide, AxTimeInForce},
64 parse::{client_order_id_to_cid, quantity_to_contracts},
65 },
66 websocket::messages::{AxOrdersWsMessage, AxWsPlaceOrder, OrderMetadata},
67};
68
69pub type AxOrdersWsResult<T> = Result<T, AxOrdersWsClientError>;
71
72#[derive(Debug, Clone)]
74pub struct OrdersCaches {
75 pub orders_metadata: Arc<DashMap<ClientOrderId, OrderMetadata>>,
77 pub venue_to_client_id: Arc<DashMap<VenueOrderId, ClientOrderId>>,
79 pub cid_to_client_order_id: Arc<DashMap<u64, ClientOrderId>>,
81}
82
83impl Default for OrdersCaches {
84 fn default() -> Self {
85 Self {
86 orders_metadata: Arc::new(DashMap::new()),
87 venue_to_client_id: Arc::new(DashMap::new()),
88 cid_to_client_order_id: Arc::new(DashMap::new()),
89 }
90 }
91}
92
93#[derive(Debug, Clone)]
95pub enum AxOrdersWsClientError {
96 Transport(String),
98 ChannelError(String),
100 AuthenticationError(String),
102 ClientError(String),
104}
105
106impl core::fmt::Display for AxOrdersWsClientError {
107 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
108 match self {
109 Self::Transport(msg) => write!(f, "Transport error: {msg}"),
110 Self::ChannelError(msg) => write!(f, "Channel error: {msg}"),
111 Self::AuthenticationError(msg) => write!(f, "Authentication error: {msg}"),
112 Self::ClientError(msg) => write!(f, "Client error: {msg}"),
113 }
114 }
115}
116
117impl std::error::Error for AxOrdersWsClientError {}
118
119impl From<&'static str> for AxOrdersWsClientError {
120 fn from(msg: &'static str) -> Self {
121 Self::ClientError(msg.to_string())
122 }
123}
124
125pub struct AxOrdersWebSocketClient {
130 clock: &'static AtomicTime,
131 url: String,
132 heartbeat: Option<u64>,
133 reconnect_headers: Arc<Mutex<Option<ReconnectHeaders>>>,
134 connection_mode: Arc<ArcSwap<AtomicU8>>,
135 cmd_tx: Arc<tokio::sync::RwLock<tokio::sync::mpsc::UnboundedSender<HandlerCommand>>>,
136 out_rx: Option<Arc<tokio::sync::mpsc::UnboundedReceiver<AxOrdersWsMessage>>>,
137 signal: Arc<AtomicBool>,
138 cancellation_token: Arc<ArcSwap<CancellationToken>>,
139 task_handle: Arc<SharedTaskSlot<()>>,
140 connect_lock: Arc<tokio::sync::Mutex<()>>,
141 auth_tracker: AuthTracker,
142 instruments_cache: Arc<AtomicMap<Ustr, InstrumentAny>>,
143 caches: OrdersCaches,
144 request_id_counter: Arc<AtomicI64>,
145 account_id: AccountId,
146 trader_id: TraderId,
147 transport_backend: TransportBackend,
148 proxy_url: Option<String>,
149 socket_control: Option<SocketControl>,
150}
151
152impl Debug for AxOrdersWebSocketClient {
153 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
154 f.debug_struct(stringify!(AxOrdersWebSocketClient))
155 .field("url", &self.url)
156 .field("heartbeat", &self.heartbeat)
157 .field("account_id", &self.account_id)
158 .finish()
159 }
160}
161
162impl Clone for AxOrdersWebSocketClient {
163 fn clone(&self) -> Self {
164 Self {
165 clock: self.clock,
166 url: self.url.clone(),
167 heartbeat: self.heartbeat,
168 reconnect_headers: Arc::clone(&self.reconnect_headers),
169 connection_mode: Arc::clone(&self.connection_mode),
170 cmd_tx: Arc::clone(&self.cmd_tx),
171 out_rx: None, signal: Arc::clone(&self.signal),
173 cancellation_token: Arc::clone(&self.cancellation_token),
174 task_handle: Arc::clone(&self.task_handle),
175 connect_lock: Arc::clone(&self.connect_lock),
176 auth_tracker: self.auth_tracker.clone(),
177 instruments_cache: Arc::clone(&self.instruments_cache),
178 caches: self.caches.clone(),
179 request_id_counter: Arc::clone(&self.request_id_counter),
180 account_id: self.account_id,
181 trader_id: self.trader_id,
182 transport_backend: self.transport_backend,
183 proxy_url: self.proxy_url.clone(),
184 socket_control: self.socket_control.clone(),
185 }
186 }
187}
188
189impl AxOrdersWebSocketClient {
190 fn initial_connect_retry_policy() -> InitialConnectRetryPolicy {
191 InitialConnectRetryPolicy {
192 max_attempts: NonZeroU32::new(5).expect("initial connect attempts must be non-zero"),
193 delay_initial: Duration::from_millis(500),
194 delay_max: Duration::from_secs(5),
195 backoff_factor: 2.0,
196 jitter_ms: 250,
197 }
198 }
199
200 #[must_use]
202 pub fn new(
203 url: String,
204 account_id: AccountId,
205 trader_id: TraderId,
206 heartbeat: u64,
207 transport_backend: TransportBackend,
208 proxy_url: Option<String>,
209 ) -> Self {
210 let (cmd_tx, _cmd_rx) = tokio::sync::mpsc::unbounded_channel::<HandlerCommand>();
211
212 let initial_mode = AtomicU8::new(ConnectionMode::Closed.as_u8());
213 let connection_mode = Arc::new(ArcSwap::from_pointee(initial_mode));
214
215 Self {
216 clock: get_atomic_clock_realtime(),
217 url,
218 heartbeat: Some(heartbeat),
219 reconnect_headers: Arc::new(Mutex::new(None)),
220 connection_mode,
221 cmd_tx: Arc::new(tokio::sync::RwLock::new(cmd_tx)),
222 out_rx: None,
223 signal: Arc::new(AtomicBool::new(false)),
224 cancellation_token: Arc::new(ArcSwap::from_pointee(CancellationToken::new())),
225 task_handle: Arc::new(SharedTaskSlot::new()),
226 connect_lock: Arc::new(tokio::sync::Mutex::new(())),
227 auth_tracker: AuthTracker::default(),
228 instruments_cache: Arc::new(AtomicMap::new()),
229 caches: OrdersCaches::default(),
230 request_id_counter: Arc::new(AtomicI64::new(1)),
231 account_id,
232 trader_id,
233 transport_backend,
234 proxy_url,
235 socket_control: None,
236 }
237 }
238
239 #[must_use]
241 pub fn with_socket_control(mut self, control: SocketControl) -> Self {
242 self.socket_control = Some(control);
243 self
244 }
245
246 fn generate_ts_init(&self) -> UnixNanos {
247 self.clock.get_time_ns()
248 }
249
250 #[must_use]
252 pub fn url(&self) -> &str {
253 &self.url
254 }
255
256 #[must_use]
258 pub fn account_id(&self) -> AccountId {
259 self.account_id
260 }
261
262 #[must_use]
264 pub fn is_active(&self) -> bool {
265 let connection_mode_arc = self.connection_mode.load();
266 ConnectionMode::from_atomic(&connection_mode_arc).is_active()
267 && !self.signal.load(Ordering::Acquire)
268 }
269
270 #[must_use]
272 pub fn is_closed(&self) -> bool {
273 let connection_mode_arc = self.connection_mode.load();
274 ConnectionMode::from_atomic(&connection_mode_arc).is_closed()
275 || self.signal.load(Ordering::Acquire)
276 }
277
278 fn next_request_id(&self) -> i64 {
280 self.request_id_counter.fetch_add(1, Ordering::Relaxed)
281 }
282
283 pub fn cache_instrument(&self, instrument: InstrumentAny) {
285 let symbol = instrument.symbol().inner();
286 self.instruments_cache.insert(symbol, instrument);
287 }
288
289 pub fn cache_instruments(&self, instruments: &[InstrumentAny]) {
291 self.instruments_cache.rcu(|m| {
292 for inst in instruments {
293 m.insert(inst.symbol().inner(), inst.clone());
294 }
295 });
296 }
297
298 pub fn update_auth_token(&self, token: &str) -> AxOrdersWsResult<()> {
306 let value = format!("Bearer {token}");
307
308 if let Some(headers) = self.reconnect_headers.lock().as_ref() {
309 headers
310 .update("Authorization", &value)
311 .map_err(|e| AxOrdersWsClientError::Transport(e.to_string()))?;
312 }
313 Ok(())
314 }
315
316 #[must_use]
318 pub fn get_cached_instrument(&self, symbol: &Ustr) -> Option<InstrumentAny> {
319 self.instruments_cache.get_cloned(symbol)
320 }
321
322 #[must_use]
324 pub fn caches(&self) -> &OrdersCaches {
325 &self.caches
326 }
327
328 #[must_use]
330 pub fn instruments_cache(&self) -> Arc<AtomicMap<Ustr, InstrumentAny>> {
331 Arc::clone(&self.instruments_cache)
332 }
333
334 #[must_use]
336 pub fn orders_metadata(&self) -> &Arc<DashMap<ClientOrderId, OrderMetadata>> {
337 &self.caches.orders_metadata
338 }
339
340 #[must_use]
342 pub fn cid_to_client_order_id(&self) -> &Arc<DashMap<u64, ClientOrderId>> {
343 &self.caches.cid_to_client_order_id
344 }
345
346 #[must_use]
348 pub fn resolve_cid(&self, cid: u64) -> Option<ClientOrderId> {
349 self.caches.cid_to_client_order_id.get(&cid).map(|v| *v)
350 }
351
352 pub fn register_external_order(
359 &self,
360 client_order_id: ClientOrderId,
361 venue_order_id: VenueOrderId,
362 instrument_id: InstrumentId,
363 strategy_id: StrategyId,
364 ) -> bool {
365 if self.caches.orders_metadata.contains_key(&client_order_id) {
366 return true;
367 }
368
369 let symbol = instrument_id.symbol.inner();
371 let Some(instrument) = self.get_cached_instrument(&symbol) else {
372 log::warn!(
373 "Cannot register external order {client_order_id}: \
374 instrument {instrument_id} not in cache"
375 );
376 return false;
377 };
378
379 let metadata = OrderMetadata {
380 trader_id: self.trader_id,
381 strategy_id,
382 instrument_id,
383 client_order_id,
384 venue_order_id: Some(venue_order_id),
385 ts_init: self.generate_ts_init(),
386 size_precision: instrument.size_precision(),
387 price_precision: instrument.price_precision(),
388 quote_currency: instrument.quote_currency(),
389 };
390
391 self.caches
392 .orders_metadata
393 .insert(client_order_id, metadata);
394 self.caches
395 .venue_to_client_id
396 .insert(venue_order_id, client_order_id);
397
398 log::debug!(
399 "Registered external order {client_order_id} ({venue_order_id}) for {instrument_id} [{strategy_id}]"
400 );
401
402 true
403 }
404
405 pub async fn connect(&mut self, bearer_token: &str) -> AxOrdersWsResult<()> {
415 let connect_lock = Arc::clone(&self.connect_lock);
416 let _guard = connect_lock.lock().await;
417
418 if !self.task_handle.is_empty() && !self.task_handle.is_finished() {
419 return Err(AxOrdersWsClientError::ClientError(
420 "WebSocket handler is already running".to_string(),
421 ));
422 }
423
424 if let Some(outcome) = self
425 .task_handle
426 .finish(Duration::from_secs(2), Duration::from_secs(2))
427 .await
428 {
429 match outcome {
430 TaskJoinOutcome::Completed(()) | TaskJoinOutcome::Aborted => {}
431 TaskJoinOutcome::Failed(error) => {
432 return Err(AxOrdersWsClientError::ClientError(format!(
433 "Previous WebSocket handler failed: {error}"
434 )));
435 }
436 TaskJoinOutcome::Incomplete => {
437 return Err(AxOrdersWsClientError::ClientError(
438 "Previous WebSocket handler did not stop within shutdown bounds"
439 .to_string(),
440 ));
441 }
442 }
443 }
444
445 self.signal.store(false, Ordering::Release);
446 let cancellation_token = CancellationToken::new();
447 self.cancellation_token
448 .store(Arc::new(cancellation_token.clone()));
449
450 let (raw_handler, raw_rx) = channel_message_handler();
451
452 let ping_handler: PingHandler = Arc::new(move |_payload: Vec<u8>| {
454 });
456
457 let config = WebSocketConfig {
458 url: self.url.clone(),
459 headers: vec![
460 (USER_AGENT.to_string(), NAUTILUS_USER_AGENT.to_string()),
461 (
462 "Authorization".to_string(),
463 format!("Bearer {bearer_token}"),
464 ),
465 ],
466 heartbeat_interval_secs: self.heartbeat,
467 heartbeat_payload: None, connect_timeout_ms: Some(5_000),
469 reconnect_delay_initial_ms: Some(500),
470 reconnect_delay_max_ms: Some(5_000),
471 reconnect_backoff_factor: Some(1.5),
472 reconnect_jitter_ms: Some(250),
473 reconnect_max_attempts: None,
474 heartbeat_timeout_secs: None,
475 idle_timeout_ms: None,
476 backend: self.transport_backend,
477 proxy_url: self.proxy_url.clone(),
478 };
479
480 let client = WebSocketClient::builder()
481 .config(config.clone())
482 .message_handler(raw_handler.clone())
483 .ping_handler(ping_handler.clone())
484 .initial_connect_retry_policy(Self::initial_connect_retry_policy())
485 .cancellation_token(cancellation_token)
486 .maybe_state_sink(self.socket_control.as_ref().map(SocketControl::sink))
487 .connect()
488 .await
489 .map_err(|e| {
490 AxOrdersWsClientError::Transport(format!("Failed to connect to {}: {e}", self.url))
491 })?;
492
493 self.connection_mode.store(client.connection_mode_atomic());
494 let reconnect_handle = client.reconnect_handle();
495 *self.reconnect_headers.lock() = Some(client.reconnect_headers());
496
497 let (out_tx, out_rx) = tokio::sync::mpsc::unbounded_channel::<AxOrdersWsMessage>();
498 self.out_rx = Some(Arc::new(out_rx));
499
500 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel::<HandlerCommand>();
501 *self.cmd_tx.write().await = cmd_tx.clone();
502
503 self.send_cmd(HandlerCommand::SetClient(client)).await?;
504
505 self.send_cmd(HandlerCommand::SessionAuthenticated).await?;
506
507 let signal = Arc::clone(&self.signal);
508 let auth_tracker = self.auth_tracker.clone();
509 let orders_metadata = Arc::clone(&self.caches.orders_metadata);
510 let venue_to_client_order_id = Arc::clone(&self.caches.venue_to_client_id);
511 let cid_to_client_order_id = Arc::clone(&self.caches.cid_to_client_order_id);
512
513 if let Err(e) = self.task_handle.spawn(async move {
514 let mut handler = AxOrdersWsFeedHandler::new(
515 signal.clone(),
516 cmd_rx,
517 raw_rx,
518 auth_tracker.clone(),
519 orders_metadata,
520 venue_to_client_order_id,
521 cid_to_client_order_id,
522 );
523
524 while let Some(msg) = handler.next().await {
525 if matches!(msg, AxOrdersWsMessage::Reconnected) {
526 log::info!("WebSocket reconnected, authentication will be restored");
527 }
528
529 if out_tx.send(msg).is_err() {
530 log::debug!("Output channel closed");
531 break;
532 }
533 }
534
535 log::debug!("Handler loop exited");
536 }) {
537 self.out_rx = None;
538 return Err(AxOrdersWsClientError::Transport(format!(
539 "Failed to start WebSocket handler task: {e}"
540 )));
541 }
542
543 if let Some(control) = &self.socket_control {
544 control.register(move || reconnect_handle.request_reconnect());
545 }
546
547 Ok(())
548 }
549
550 #[expect(clippy::too_many_arguments)]
562 pub async fn submit_order(
563 &self,
564 trader_id: TraderId,
565 strategy_id: StrategyId,
566 instrument_id: InstrumentId,
567 client_order_id: ClientOrderId,
568 order_side: OrderSide,
569 quantity: Quantity,
570 time_in_force: TimeInForce,
571 price: Price,
572 post_only: bool,
573 ) -> AxOrdersWsResult<i64> {
574 let symbol = instrument_id.symbol.inner();
576 let instrument = self.get_cached_instrument(&symbol).ok_or_else(|| {
577 AxOrdersWsClientError::ClientError(
578 InstrumentLookupError::not_found(instrument_id).to_string(),
579 )
580 })?;
581
582 let ax_side = AxOrderSide::from(order_side);
583
584 let qty_contracts = quantity_to_contracts(quantity)
585 .map_err(|e| AxOrdersWsClientError::ClientError(e.to_string()))?;
586
587 let request_id = self.next_request_id();
588 let ax_tif = AxTimeInForce::try_from(time_in_force)?;
589 let cid = client_order_id_to_cid(&client_order_id);
590
591 reserve_cid_mapping(&self.caches, cid, client_order_id)?;
592
593 let metadata = OrderMetadata {
595 trader_id,
596 strategy_id,
597 instrument_id,
598 client_order_id,
599 venue_order_id: None,
600 ts_init: self.generate_ts_init(),
601 size_precision: instrument.size_precision(),
602 price_precision: instrument.price_precision(),
603 quote_currency: instrument.quote_currency(),
604 };
605 self.caches
606 .orders_metadata
607 .insert(client_order_id, metadata);
608
609 let order = AxWsPlaceOrder {
610 rid: request_id,
611 t: AxOrderRequestType::PlaceOrder,
612 s: symbol,
613 d: ax_side,
614 q: qty_contracts,
615 p: price.as_decimal(),
616 tif: ax_tif,
617 po: post_only,
618 tag: Some(AX_NAUTILUS_TAG.to_string()),
619 cid: Some(cid),
620 };
621
622 let order_info = WsOrderInfo {
623 client_order_id,
624 symbol,
625 cid,
626 };
627
628 let result = self
629 .send_cmd(HandlerCommand::PlaceOrder {
630 request_id,
631 order,
632 order_info,
633 })
634 .await;
635
636 if result.is_err() {
637 self.caches.orders_metadata.remove(&client_order_id);
638 self.caches.cid_to_client_order_id.remove(&cid);
639 }
640
641 result?;
642 Ok(request_id)
643 }
644
645 pub async fn cancel_order(
653 &self,
654 client_order_id: ClientOrderId,
655 venue_order_id: Option<VenueOrderId>,
656 ) -> AxOrdersWsResult<i64> {
657 let order_id = venue_order_id.map(|v| v.to_string()).ok_or_else(|| {
658 AxOrdersWsClientError::ClientError(format!(
659 "Cannot cancel order {client_order_id}: missing venue_order_id"
660 ))
661 })?;
662
663 let request_id = self.next_request_id();
664
665 self.send_cmd(HandlerCommand::CancelOrder {
666 request_id,
667 order_id,
668 })
669 .await?;
670
671 Ok(request_id)
672 }
673
674 pub async fn get_open_orders(&self) -> AxOrdersWsResult<i64> {
680 let request_id = self.next_request_id();
681
682 self.send_cmd(HandlerCommand::GetOpenOrders { request_id })
683 .await?;
684
685 Ok(request_id)
686 }
687
688 pub fn stream(&mut self) -> impl futures_util::Stream<Item = AxOrdersWsMessage> + 'static {
694 let rx = self
695 .out_rx
696 .take()
697 .expect("Stream receiver already taken or client not connected - stream() can only be called once");
698 let mut rx = Arc::try_unwrap(rx).expect(
699 "Cannot take ownership of stream - client was cloned and other references exist",
700 );
701 async_stream::stream! {
702 while let Some(msg) = rx.recv().await {
703 yield msg;
704 }
705 }
706 }
707
708 pub(crate) fn begin_shutdown(&self) {
709 self.cancellation_token.load().cancel();
710 self.signal.store(true, Ordering::Release);
711 }
712
713 pub async fn disconnect(&self) {
715 log::debug!("Disconnecting WebSocket");
716 let _ = self.send_cmd(HandlerCommand::Disconnect).await;
717 }
718
719 pub async fn close(&mut self) -> anyhow::Result<()> {
725 let connect_lock = Arc::clone(&self.connect_lock);
726 let _guard = connect_lock.lock().await;
727 log::debug!("Closing WebSocket client");
728
729 self.cancellation_token.load().cancel();
731 let _ = self.send_cmd(HandlerCommand::Disconnect).await;
732 tokio::time::sleep(Duration::from_millis(50)).await;
733 self.signal.store(true, Ordering::Release);
734
735 let outcome = self
736 .task_handle
737 .finish(Duration::from_secs(2), Duration::from_secs(2))
738 .await;
739 *self.reconnect_headers.lock() = None;
740
741 if let Some(control) = &self.socket_control {
742 control.deregister();
743 }
744
745 match outcome {
746 None | Some(TaskJoinOutcome::Completed(()) | TaskJoinOutcome::Aborted) => Ok(()),
747 Some(TaskJoinOutcome::Failed(error)) => Err(anyhow::anyhow!(
748 "Architect AX orders WebSocket handler failed: {error}"
749 )),
750 Some(TaskJoinOutcome::Incomplete) => Err(anyhow::anyhow!(
751 "Architect AX orders WebSocket handler did not stop after abort"
752 )),
753 }
754 }
755
756 async fn send_cmd(&self, cmd: HandlerCommand) -> AxOrdersWsResult<()> {
757 let guard = self.cmd_tx.read().await;
758 guard
759 .send(cmd)
760 .map_err(|e| AxOrdersWsClientError::ChannelError(e.to_string()))
761 }
762}
763
764impl Drop for AxOrdersWebSocketClient {
765 fn drop(&mut self) {
766 if Arc::strong_count(&self.task_handle) == 1 && !self.task_handle.is_empty() {
767 self.cancellation_token.load().cancel();
768 self.signal.store(true, Ordering::Release);
769 self.task_handle.abort();
770
771 if let Some(control) = &self.socket_control {
772 control.deregister();
773 }
774 }
775 }
776}
777
778fn reserve_cid_mapping(
779 caches: &OrdersCaches,
780 cid: u64,
781 client_order_id: ClientOrderId,
782) -> AxOrdersWsResult<()> {
783 match caches.cid_to_client_order_id.entry(cid) {
784 Entry::Vacant(entry) => {
785 entry.insert(client_order_id);
786 Ok(())
787 }
788 Entry::Occupied(entry) => Err(AxOrdersWsClientError::ClientError(format!(
789 "AX cid {cid} is already mapped to {}",
790 entry.get(),
791 ))),
792 }
793}
794
795#[cfg(test)]
796mod tests {
797 use std::sync::Arc;
798
799 use rstest::rstest;
800
801 use super::*;
802
803 #[tokio::test]
804 async fn test_drop_aborts_handler_task() {
805 let client = AxOrdersWebSocketClient::new(
806 "wss://example.com/orders/ws".to_string(),
807 AccountId::from("AX-001"),
808 TraderId::from("TRADER-001"),
809 30,
810 TransportBackend::default(),
811 None,
812 );
813 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
814 let handle = tokio::spawn(async move {
815 started_tx.send(()).expect("started receiver");
816 std::future::pending::<()>().await;
817 });
818 let abort_handle = handle.abort_handle();
819 client.task_handle.insert(handle);
820 started_rx.await.expect("handler task started");
821
822 drop(client);
823
824 tokio::time::timeout(Duration::from_secs(1), async {
825 while !abort_handle.is_finished() {
826 tokio::task::yield_now().await;
827 }
828 })
829 .await
830 .expect("handler task aborted");
831 }
832
833 #[rstest]
834 fn test_reserve_cid_mapping_rejects_collision() {
835 let caches = OrdersCaches::default();
836 let cid = 123;
837 let first_client_order_id = ClientOrderId::from("CID-123-A");
838 let second_client_order_id = ClientOrderId::from("CID-123-B");
839
840 reserve_cid_mapping(&caches, cid, first_client_order_id).unwrap();
841 let result = reserve_cid_mapping(&caches, cid, second_client_order_id);
842
843 assert!(matches!(
844 result,
845 Err(AxOrdersWsClientError::ClientError(msg))
846 if msg == "AX cid 123 is already mapped to CID-123-A"
847 ));
848 assert_eq!(
849 caches
850 .cid_to_client_order_id
851 .get(&cid)
852 .map(|client_order_id| *client_order_id),
853 Some(first_client_order_id),
854 );
855 }
856
857 #[tokio::test]
858 async fn test_cancel_order_rejects_without_venue_order_id() {
859 let client = AxOrdersWebSocketClient::new(
860 "wss://example.com/orders/ws".to_string(),
861 AccountId::from("AX-001"),
862 TraderId::from("TRADER-001"),
863 30,
864 TransportBackend::default(),
865 None,
866 );
867 let client_order_id = ClientOrderId::from("CID-123");
868
869 let result = client.cancel_order(client_order_id, None).await;
870
871 assert!(matches!(
872 result,
873 Err(AxOrdersWsClientError::ClientError(msg))
874 if msg.contains("missing venue_order_id")
875 ));
876 }
877
878 #[tokio::test]
879 async fn test_cancel_order_sends_known_venue_order_id() {
880 let mut client = AxOrdersWebSocketClient::new(
881 "wss://example.com/orders/ws".to_string(),
882 AccountId::from("AX-001"),
883 TraderId::from("TRADER-001"),
884 30,
885 TransportBackend::default(),
886 None,
887 );
888
889 let (cmd_tx, mut cmd_rx) = tokio::sync::mpsc::unbounded_channel::<HandlerCommand>();
890 client.cmd_tx = Arc::new(tokio::sync::RwLock::new(cmd_tx));
891
892 let client_order_id = ClientOrderId::from("CID-456");
893 let venue_order_id = VenueOrderId::from("V-ORDER-789");
894
895 let request_id = client
896 .cancel_order(client_order_id, Some(venue_order_id))
897 .await
898 .unwrap();
899
900 assert_eq!(request_id, 1);
901 let cmd = cmd_rx.recv().await.unwrap();
902 match cmd {
903 HandlerCommand::CancelOrder {
904 request_id,
905 order_id,
906 } => {
907 assert_eq!(request_id, 1);
908 assert_eq!(order_id, "V-ORDER-789");
909 }
910 other => panic!("unexpected command: {other:?}"),
911 }
912 }
913}