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 nanos::UnixNanos,
34 string::secret::SecretString,
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::create_standard_nautilus_headers,
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<SecretString>,
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: proxy_url.map(SecretString::from),
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 mut headers = create_standard_nautilus_headers();
458 headers.push((
459 "Authorization".to_string(),
460 format!("Bearer {bearer_token}"),
461 ));
462
463 let config = WebSocketConfig {
464 url: self.url.clone(),
465 headers,
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
478 .proxy_url
479 .as_ref()
480 .map(|url| url.expose_secret().to_owned()),
481 };
482
483 let client = WebSocketClient::builder()
484 .config(config.clone())
485 .message_handler(raw_handler.clone())
486 .ping_handler(ping_handler.clone())
487 .initial_connect_retry_policy(Self::initial_connect_retry_policy())
488 .cancellation_token(cancellation_token)
489 .maybe_state_sink(self.socket_control.as_ref().map(SocketControl::sink))
490 .connect()
491 .await
492 .map_err(|e| {
493 AxOrdersWsClientError::Transport(format!("Failed to connect to {}: {e}", self.url))
494 })?;
495
496 self.connection_mode.store(client.connection_mode_atomic());
497 let reconnect_handle = client.reconnect_handle();
498 *self.reconnect_headers.lock() = Some(client.reconnect_headers());
499
500 let (out_tx, out_rx) = tokio::sync::mpsc::unbounded_channel::<AxOrdersWsMessage>();
501 self.out_rx = Some(Arc::new(out_rx));
502
503 let (cmd_tx, cmd_rx) = tokio::sync::mpsc::unbounded_channel::<HandlerCommand>();
504 *self.cmd_tx.write().await = cmd_tx.clone();
505
506 self.send_cmd(HandlerCommand::SetClient(client)).await?;
507
508 self.send_cmd(HandlerCommand::SessionAuthenticated).await?;
509
510 let signal = Arc::clone(&self.signal);
511 let auth_tracker = self.auth_tracker.clone();
512 let orders_metadata = Arc::clone(&self.caches.orders_metadata);
513 let venue_to_client_order_id = Arc::clone(&self.caches.venue_to_client_id);
514 let cid_to_client_order_id = Arc::clone(&self.caches.cid_to_client_order_id);
515
516 if let Err(e) = self.task_handle.spawn(async move {
517 let mut handler = AxOrdersWsFeedHandler::new(
518 signal.clone(),
519 cmd_rx,
520 raw_rx,
521 auth_tracker.clone(),
522 orders_metadata,
523 venue_to_client_order_id,
524 cid_to_client_order_id,
525 );
526
527 while let Some(msg) = handler.next().await {
528 if matches!(msg, AxOrdersWsMessage::Reconnected) {
529 log::info!("WebSocket reconnected, authentication will be restored");
530 }
531
532 if out_tx.send(msg).is_err() {
533 log::debug!("Output channel closed");
534 break;
535 }
536 }
537
538 log::debug!("Handler loop exited");
539 }) {
540 self.out_rx = None;
541 return Err(AxOrdersWsClientError::Transport(format!(
542 "Failed to start WebSocket handler task: {e}"
543 )));
544 }
545
546 if let Some(control) = &self.socket_control {
547 control.register(move || reconnect_handle.request_reconnect());
548 }
549
550 Ok(())
551 }
552
553 #[expect(clippy::too_many_arguments)]
565 pub async fn submit_order(
566 &self,
567 trader_id: TraderId,
568 strategy_id: StrategyId,
569 instrument_id: InstrumentId,
570 client_order_id: ClientOrderId,
571 order_side: OrderSide,
572 quantity: Quantity,
573 time_in_force: TimeInForce,
574 price: Price,
575 post_only: bool,
576 ) -> AxOrdersWsResult<i64> {
577 let symbol = instrument_id.symbol.inner();
579 let instrument = self.get_cached_instrument(&symbol).ok_or_else(|| {
580 AxOrdersWsClientError::ClientError(
581 InstrumentLookupError::not_found(instrument_id).to_string(),
582 )
583 })?;
584
585 let ax_side = AxOrderSide::from(order_side);
586
587 let qty_contracts = quantity_to_contracts(quantity)
588 .map_err(|e| AxOrdersWsClientError::ClientError(e.to_string()))?;
589
590 let request_id = self.next_request_id();
591 let ax_tif = AxTimeInForce::try_from(time_in_force)?;
592 let cid = client_order_id_to_cid(&client_order_id);
593
594 reserve_cid_mapping(&self.caches, cid, client_order_id)?;
595
596 let metadata = OrderMetadata {
598 trader_id,
599 strategy_id,
600 instrument_id,
601 client_order_id,
602 venue_order_id: None,
603 ts_init: self.generate_ts_init(),
604 size_precision: instrument.size_precision(),
605 price_precision: instrument.price_precision(),
606 quote_currency: instrument.quote_currency(),
607 };
608 self.caches
609 .orders_metadata
610 .insert(client_order_id, metadata);
611
612 let order = AxWsPlaceOrder {
613 rid: request_id,
614 t: AxOrderRequestType::PlaceOrder,
615 s: symbol,
616 d: ax_side,
617 q: qty_contracts,
618 p: price.as_decimal(),
619 tif: ax_tif,
620 po: post_only,
621 tag: Some(AX_NAUTILUS_TAG.to_string()),
622 cid: Some(cid),
623 };
624
625 let order_info = WsOrderInfo {
626 client_order_id,
627 symbol,
628 cid,
629 };
630
631 let result = self
632 .send_cmd(HandlerCommand::PlaceOrder {
633 request_id,
634 order,
635 order_info,
636 })
637 .await;
638
639 if result.is_err() {
640 self.caches.orders_metadata.remove(&client_order_id);
641 self.caches.cid_to_client_order_id.remove(&cid);
642 }
643
644 result?;
645 Ok(request_id)
646 }
647
648 pub async fn cancel_order(
656 &self,
657 client_order_id: ClientOrderId,
658 venue_order_id: Option<VenueOrderId>,
659 ) -> AxOrdersWsResult<i64> {
660 let order_id = venue_order_id.map(|v| v.to_string()).ok_or_else(|| {
661 AxOrdersWsClientError::ClientError(format!(
662 "Cannot cancel order {client_order_id}: missing venue_order_id"
663 ))
664 })?;
665
666 let request_id = self.next_request_id();
667
668 self.send_cmd(HandlerCommand::CancelOrder {
669 request_id,
670 order_id,
671 })
672 .await?;
673
674 Ok(request_id)
675 }
676
677 pub async fn get_open_orders(&self) -> AxOrdersWsResult<i64> {
683 let request_id = self.next_request_id();
684
685 self.send_cmd(HandlerCommand::GetOpenOrders { request_id })
686 .await?;
687
688 Ok(request_id)
689 }
690
691 pub fn stream(&mut self) -> impl futures_util::Stream<Item = AxOrdersWsMessage> + 'static {
697 let rx = self
698 .out_rx
699 .take()
700 .expect("Stream receiver already taken or client not connected - stream() can only be called once");
701 let mut rx = Arc::try_unwrap(rx).expect(
702 "Cannot take ownership of stream - client was cloned and other references exist",
703 );
704 async_stream::stream! {
705 while let Some(msg) = rx.recv().await {
706 yield msg;
707 }
708 }
709 }
710
711 pub(crate) fn begin_shutdown(&self) {
712 self.cancellation_token.load().cancel();
713 self.signal.store(true, Ordering::Release);
714 }
715
716 pub async fn disconnect(&self) {
718 log::debug!("Disconnecting WebSocket");
719 let _ = self.send_cmd(HandlerCommand::Disconnect).await;
720 }
721
722 pub async fn close(&mut self) -> anyhow::Result<()> {
728 let connect_lock = Arc::clone(&self.connect_lock);
729 let _guard = connect_lock.lock().await;
730 log::debug!("Closing WebSocket client");
731
732 self.cancellation_token.load().cancel();
734 let _ = self.send_cmd(HandlerCommand::Disconnect).await;
735 tokio::time::sleep(Duration::from_millis(50)).await;
736 self.signal.store(true, Ordering::Release);
737
738 let outcome = self
739 .task_handle
740 .finish(Duration::from_secs(2), Duration::from_secs(2))
741 .await;
742 *self.reconnect_headers.lock() = None;
743
744 if let Some(control) = &self.socket_control {
745 control.deregister();
746 }
747
748 match outcome {
749 None | Some(TaskJoinOutcome::Completed(()) | TaskJoinOutcome::Aborted) => Ok(()),
750 Some(TaskJoinOutcome::Failed(error)) => Err(anyhow::anyhow!(
751 "Architect AX orders WebSocket handler failed: {error}"
752 )),
753 Some(TaskJoinOutcome::Incomplete) => Err(anyhow::anyhow!(
754 "Architect AX orders WebSocket handler did not stop after abort"
755 )),
756 }
757 }
758
759 async fn send_cmd(&self, cmd: HandlerCommand) -> AxOrdersWsResult<()> {
760 let guard = self.cmd_tx.read().await;
761 guard
762 .send(cmd)
763 .map_err(|e| AxOrdersWsClientError::ChannelError(e.to_string()))
764 }
765}
766
767impl Drop for AxOrdersWebSocketClient {
768 fn drop(&mut self) {
769 if Arc::strong_count(&self.task_handle) == 1 && !self.task_handle.is_empty() {
770 self.cancellation_token.load().cancel();
771 self.signal.store(true, Ordering::Release);
772 self.task_handle.abort();
773
774 if let Some(control) = &self.socket_control {
775 control.deregister();
776 }
777 }
778 }
779}
780
781fn reserve_cid_mapping(
782 caches: &OrdersCaches,
783 cid: u64,
784 client_order_id: ClientOrderId,
785) -> AxOrdersWsResult<()> {
786 match caches.cid_to_client_order_id.entry(cid) {
787 Entry::Vacant(entry) => {
788 entry.insert(client_order_id);
789 Ok(())
790 }
791 Entry::Occupied(entry) => Err(AxOrdersWsClientError::ClientError(format!(
792 "AX cid {cid} is already mapped to {}",
793 entry.get(),
794 ))),
795 }
796}
797
798#[cfg(test)]
799mod tests {
800 use std::sync::Arc;
801
802 use rstest::rstest;
803
804 use super::*;
805
806 #[tokio::test]
807 async fn test_drop_aborts_handler_task() {
808 let client = AxOrdersWebSocketClient::new(
809 "wss://example.com/orders/ws".to_string(),
810 AccountId::from("AX-001"),
811 TraderId::from("TRADER-001"),
812 30,
813 TransportBackend::default(),
814 None,
815 );
816 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
817 let handle = tokio::spawn(async move {
818 started_tx.send(()).expect("started receiver");
819 std::future::pending::<()>().await;
820 });
821 let abort_handle = handle.abort_handle();
822 client.task_handle.insert(handle);
823 started_rx.await.expect("handler task started");
824
825 drop(client);
826
827 tokio::time::timeout(Duration::from_secs(1), async {
828 while !abort_handle.is_finished() {
829 tokio::task::yield_now().await;
830 }
831 })
832 .await
833 .expect("handler task aborted");
834 }
835
836 #[rstest]
837 fn test_reserve_cid_mapping_rejects_collision() {
838 let caches = OrdersCaches::default();
839 let cid = 123;
840 let first_client_order_id = ClientOrderId::from("CID-123-A");
841 let second_client_order_id = ClientOrderId::from("CID-123-B");
842
843 reserve_cid_mapping(&caches, cid, first_client_order_id).unwrap();
844 let result = reserve_cid_mapping(&caches, cid, second_client_order_id);
845
846 assert!(matches!(
847 result,
848 Err(AxOrdersWsClientError::ClientError(msg))
849 if msg == "AX cid 123 is already mapped to CID-123-A"
850 ));
851 assert_eq!(
852 caches
853 .cid_to_client_order_id
854 .get(&cid)
855 .map(|client_order_id| *client_order_id),
856 Some(first_client_order_id),
857 );
858 }
859
860 #[tokio::test]
861 async fn test_cancel_order_rejects_without_venue_order_id() {
862 let client = AxOrdersWebSocketClient::new(
863 "wss://example.com/orders/ws".to_string(),
864 AccountId::from("AX-001"),
865 TraderId::from("TRADER-001"),
866 30,
867 TransportBackend::default(),
868 None,
869 );
870 let client_order_id = ClientOrderId::from("CID-123");
871
872 let result = client.cancel_order(client_order_id, None).await;
873
874 assert!(matches!(
875 result,
876 Err(AxOrdersWsClientError::ClientError(msg))
877 if msg.contains("missing venue_order_id")
878 ));
879 }
880
881 #[tokio::test]
882 async fn test_cancel_order_sends_known_venue_order_id() {
883 let mut client = AxOrdersWebSocketClient::new(
884 "wss://example.com/orders/ws".to_string(),
885 AccountId::from("AX-001"),
886 TraderId::from("TRADER-001"),
887 30,
888 TransportBackend::default(),
889 None,
890 );
891
892 let (cmd_tx, mut cmd_rx) = tokio::sync::mpsc::unbounded_channel::<HandlerCommand>();
893 client.cmd_tx = Arc::new(tokio::sync::RwLock::new(cmd_tx));
894
895 let client_order_id = ClientOrderId::from("CID-456");
896 let venue_order_id = VenueOrderId::from("V-ORDER-789");
897
898 let request_id = client
899 .cancel_order(client_order_id, Some(venue_order_id))
900 .await
901 .unwrap();
902
903 assert_eq!(request_id, 1);
904 let cmd = cmd_rx.recv().await.unwrap();
905 match cmd {
906 HandlerCommand::CancelOrder {
907 request_id,
908 order_id,
909 } => {
910 assert_eq!(request_id, 1);
911 assert_eq!(order_id, "V-ORDER-789");
912 }
913 other => panic!("unexpected command: {other:?}"),
914 }
915 }
916}