1use std::{
17 sync::{
18 Arc,
19 atomic::{AtomicU64, Ordering},
20 },
21 time::Duration,
22};
23
24use ahash::AHashMap;
25use derive_builder::Builder;
26use futures_util::future::BoxFuture;
27use nautilus_common::live::get_runtime;
28use nautilus_live::task::TaskGroup;
29use tokio::{
30 sync::{Mutex, OwnedSemaphorePermit, Semaphore, mpsc, oneshot},
31 time,
32};
33use tokio_util::sync::CancellationToken;
34
35use crate::{
36 common::{consts::HYPERLIQUID_WS_POST_INFLIGHT_MAX, enums::HyperliquidInfoRequestType},
37 http::{
38 error::{Error, Result},
39 models::{HyperliquidFills, HyperliquidL2Book, HyperliquidOrderStatus},
40 },
41 websocket::messages::{
42 ActionRequest, CancelByCloidRequest, CancelRequest, HyperliquidWsRequest, ModifyRequest,
43 OrderRequest, OrderTypeRequest, PostRequest, PostResponse, TimeInForceRequest, TpSlRequest,
44 },
45};
46
47#[derive(Debug)]
48struct Waiter {
49 tx: oneshot::Sender<PostResponse>,
50 cancellation_token: CancellationToken,
51 _permit: OwnedSemaphorePermit,
53}
54
55#[derive(Debug)]
56pub struct PostRouter {
57 inner: Mutex<AHashMap<u64, Waiter>>,
58 inflight: Arc<Semaphore>, }
60
61impl Default for PostRouter {
62 fn default() -> Self {
63 Self {
64 inner: Mutex::new(AHashMap::new()),
65 inflight: Arc::new(Semaphore::new(HYPERLIQUID_WS_POST_INFLIGHT_MAX)),
66 }
67 }
68}
69
70impl PostRouter {
71 pub fn new() -> Arc<Self> {
72 Arc::new(Self::default())
73 }
74
75 pub(super) fn with_inflight(inflight: Arc<Semaphore>) -> Arc<Self> {
76 Arc::new(Self {
77 inner: Mutex::new(AHashMap::new()),
78 inflight,
79 })
80 }
81
82 pub async fn register(&self, id: u64) -> Result<oneshot::Receiver<PostResponse>> {
84 self.register_waiter(id, &CancellationToken::new()).await
85 }
86
87 pub(super) async fn register_with_cancellation(
88 self: &Arc<Self>,
89 id: u64,
90 cancellation_token: &CancellationToken,
91 ) -> Result<oneshot::Receiver<PostResponse>> {
92 let rx = self.register_waiter(id, cancellation_token).await?;
93 let post_router = Arc::clone(self);
94 let cancellation_token = cancellation_token.clone();
95 get_runtime().spawn(async move {
96 cancellation_token.cancelled().await;
97 post_router
98 .cancel_registration(id, &cancellation_token)
99 .await;
100 });
101
102 Ok(rx)
103 }
104
105 async fn register_waiter(
106 &self,
107 id: u64,
108 cancellation_token: &CancellationToken,
109 ) -> Result<oneshot::Receiver<PostResponse>> {
110 let permit = self
112 .inflight
113 .clone()
114 .acquire_owned()
115 .await
116 .map_err(|_| Error::transport("post router semaphore closed"))?;
117
118 let (tx, rx) = oneshot::channel::<PostResponse>();
119 let mut map = self.inner.lock().await;
120 if map.contains_key(&id) {
121 return Err(Error::transport(format!("post id {id} already registered")));
122 }
123 map.insert(
124 id,
125 Waiter {
126 tx,
127 cancellation_token: cancellation_token.clone(),
128 _permit: permit,
129 },
130 );
131 Ok(rx)
132 }
133
134 pub async fn complete(&self, resp: PostResponse) {
136 let id = resp.id;
137 let waiter = {
138 let mut map = self.inner.lock().await;
139 map.remove(&id)
140 };
141
142 if let Some(waiter) = waiter {
143 waiter.cancellation_token.cancel();
144 if waiter.tx.send(resp).is_err() {
145 log::warn!("Post waiter dropped before delivery: id={id}");
146 }
147 } else {
149 log::warn!("Post response with unknown id (late/duplicate?): id={id}");
150 }
151 }
152
153 pub async fn cancel(&self, id: u64) {
155 let waiter = self.inner.lock().await.remove(&id);
156 if let Some(waiter) = waiter {
157 waiter.cancellation_token.cancel();
158 }
159 }
161
162 pub(super) async fn cancel_registration(
163 &self,
164 id: u64,
165 cancellation_token: &CancellationToken,
166 ) {
167 let waiter = {
168 let mut map = self.inner.lock().await;
169 if map
170 .get(&id)
171 .is_some_and(|waiter| &waiter.cancellation_token == cancellation_token)
172 {
173 map.remove(&id)
174 } else {
175 None
176 }
177 };
178
179 if let Some(waiter) = waiter {
180 waiter.cancellation_token.cancel();
181 }
182 }
183
184 pub async fn await_with_timeout(
186 &self,
187 id: u64,
188 rx: oneshot::Receiver<PostResponse>,
189 timeout: Duration,
190 ) -> Result<PostResponse> {
191 match time::timeout(timeout, rx).await {
192 Ok(Ok(resp)) => Ok(resp),
193 Ok(Err(_closed)) => {
194 self.cancel(id).await;
195 Err(Error::transport("post response channel closed"))
196 }
197 Err(_elapsed) => {
198 self.cancel(id).await;
199 Err(Error::Timeout)
200 }
201 }
202 }
203}
204
205#[derive(Debug)]
206pub struct PostIds(AtomicU64);
207
208impl PostIds {
209 pub fn new(start: u64) -> Self {
210 Self(AtomicU64::new(start))
211 }
212 pub fn next(&self) -> u64 {
213 self.0.fetch_add(1, Ordering::Relaxed)
214 }
215}
216
217#[derive(Debug, Clone, Copy, PartialEq, Eq)]
218pub enum PostLane {
219 Alo, Normal, }
222
223#[derive(Debug)]
224pub struct ScheduledPost {
225 pub id: u64,
226 pub request: PostRequest,
227 pub lane: PostLane,
228}
229
230#[derive(Debug)]
231pub struct PostBatcher {
232 tx_alo: mpsc::Sender<ScheduledPost>,
233 tx_normal: mpsc::Sender<ScheduledPost>,
234 _tasks: TaskGroup,
235}
236
237impl PostBatcher {
238 pub fn new<F>(send_fn: F) -> Self
244 where
245 F: Send + 'static + Clone + FnMut(HyperliquidWsRequest) -> BoxFuture<'static, Result<()>>,
246 {
247 let (tx_alo, rx_alo) = mpsc::channel::<ScheduledPost>(1024);
248 let (tx_normal, rx_normal) = mpsc::channel::<ScheduledPost>(4096);
249 let tasks = TaskGroup::new();
250
251 tasks
253 .spawn(Self::run_lane(
254 "ALO",
255 rx_alo,
256 Duration::from_millis(100),
257 send_fn.clone(),
258 ))
259 .expect("new post batcher accepts ALO lane task");
260
261 tasks
263 .spawn(Self::run_lane(
264 "NORMAL",
265 rx_normal,
266 Duration::from_millis(50),
267 send_fn,
268 ))
269 .expect("new post batcher accepts normal lane task");
270
271 Self {
272 tx_alo,
273 tx_normal,
274 _tasks: tasks,
275 }
276 }
277
278 async fn run_lane<F>(
279 lane_name: &'static str,
280 mut rx: mpsc::Receiver<ScheduledPost>,
281 tick: Duration,
282 mut send_fn: F,
283 ) where
284 F: Send + 'static + FnMut(HyperliquidWsRequest) -> BoxFuture<'static, Result<()>>,
285 {
286 let mut pend: Vec<ScheduledPost> = Vec::with_capacity(128);
287 let mut interval = time::interval(tick);
288 interval.set_missed_tick_behavior(time::MissedTickBehavior::Delay);
289
290 loop {
291 tokio::select! {
292 maybe_item = rx.recv() => {
293 match maybe_item {
294 Some(item) => pend.push(item),
295 None => break, }
297 }
298 _ = interval.tick() => {
299 if pend.is_empty() { continue; }
300 let to_send = std::mem::take(&mut pend);
301 for item in to_send {
302 let req = HyperliquidWsRequest::Post { id: item.id, request: item.request.clone() };
303 if let Err(e) = send_fn(req).await {
304 log::error!("Failed to send post: lane={lane_name}, id={}, {e}", item.id);
305 }
306 }
307 }
308 }
309 }
310 log::debug!("Post lane terminated: lane={lane_name}");
311 }
312
313 pub async fn enqueue(&self, item: ScheduledPost) -> Result<()> {
314 match item.lane {
315 PostLane::Alo => self
316 .tx_alo
317 .send(item)
318 .await
319 .map_err(|_| Error::transport("ALO lane closed")),
320 PostLane::Normal => self
321 .tx_normal
322 .send(item)
323 .await
324 .map_err(|_| Error::transport("NORMAL lane closed")),
325 }
326 }
327}
328
329pub fn lane_for_action(action: &ActionRequest) -> PostLane {
331 match action {
332 ActionRequest::Order { orders, .. } => {
333 if orders.is_empty() {
334 return PostLane::Normal;
335 }
336 let all_alo = orders.iter().all(|o| {
337 matches!(
338 o.t,
339 OrderTypeRequest::Limit {
340 tif: TimeInForceRequest::Alo
341 }
342 )
343 });
344
345 if all_alo {
346 PostLane::Alo
347 } else {
348 PostLane::Normal
349 }
350 }
351 _ => PostLane::Normal,
352 }
353}
354
355#[derive(Debug, Clone, Copy, Default)]
356pub enum Grouping {
357 #[default]
358 Na,
359 NormalTpsl,
360 PositionTpsl,
361}
362impl Grouping {
363 pub fn as_str(&self) -> &'static str {
364 match self {
365 Self::Na => "na",
366 Self::NormalTpsl => "normalTpsl",
367 Self::PositionTpsl => "positionTpsl",
368 }
369 }
370}
371
372#[derive(Debug, Clone, Builder)]
374pub struct LimitOrderParams {
375 pub asset: u32,
376 pub is_buy: bool,
377 pub px: String,
378 pub sz: String,
379 pub reduce_only: bool,
380 pub tif: TimeInForceRequest,
381 pub cloid: Option<String>,
382}
383
384#[derive(Debug, Clone, Builder)]
386pub struct TriggerOrderParams {
387 pub asset: u32,
388 pub is_buy: bool,
389 pub px: String,
390 pub sz: String,
391 pub reduce_only: bool,
392 pub is_market: bool,
393 pub trigger_px: String,
394 pub tpsl: TpSlRequest,
395 pub cloid: Option<String>,
396}
397
398#[derive(Debug, Default)]
400pub struct OrderBuilder {
401 orders: Vec<OrderRequest>,
402 grouping: Grouping,
403}
404
405impl OrderBuilder {
406 pub fn new() -> Self {
407 Self::default()
408 }
409
410 #[must_use]
411 pub fn grouping(mut self, g: Grouping) -> Self {
412 self.grouping = g;
413 self
414 }
415
416 #[expect(clippy::too_many_arguments)]
418 #[must_use]
419 pub fn push_limit(
420 self,
421 asset: u32,
422 is_buy: bool,
423 px: &(impl ToString + ?Sized),
424 sz: &(impl ToString + ?Sized),
425 reduce_only: bool,
426 tif: TimeInForceRequest,
427 cloid: Option<String>,
428 ) -> Self {
429 let params = LimitOrderParams {
430 asset,
431 is_buy,
432 px: px.to_string(),
433 sz: sz.to_string(),
434 reduce_only,
435 tif,
436 cloid,
437 };
438 self.push_limit_order(params)
439 }
440
441 #[must_use]
443 pub fn push_limit_order(mut self, params: LimitOrderParams) -> Self {
444 self.orders.push(OrderRequest {
445 a: params.asset,
446 b: params.is_buy,
447 p: params.px,
448 s: params.sz,
449 r: params.reduce_only,
450 t: OrderTypeRequest::Limit { tif: params.tif },
451 c: params.cloid,
452 });
453 self
454 }
455
456 #[expect(clippy::too_many_arguments)]
458 #[must_use]
459 pub fn push_trigger(
460 self,
461 asset: u32,
462 is_buy: bool,
463 px: &(impl ToString + ?Sized),
464 sz: &(impl ToString + ?Sized),
465 reduce_only: bool,
466 is_market: bool,
467 trigger_px: &(impl ToString + ?Sized),
468 tpsl: TpSlRequest,
469 cloid: Option<String>,
470 ) -> Self {
471 let params = TriggerOrderParams {
472 asset,
473 is_buy,
474 px: px.to_string(),
475 sz: sz.to_string(),
476 reduce_only,
477 is_market,
478 trigger_px: trigger_px.to_string(),
479 tpsl,
480 cloid,
481 };
482 self.push_trigger_order(params)
483 }
484
485 #[must_use]
487 pub fn push_trigger_order(mut self, params: TriggerOrderParams) -> Self {
488 self.orders.push(OrderRequest {
489 a: params.asset,
490 b: params.is_buy,
491 p: params.px,
492 s: params.sz,
493 r: params.reduce_only,
494 t: OrderTypeRequest::Trigger {
495 is_market: params.is_market,
496 trigger_px: params.trigger_px,
497 tpsl: params.tpsl,
498 },
499 c: params.cloid,
500 });
501 self
502 }
503 pub fn build(self) -> ActionRequest {
504 ActionRequest::Order {
505 orders: self.orders,
506 grouping: self.grouping.as_str().to_string(),
507 }
508 }
509
510 pub fn single_limit_order(params: LimitOrderParams) -> ActionRequest {
527 Self::new().push_limit_order(params).build()
528 }
529
530 pub fn single_trigger_order(params: TriggerOrderParams) -> ActionRequest {
549 Self::new().push_trigger_order(params).build()
550 }
551}
552
553pub fn cancel_many(cancels: Vec<(u32, u64)>) -> ActionRequest {
554 ActionRequest::Cancel {
555 cancels: cancels
556 .into_iter()
557 .map(|(a, o)| CancelRequest { a, o })
558 .collect(),
559 fast: None,
560 }
561}
562pub fn cancel_by_cloid(asset: u32, cloid: impl Into<String>) -> ActionRequest {
563 ActionRequest::CancelByCloid {
564 cancels: vec![CancelByCloidRequest {
565 asset,
566 cloid: cloid.into(),
567 }],
568 fast: None,
569 }
570}
571pub fn modify(oid: u64, new_order: OrderRequest) -> ActionRequest {
572 ActionRequest::Modify {
573 modifies: vec![ModifyRequest {
574 oid,
575 order: new_order,
576 }],
577 }
578}
579
580pub fn info_l2_book(coin: &str) -> PostRequest {
581 PostRequest::Info {
582 payload: serde_json::json!({"type": HyperliquidInfoRequestType::L2Book.as_str(), "coin": coin}),
583 }
584}
585
586pub fn info_all_mids() -> PostRequest {
587 PostRequest::Info {
588 payload: serde_json::json!({"type": HyperliquidInfoRequestType::AllMids.as_str()}),
589 }
590}
591
592pub fn info_order_status(user: &str, oid: u64) -> PostRequest {
593 PostRequest::Info {
594 payload: serde_json::json!({"type": HyperliquidInfoRequestType::OrderStatus.as_str(), "user": user, "oid": oid}),
595 }
596}
597
598pub fn info_open_orders(user: &str, frontend: Option<bool>) -> PostRequest {
599 let mut body =
600 serde_json::json!({"type": HyperliquidInfoRequestType::OpenOrders.as_str(), "user": user});
601
602 if let Some(fe) = frontend {
603 body["frontend"] = serde_json::json!(fe);
604 }
605 PostRequest::Info { payload: body }
606}
607
608pub fn info_user_fills(user: &str, aggregate_by_time: Option<bool>) -> PostRequest {
609 let mut body =
610 serde_json::json!({"type": HyperliquidInfoRequestType::UserFills.as_str(), "user": user});
611
612 if let Some(agg) = aggregate_by_time {
613 body["aggregateByTime"] = serde_json::json!(agg);
614 }
615 PostRequest::Info { payload: body }
616}
617
618pub fn info_user_rate_limit(user: &str) -> PostRequest {
619 PostRequest::Info {
620 payload: serde_json::json!({"type": HyperliquidInfoRequestType::UserRateLimit.as_str(), "user": user}),
621 }
622}
623
624pub fn info_candle(coin: &str, interval: &str) -> PostRequest {
625 PostRequest::Info {
626 payload: serde_json::json!({"type": HyperliquidInfoRequestType::Candle.as_str(), "coin": coin, "interval": interval}),
627 }
628}
629
630pub fn parse_l2_book(payload: &serde_json::Value) -> Result<HyperliquidL2Book> {
631 serde_json::from_value(payload.clone()).map_err(Error::Serde)
632}
633pub fn parse_user_fills(payload: &serde_json::Value) -> Result<HyperliquidFills> {
634 serde_json::from_value(payload.clone()).map_err(Error::Serde)
635}
636pub fn parse_order_status(payload: &serde_json::Value) -> Result<HyperliquidOrderStatus> {
637 serde_json::from_value(payload.clone()).map_err(Error::Serde)
638}
639
640#[derive(Debug)]
642pub enum ActionOutcome<'a> {
643 Resting {
644 oid: u64,
645 },
646 Filled {
647 total_sz: &'a str,
648 avg_px: &'a str,
649 oid: Option<u64>,
650 },
651 Error {
652 msg: &'a str,
653 },
654 Unknown(&'a serde_json::Value),
655}
656pub fn classify_action_payload(payload: &serde_json::Value) -> ActionOutcome<'_> {
657 if let Some(oid) = payload.get("oid").and_then(|v| v.as_u64()) {
658 if let (Some(total_sz), Some(avg_px)) = (
659 payload.get("totalSz").and_then(|v| v.as_str()),
660 payload.get("avgPx").and_then(|v| v.as_str()),
661 ) {
662 return ActionOutcome::Filled {
663 total_sz,
664 avg_px,
665 oid: Some(oid),
666 };
667 }
668 return ActionOutcome::Resting { oid };
669 }
670
671 if let (Some(total_sz), Some(avg_px)) = (
672 payload.get("totalSz").and_then(|v| v.as_str()),
673 payload.get("avgPx").and_then(|v| v.as_str()),
674 ) {
675 return ActionOutcome::Filled {
676 total_sz,
677 avg_px,
678 oid: None,
679 };
680 }
681
682 if let Some(msg) = payload
683 .get("error")
684 .and_then(|v| v.as_str())
685 .or_else(|| payload.get("message").and_then(|v| v.as_str()))
686 {
687 return ActionOutcome::Error { msg };
688 }
689 ActionOutcome::Unknown(payload)
690}
691
692#[derive(Clone, Debug)]
693pub struct WsSender {
694 inner: mpsc::Sender<HyperliquidWsRequest>,
695}
696
697impl WsSender {
698 pub fn new(tx: mpsc::Sender<HyperliquidWsRequest>) -> Self {
699 Self { inner: tx }
700 }
701
702 pub async fn send(&self, req: HyperliquidWsRequest) -> Result<()> {
703 self.inner
704 .send(req)
705 .await
706 .map_err(|_| Error::transport("WebSocket sender closed"))
707 }
708}
709
710#[cfg(test)]
711mod tests {
712 use std::sync::atomic::AtomicUsize;
713
714 use nautilus_common::{live::get_runtime, testing::wait_until_async};
715 use rstest::rstest;
716 use tokio::{
717 sync::oneshot,
718 time::{Duration, timeout},
719 };
720
721 use super::*;
722 use crate::{
723 common::consts::HYPERLIQUID_WS_POST_INFLIGHT_MAX,
724 websocket::messages::{
725 ActionRequest, CancelByCloidRequest, CancelRequest, HyperliquidWsRequest, OrderRequest,
726 OrderRequestBuilder, OrderTypeRequest, PostResponsePayload, TimeInForceRequest,
727 },
728 };
729
730 struct DropCounter(Arc<AtomicUsize>);
731
732 impl Drop for DropCounter {
733 fn drop(&mut self) {
734 self.0.fetch_add(1, Ordering::Relaxed);
735 }
736 }
737
738 fn mk_limit_alo(asset: u32) -> OrderRequest {
739 OrderRequest {
740 a: asset,
741 b: true,
742 p: "1".to_string(),
743 s: "1".to_string(),
744 r: false,
745 t: OrderTypeRequest::Limit {
746 tif: TimeInForceRequest::Alo,
747 },
748 c: None,
749 }
750 }
751
752 fn mk_limit_gtc(asset: u32) -> OrderRequest {
753 OrderRequest {
754 a: asset,
755 b: true,
756 p: "1".to_string(),
757 s: "1".to_string(),
758 r: false,
759 t: OrderTypeRequest::Limit {
760 tif: TimeInForceRequest::Gtc,
762 },
763 c: None,
764 }
765 }
766
767 #[rstest]
768 #[tokio::test]
769 async fn test_ws_sender_forwards_and_reports_closed_channel() {
770 let (tx, mut rx) = mpsc::channel(1);
771 let sender = WsSender::new(tx);
772
773 sender.send(HyperliquidWsRequest::Ping).await.unwrap();
774 assert!(matches!(rx.recv().await, Some(HyperliquidWsRequest::Ping)));
775
776 drop(rx);
777 let error = sender.send(HyperliquidWsRequest::Ping).await.unwrap_err();
778 assert_eq!(
779 error.to_string(),
780 "transport error: WebSocket sender closed"
781 );
782 }
783
784 #[rstest]
785 #[tokio::test(flavor = "multi_thread")]
786 async fn register_duplicate_id_errors() {
787 let router = PostRouter::new();
788 let _rx = router.register(42).await.expect("first register OK");
789
790 let err = router.register(42).await.expect_err("duplicate must error");
791 let msg = err.to_string().to_lowercase();
792 assert!(
793 msg.contains("already") || msg.contains("duplicate"),
794 "unexpected error: {msg}"
795 );
796 }
797
798 #[rstest]
799 #[tokio::test(flavor = "multi_thread")]
800 async fn timeout_cancels_and_allows_reregister() {
801 let router = PostRouter::new();
802 let id = 7;
803
804 let rx = router.register(id).await.unwrap();
805 let err = router
807 .await_with_timeout(id, rx, Duration::from_millis(25))
808 .await
809 .expect_err("should timeout");
810 assert!(
811 err.to_string().to_lowercase().contains("timeout")
812 || err.to_string().to_lowercase().contains("closed"),
813 "unexpected error kind: {err}"
814 );
815
816 let _rx2 = router
818 .register(id)
819 .await
820 .expect("id should be reusable after timeout cancel");
821 }
822
823 #[rstest]
824 #[tokio::test]
825 async fn complete_cancels_registration_cleanup_and_allows_reregister() {
826 let router = PostRouter::new();
827 let id = 8;
828 let cancellation_token = CancellationToken::new();
829 let rx = router
830 .register_with_cancellation(id, &cancellation_token)
831 .await
832 .unwrap();
833
834 router
835 .complete(PostResponse {
836 id,
837 response: PostResponsePayload::Info {
838 payload: serde_json::json!({"status": "ok"}),
839 },
840 })
841 .await;
842 let response = rx.await.unwrap();
843
844 assert_eq!(response.id, id);
845 assert!(matches!(
846 response.response,
847 PostResponsePayload::Info { .. }
848 ));
849 assert!(cancellation_token.is_cancelled());
850 router
851 .register(id)
852 .await
853 .expect("id should be reusable after completion");
854 }
855
856 #[rstest]
857 #[tokio::test(flavor = "multi_thread")]
858 async fn inflight_cap_blocks_then_unblocks() {
859 let router = PostRouter::new();
860
861 let mut rxs = Vec::with_capacity(HYPERLIQUID_WS_POST_INFLIGHT_MAX);
863 for i in 0..HYPERLIQUID_WS_POST_INFLIGHT_MAX {
864 let rx = router.register(i as u64).await.unwrap();
865 rxs.push(rx); }
867
868 let router2 = Arc::clone(&router);
870 let (entered_tx, entered_rx) = oneshot::channel::<()>();
871 let (done_tx, done_rx) = oneshot::channel::<()>();
872 let (check_tx, check_rx) = oneshot::channel::<()>(); get_runtime().spawn(async move {
875 let _ = entered_tx.send(());
876 let _rx = router2.register(9_999_999).await.unwrap();
877 let _ = done_tx.send(());
878 });
879
880 entered_rx.await.unwrap();
882
883 get_runtime().spawn(async move {
885 if done_rx.await.is_ok() {
886 let _ = check_tx.send(());
887 }
888 });
889
890 assert!(
891 timeout(Duration::from_millis(50), check_rx).await.is_err(),
892 "should still be blocked while at cap"
893 );
894
895 router.cancel(0).await;
897
898 tokio::time::sleep(Duration::from_millis(100)).await;
900 }
901
902 #[rstest(
903 orders, expected,
904 case::all_alo(vec![mk_limit_alo(0), mk_limit_alo(1)], PostLane::Alo),
905 case::mixed_alo_gtc(vec![mk_limit_alo(0), mk_limit_gtc(1)], PostLane::Normal),
906 case::all_gtc(vec![mk_limit_gtc(0), mk_limit_gtc(1)], PostLane::Normal),
907 case::empty(vec![], PostLane::Normal),
908 )]
909 fn lane_classifier_cases(orders: Vec<OrderRequest>, expected: PostLane) {
910 let action = ActionRequest::Order {
911 orders,
912 grouping: "na".to_string(),
913 };
914 assert_eq!(lane_for_action(&action), expected);
915 }
916
917 #[rstest]
918 fn test_order_request_builder() {
919 let order = OrderRequestBuilder::default()
921 .a(0)
922 .b(true)
923 .p("40000.0".to_string())
924 .s("0.01".to_string())
925 .r(false)
926 .t(OrderTypeRequest::Limit {
927 tif: TimeInForceRequest::Gtc,
928 })
929 .c(Some("test-order-1".to_string()))
930 .build()
931 .expect("should build order");
932
933 assert_eq!(order.a, 0);
934 assert!(order.b);
935 assert_eq!(order.p, "40000.0");
936 assert_eq!(order.s, "0.01");
937 assert!(!order.r);
938 assert_eq!(order.c, Some("test-order-1".to_string()));
939 }
940
941 #[rstest]
942 fn test_limit_order_params_builder() {
943 let params = LimitOrderParamsBuilder::default()
945 .asset(0)
946 .is_buy(true)
947 .px("40000.0".to_string())
948 .sz("0.01".to_string())
949 .reduce_only(false)
950 .tif(TimeInForceRequest::Alo)
951 .cloid(Some("test-limit-1".to_string()))
952 .build()
953 .expect("should build limit params");
954
955 assert_eq!(params.asset, 0);
956 assert!(params.is_buy);
957 assert_eq!(params.px, "40000.0");
958 assert_eq!(params.sz, "0.01");
959 assert!(!params.reduce_only);
960 assert_eq!(params.cloid, Some("test-limit-1".to_string()));
961 }
962
963 #[rstest]
964 fn test_trigger_order_params_builder() {
965 let params = TriggerOrderParamsBuilder::default()
967 .asset(1)
968 .is_buy(false)
969 .px("39000.0".to_string())
970 .sz("0.02".to_string())
971 .reduce_only(false)
972 .is_market(true)
973 .trigger_px("39500.0".to_string())
974 .tpsl(TpSlRequest::Sl)
975 .cloid(Some("test-trigger-1".to_string()))
976 .build()
977 .expect("should build trigger params");
978
979 assert_eq!(params.asset, 1);
980 assert!(!params.is_buy);
981 assert_eq!(params.px, "39000.0");
982 assert!(params.is_market);
983 assert_eq!(params.trigger_px, "39500.0");
984 }
985
986 #[rstest]
987 fn test_order_builder_single_limit_convenience() {
988 let params = LimitOrderParamsBuilder::default()
990 .asset(0)
991 .is_buy(true)
992 .px("40000.0".to_string())
993 .sz("0.01".to_string())
994 .reduce_only(false)
995 .tif(TimeInForceRequest::Gtc)
996 .cloid(None)
997 .build()
998 .unwrap();
999
1000 let action = OrderBuilder::single_limit_order(params);
1001
1002 match action {
1003 ActionRequest::Order { orders, grouping } => {
1004 assert_eq!(orders.len(), 1);
1005 assert_eq!(orders[0].a, 0);
1006 assert!(orders[0].b);
1007 assert_eq!(grouping, "na");
1008 }
1009 _ => panic!("Expected ActionRequest::Order variant"),
1010 }
1011 }
1012
1013 #[rstest]
1014 fn test_order_builder_single_trigger_convenience() {
1015 let params = TriggerOrderParamsBuilder::default()
1017 .asset(1)
1018 .is_buy(false)
1019 .px("39000.0".to_string())
1020 .sz("0.02".to_string())
1021 .reduce_only(false)
1022 .is_market(true)
1023 .trigger_px("39500.0".to_string())
1024 .tpsl(TpSlRequest::Sl)
1025 .cloid(Some("sl-order".to_string()))
1026 .build()
1027 .unwrap();
1028
1029 let action = OrderBuilder::single_trigger_order(params);
1030
1031 match action {
1032 ActionRequest::Order { orders, grouping } => {
1033 assert_eq!(orders.len(), 1);
1034 assert_eq!(orders[0].a, 1);
1035 assert_eq!(orders[0].c, Some("sl-order".to_string()));
1036 assert_eq!(grouping, "na");
1037 }
1038 _ => panic!("Expected ActionRequest::Order variant"),
1039 }
1040 }
1041
1042 #[rstest]
1043 fn test_order_builder_batch_orders() {
1044 let params1 = LimitOrderParams {
1046 asset: 0,
1047 is_buy: true,
1048 px: "40000.0".to_string(),
1049 sz: "0.01".to_string(),
1050 reduce_only: false,
1051 tif: TimeInForceRequest::Gtc,
1052 cloid: Some("order-1".to_string()),
1053 };
1054
1055 let params2 = LimitOrderParams {
1056 asset: 1,
1057 is_buy: false,
1058 px: "2000.0".to_string(),
1059 sz: "0.5".to_string(),
1060 reduce_only: false,
1061 tif: TimeInForceRequest::Ioc,
1062 cloid: Some("order-2".to_string()),
1063 };
1064
1065 let action = OrderBuilder::new()
1066 .grouping(Grouping::NormalTpsl)
1067 .push_limit_order(params1)
1068 .push_limit_order(params2)
1069 .build();
1070
1071 match action {
1072 ActionRequest::Order { orders, grouping } => {
1073 assert_eq!(orders.len(), 2);
1074 assert_eq!(orders[0].c, Some("order-1".to_string()));
1075 assert_eq!(orders[1].c, Some("order-2".to_string()));
1076 assert_eq!(grouping, "normalTpsl");
1077 }
1078 _ => panic!("Expected ActionRequest::Order variant"),
1079 }
1080 }
1081
1082 #[rstest]
1083 fn test_action_request_constructors() {
1084 let order1 = mk_limit_gtc(0);
1086 let order2 = mk_limit_gtc(1);
1087 let action = ActionRequest::order(vec![order1, order2], "na");
1088
1089 match action {
1090 ActionRequest::Order { orders, grouping } => {
1091 assert_eq!(orders.len(), 2);
1092 assert_eq!(grouping, "na");
1093 }
1094 _ => panic!("Expected ActionRequest::Order variant"),
1095 }
1096
1097 let cancels = vec![CancelRequest { a: 0, o: 12345 }];
1099 let action = ActionRequest::cancel(cancels);
1100 assert!(matches!(action, ActionRequest::Cancel { .. }));
1101
1102 let cancels = vec![CancelByCloidRequest {
1104 asset: 0,
1105 cloid: "order-1".to_string(),
1106 }];
1107 let action = ActionRequest::cancel_by_cloid(cancels);
1108 assert!(matches!(action, ActionRequest::CancelByCloid { .. }));
1109 }
1110
1111 #[rstest]
1112 #[tokio::test(flavor = "multi_thread")]
1113 async fn batcher_sends_on_tick() {
1114 let sent: Arc<tokio::sync::Mutex<Vec<u64>>> = Arc::new(tokio::sync::Mutex::new(Vec::new()));
1116 let sent_closure = sent.clone();
1117
1118 let send_fn = move |req: HyperliquidWsRequest| -> BoxFuture<'static, Result<()>> {
1119 let sent_inner = sent_closure.clone();
1120 Box::pin(async move {
1121 if let HyperliquidWsRequest::Post { id, .. } = req {
1122 sent_inner.lock().await.push(id);
1123 }
1124 Ok(())
1125 })
1126 };
1127
1128 let batcher = PostBatcher::new(send_fn);
1129
1130 for id in 1..=5u64 {
1132 batcher
1133 .enqueue(ScheduledPost {
1134 id,
1135 request: info_all_mids(),
1136 lane: PostLane::Normal,
1137 })
1138 .await
1139 .unwrap();
1140 }
1141
1142 let sent_check = sent.clone();
1144 wait_until_async(
1145 || {
1146 let sent_inner = sent_check.clone();
1147 async move { sent_inner.lock().await.len() == 5 }
1148 },
1149 Duration::from_secs(2),
1150 )
1151 .await;
1152
1153 let actual = sent.lock().await.clone();
1154 assert_eq!(actual, vec![1, 2, 3, 4, 5]);
1155 }
1156
1157 #[rstest]
1158 #[tokio::test]
1159 async fn test_batcher_drop_aborts_lane_tasks() {
1160 let started = Arc::new(AtomicUsize::new(0));
1161 let dropped = Arc::new(AtomicUsize::new(0));
1162 let started_send = Arc::clone(&started);
1163 let dropped_send = Arc::clone(&dropped);
1164 let send_fn = move |_req: HyperliquidWsRequest| -> BoxFuture<'static, Result<()>> {
1165 let started = Arc::clone(&started_send);
1166 let dropped = Arc::clone(&dropped_send);
1167 Box::pin(async move {
1168 let _drop_counter = DropCounter(dropped);
1169 started.fetch_add(1, Ordering::Relaxed);
1170 std::future::pending::<Result<()>>().await
1171 })
1172 };
1173 let batcher = PostBatcher::new(send_fn);
1174
1175 for (id, lane) in [(1, PostLane::Alo), (2, PostLane::Normal)] {
1176 batcher
1177 .enqueue(ScheduledPost {
1178 id,
1179 request: info_all_mids(),
1180 lane,
1181 })
1182 .await
1183 .unwrap();
1184 }
1185 let started_check = Arc::clone(&started);
1186 wait_until_async(
1187 || {
1188 let started = Arc::clone(&started_check);
1189 async move { started.load(Ordering::Relaxed) == 2 }
1190 },
1191 Duration::from_secs(1),
1192 )
1193 .await;
1194
1195 drop(batcher);
1196
1197 let dropped_check = Arc::clone(&dropped);
1198 wait_until_async(
1199 || {
1200 let dropped = Arc::clone(&dropped_check);
1201 async move { dropped.load(Ordering::Relaxed) == 2 }
1202 },
1203 Duration::from_secs(1),
1204 )
1205 .await;
1206 assert_eq!(dropped.load(Ordering::Relaxed), 2);
1207 }
1208}