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    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
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<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, // 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: proxy_url.map(SecretString::from),
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 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, // 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
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    /// Submits the AX priced order shape using Nautilus domain types.
554    ///
555    /// This method handles conversion from Nautilus domain types to AX-specific
556    /// types and stores order metadata for event correlation.
557    ///
558    /// # Errors
559    ///
560    /// Returns an error if:
561    /// - The time-in-force is not supported.
562    /// - The instrument is not found in the cache.
563    /// - The order command cannot be sent.
564    #[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        // Get instrument from cache for precision
578        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        // Store order metadata for event correlation (after validation to avoid stale entries)
597        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    /// Cancels an order via WebSocket.
649    ///
650    /// Requires a known `venue_order_id`.
651    ///
652    /// # Errors
653    ///
654    /// Returns an error if the cancel command cannot be sent.
655    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    /// Requests open orders via WebSocket.
678    ///
679    /// # Errors
680    ///
681    /// Returns an error if the request command cannot be sent.
682    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    /// Returns a stream of WebSocket messages.
692    ///
693    /// # Panics
694    ///
695    /// Panics if called before `connect()` or if the stream has already been taken.
696    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    /// Disconnects the WebSocket connection gracefully.
717    pub async fn disconnect(&self) {
718        log::debug!("Disconnecting WebSocket");
719        let _ = self.send_cmd(HandlerCommand::Disconnect).await;
720    }
721
722    /// Closes the WebSocket connection and cleans up resources.
723    ///
724    /// # Errors
725    ///
726    /// Returns an error if the handler task fails or does not stop after abort.
727    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        // Send disconnect first to allow graceful cleanup before signal
733        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}