Skip to main content

nautilus_architect_ax/websocket/orders/
client.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 client for Ax.
17
18use 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
69/// Result type for Ax orders WebSocket operations.
70pub type AxOrdersWsResult<T> = Result<T, AxOrdersWsClientError>;
71
72/// Shared caches for order state tracking between the client and consumers.
73#[derive(Debug, Clone)]
74pub struct OrdersCaches {
75    /// Maps client order IDs to order metadata.
76    pub orders_metadata: Arc<DashMap<ClientOrderId, OrderMetadata>>,
77    /// Maps venue order IDs to client order IDs.
78    pub venue_to_client_id: Arc<DashMap<VenueOrderId, ClientOrderId>>,
79    /// Maps AX cid values to client order IDs.
80    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/// Error type for the Ax orders WebSocket client.
94#[derive(Debug, Clone)]
95pub enum AxOrdersWsClientError {
96    /// Transport/connection error.
97    Transport(String),
98    /// Channel send error.
99    ChannelError(String),
100    /// Authentication error.
101    AuthenticationError(String),
102    /// Client-side validation error.
103    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
125/// Orders WebSocket client for Ax.
126///
127/// Provides authenticated order management including placing, canceling,
128/// and monitoring order status via WebSocket.
129pub 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, // Each clone gets its own receiver
172            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    /// Creates a new Ax orders WebSocket client.
201    #[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    /// Configures socket state reporting and reconnect control.
240    #[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    /// Returns the WebSocket URL.
251    #[must_use]
252    pub fn url(&self) -> &str {
253        &self.url
254    }
255
256    /// Returns the account ID.
257    #[must_use]
258    pub fn account_id(&self) -> AccountId {
259        self.account_id
260    }
261
262    /// Returns whether the client is currently connected and active.
263    #[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    /// Returns whether the client is closed.
271    #[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    /// Generates a unique request ID.
279    fn next_request_id(&self) -> i64 {
280        self.request_id_counter.fetch_add(1, Ordering::Relaxed)
281    }
282
283    /// Caches an instrument for use during message parsing.
284    pub fn cache_instrument(&self, instrument: InstrumentAny) {
285        let symbol = instrument.symbol().inner();
286        self.instruments_cache.insert(symbol, instrument);
287    }
288
289    /// Caches multiple instruments for use during message parsing.
290    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    /// Updates the token used by future automatic reconnect attempts.
299    ///
300    /// Updating the token does not interrupt the active WebSocket connection.
301    ///
302    /// # Errors
303    ///
304    /// Returns an error if the reconnect header cannot be updated.
305    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    /// Returns a cached instrument by symbol.
317    #[must_use]
318    pub fn get_cached_instrument(&self, symbol: &Ustr) -> Option<InstrumentAny> {
319        self.instruments_cache.get_cloned(symbol)
320    }
321
322    /// Returns the shared order caches.
323    #[must_use]
324    pub fn caches(&self) -> &OrdersCaches {
325        &self.caches
326    }
327
328    /// Returns the instruments cache.
329    #[must_use]
330    pub fn instruments_cache(&self) -> Arc<AtomicMap<Ustr, InstrumentAny>> {
331        Arc::clone(&self.instruments_cache)
332    }
333
334    /// Returns the orders metadata cache.
335    #[must_use]
336    pub fn orders_metadata(&self) -> &Arc<DashMap<ClientOrderId, OrderMetadata>> {
337        &self.caches.orders_metadata
338    }
339
340    /// Returns the cid to client order ID mapping for order correlation.
341    #[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    /// Resolves a cid to a ClientOrderId if the mapping exists.
347    #[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    /// Registers an external order with the WebSocket handler for event tracking.
353    ///
354    /// This allows the handler to create proper events (e.g., OrderCanceled, OrderFilled)
355    /// for orders that were reconciled externally and not submitted through this client.
356    ///
357    /// Returns `false` if the instrument is not cached (registration skipped).
358    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        // Required for correct precision on fills
370        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    /// Establishes the WebSocket connection with authentication.
406    ///
407    /// # Arguments
408    ///
409    /// * `bearer_token` - The bearer token for authentication.
410    ///
411    /// # Errors
412    ///
413    /// Returns an error if the connection cannot be established.
414    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        // No-op ping handler: handler owns the WebSocketClient and responds to pings directly
453        let ping_handler: PingHandler = Arc::new(move |_payload: Vec<u8>| {
454            // Handler responds to pings internally via select! loop
455        });
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, // Ax server sends heartbeats
468            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    /// Submits the AX priced order shape using Nautilus domain types.
551    ///
552    /// This method handles conversion from Nautilus domain types to AX-specific
553    /// types and stores order metadata for event correlation.
554    ///
555    /// # Errors
556    ///
557    /// Returns an error if:
558    /// - The time-in-force is not supported.
559    /// - The instrument is not found in the cache.
560    /// - The order command cannot be sent.
561    #[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        // Get instrument from cache for precision
575        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        // Store order metadata for event correlation (after validation to avoid stale entries)
594        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    /// Cancels an order via WebSocket.
646    ///
647    /// Requires a known `venue_order_id`.
648    ///
649    /// # Errors
650    ///
651    /// Returns an error if the cancel command cannot be sent.
652    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    /// Requests open orders via WebSocket.
675    ///
676    /// # Errors
677    ///
678    /// Returns an error if the request command cannot be sent.
679    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    /// Returns a stream of WebSocket messages.
689    ///
690    /// # Panics
691    ///
692    /// Panics if called before `connect()` or if the stream has already been taken.
693    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    /// Disconnects the WebSocket connection gracefully.
714    pub async fn disconnect(&self) {
715        log::debug!("Disconnecting WebSocket");
716        let _ = self.send_cmd(HandlerCommand::Disconnect).await;
717    }
718
719    /// Closes the WebSocket connection and cleans up resources.
720    ///
721    /// # Errors
722    ///
723    /// Returns an error if the handler task fails or does not stop after abort.
724    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        // Send disconnect first to allow graceful cleanup before signal
730        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}