1#![allow(unused_imports)]
15use anyhow::Context;
16use async_trait::async_trait;
17use derive_builder::Builder;
18use rust_decimal::prelude::*;
19use serde::{Deserialize, Serialize};
20use serde_json::Value;
21use std::{collections::BTreeMap, sync::Arc};
22
23use crate::common::{
24 errors::WebsocketError,
25 models::{ParamBuildError, WebsocketApiResponse},
26 utils::remove_empty_value,
27 websocket::{WebsocketApi, WebsocketMessageSendOptions},
28};
29use crate::spot::websocket_api::models;
30
31#[async_trait]
32pub trait AuthApi: Send + Sync {
33 async fn session_logon(
34 &self,
35 params: SessionLogonParams,
36 ) -> anyhow::Result<Vec<WebsocketApiResponse<Box<models::SessionLogonResponseResult>>>>;
37 async fn session_logout(
38 &self,
39 params: SessionLogoutParams,
40 ) -> anyhow::Result<Vec<WebsocketApiResponse<Box<models::SessionLogoutResponseResult>>>>;
41 async fn session_status(
42 &self,
43 params: SessionStatusParams,
44 ) -> anyhow::Result<WebsocketApiResponse<Box<models::SessionStatusResponseResult>>>;
45}
46
47#[derive(Clone)]
48pub struct AuthApiClient {
49 websocket_api_base: Arc<WebsocketApi>,
50}
51
52impl AuthApiClient {
53 pub fn new(websocket_api_base: Arc<WebsocketApi>) -> Self {
54 Self { websocket_api_base }
55 }
56}
57
58#[derive(Clone, Debug, Builder, Deserialize, Default)]
63#[builder(pattern = "owned", build_fn(error = "ParamBuildError"))]
64pub struct SessionLogonParams {
65 #[builder(setter(into), default)]
69 #[serde(rename = "id", default)]
70 pub id: Option<String>,
71 #[builder(setter(into), default)]
75 #[serde(rename = "recvWindow", default)]
76 pub recv_window: Option<rust_decimal::Decimal>,
77}
78
79impl SessionLogonParams {
80 #[must_use]
83 pub fn builder() -> SessionLogonParamsBuilder {
84 SessionLogonParamsBuilder::default()
85 }
86}
87#[derive(Clone, Debug, Builder, Deserialize, Default)]
92#[builder(pattern = "owned", build_fn(error = "ParamBuildError"))]
93pub struct SessionLogoutParams {
94 #[builder(setter(into), default)]
98 #[serde(rename = "id", default)]
99 pub id: Option<String>,
100}
101
102impl SessionLogoutParams {
103 #[must_use]
106 pub fn builder() -> SessionLogoutParamsBuilder {
107 SessionLogoutParamsBuilder::default()
108 }
109}
110#[derive(Clone, Debug, Builder, Deserialize, Default)]
115#[builder(pattern = "owned", build_fn(error = "ParamBuildError"))]
116pub struct SessionStatusParams {
117 #[builder(setter(into), default)]
121 #[serde(rename = "id", default)]
122 pub id: Option<String>,
123}
124
125impl SessionStatusParams {
126 #[must_use]
129 pub fn builder() -> SessionStatusParamsBuilder {
130 SessionStatusParamsBuilder::default()
131 }
132}
133
134#[async_trait]
135impl AuthApi for AuthApiClient {
136 async fn session_logon(
137 &self,
138 params: SessionLogonParams,
139 ) -> anyhow::Result<Vec<WebsocketApiResponse<Box<models::SessionLogonResponseResult>>>> {
140 let SessionLogonParams { id, recv_window } = params;
141
142 let mut payload: BTreeMap<String, Value> = BTreeMap::new();
143 if let Some(value) = id {
144 payload.insert("id".to_string(), serde_json::json!(value));
145 }
146 if let Some(value) = recv_window {
147 payload.insert("recvWindow".to_string(), serde_json::json!(value));
148 }
149 let payload = remove_empty_value(payload);
150
151 let response = self
152 .websocket_api_base
153 .send_message::<Box<models::SessionLogonResponseResult>>(
154 "/session.logon".trim_start_matches('/'),
155 payload,
156 WebsocketMessageSendOptions::new().signed().session_logon(),
157 )
158 .await
159 .map_err(anyhow::Error::from)?
160 .into_iter()
161 .collect();
162
163 Ok(response)
164 }
165
166 async fn session_logout(
167 &self,
168 params: SessionLogoutParams,
169 ) -> anyhow::Result<Vec<WebsocketApiResponse<Box<models::SessionLogoutResponseResult>>>> {
170 let SessionLogoutParams { id } = params;
171
172 let mut payload: BTreeMap<String, Value> = BTreeMap::new();
173 if let Some(value) = id {
174 payload.insert("id".to_string(), serde_json::json!(value));
175 }
176 let payload = remove_empty_value(payload);
177
178 let response = self
179 .websocket_api_base
180 .send_message::<Box<models::SessionLogoutResponseResult>>(
181 "/session.logout".trim_start_matches('/'),
182 payload,
183 WebsocketMessageSendOptions::new().session_logout(),
184 )
185 .await
186 .map_err(anyhow::Error::from)?
187 .into_iter()
188 .collect();
189
190 Ok(response)
191 }
192
193 async fn session_status(
194 &self,
195 params: SessionStatusParams,
196 ) -> anyhow::Result<WebsocketApiResponse<Box<models::SessionStatusResponseResult>>> {
197 let SessionStatusParams { id } = params;
198
199 let mut payload: BTreeMap<String, Value> = BTreeMap::new();
200 if let Some(value) = id {
201 payload.insert("id".to_string(), serde_json::json!(value));
202 }
203 let payload = remove_empty_value(payload);
204
205 self.websocket_api_base
206 .send_message::<Box<models::SessionStatusResponseResult>>(
207 "/session.status".trim_start_matches('/'),
208 payload,
209 WebsocketMessageSendOptions::new(),
210 )
211 .await
212 .map_err(anyhow::Error::from)?
213 .into_iter()
214 .next()
215 .ok_or(WebsocketError::NoResponse)
216 .map_err(anyhow::Error::from)
217 }
218}
219
220#[cfg(all(test, feature = "spot"))]
221mod tests {
222 use super::*;
223 use crate::TOKIO_SHARED_RT;
224 use crate::common::websocket::{WebsocketApi, WebsocketConnection, WebsocketHandler};
225 use crate::config::ConfigurationWebsocketApi;
226 use crate::errors::WebsocketError;
227 use crate::models::WebsocketApiRateLimit;
228 use serde_json::{Value, json};
229 use tokio::spawn;
230 use tokio::sync::mpsc::{UnboundedReceiver, unbounded_channel};
231 use tokio::time::{Duration, timeout};
232 use tokio_tungstenite::tungstenite::Message;
233
234 async fn setup() -> (
235 Arc<WebsocketApi>,
236 Arc<WebsocketConnection>,
237 UnboundedReceiver<Message>,
238 ) {
239 let conn = WebsocketConnection::new("test-conn");
240 let (tx, rx) = unbounded_channel::<Message>();
241 {
242 let mut conn_state = conn.state.lock().await;
243 conn_state.ws_write_tx = Some(tx);
244 }
245
246 let config = ConfigurationWebsocketApi::builder()
247 .api_key("key")
248 .api_secret("secret")
249 .build()
250 .expect("Failed to build configuration");
251 let ws_api = WebsocketApi::new(config, vec![conn.clone()]);
252 conn.set_handler(ws_api.clone() as Arc<dyn WebsocketHandler>)
253 .await;
254 ws_api.clone().connect().await.unwrap();
255
256 (ws_api, conn, rx)
257 }
258
259 #[test]
260 fn session_logon_success() {
261 TOKIO_SHARED_RT.block_on(async {
262 let (ws_api, conn, mut rx) = setup().await;
263 let client = AuthApiClient::new(ws_api.clone());
264
265 let handle = spawn(async move {
266 let params = SessionLogonParams::builder().build().unwrap();
267 client.session_logon(params).await
268 });
269
270 let sent = timeout(Duration::from_secs(1), rx.recv()).await.expect("send should occur").expect("channel closed");
271 let Message::Text(text) = sent else { panic!() };
272 let v: Value = serde_json::from_str(&text).unwrap();
273 let id = v["id"].as_str().unwrap();
274 assert_eq!(v["method"], "/session.logon".trim_start_matches('/'));
275 let mut resp_json: Value = serde_json::from_str(r#"{"id":"c174a2b1-3f51-4580-b200-8528bd237cb7","status":200,"result":{"apiKey":"vmPUZE6mv9SD5VNHk4HlWFsOr6aKE2zvsw0MuIgwCIPy6utIco14y7Ju91duEh8A","authorizedSince":1649729878532,"connectedSince":1649729873021,"returnRateLimits":false,"serverTime":1649729878630,"userDataStream":false}}"#).unwrap_or_else(|_| serde_json::json!({}));
276 resp_json["id"] = id.into();
277
278 let raw_data = resp_json.get("result").or_else(|| resp_json.get("response")).expect("no response in JSON");
279 let expected_data: Box<models::SessionLogonResponseResult> = serde_json::from_value(raw_data.clone()).expect("should parse raw response");
280 let empty_array = Value::Array(vec![]);
281 let raw_rate_limits = resp_json.get("rateLimits").unwrap_or(&empty_array);
282 let expected_rate_limits: Option<Vec<WebsocketApiRateLimit>> =
283 match raw_rate_limits.as_array() {
284 Some(arr) if arr.is_empty() => None,
285 Some(_) => Some(serde_json::from_value(raw_rate_limits.clone()).expect("should parse rateLimits array")),
286 None => None,
287 };
288
289 WebsocketHandler::on_message(&*ws_api, resp_json.to_string(), conn.clone()).await;
290
291 let response = timeout(Duration::from_secs(1), handle).await.expect("task done").expect("no panic").expect("no error");
292let response = response.into_iter().next().expect("should have response");
293
294 let response_rate_limits = response.rate_limits.clone();
295 let response_data = response.data().expect("deserialize data");
296
297 assert_eq!(response_rate_limits, expected_rate_limits);
298 assert_eq!(response_data, expected_data);
299 });
300 }
301
302 #[test]
303 fn session_logon_error_response() {
304 TOKIO_SHARED_RT.block_on(async {
305 let (ws_api, conn, mut rx) = setup().await;
306 let client = AuthApiClient::new(ws_api.clone());
307
308 let handle = tokio::spawn(async move {
309 let params = SessionLogonParams::builder().build().unwrap();
310 client.session_logon(params).await
311 });
312
313 let sent = timeout(Duration::from_secs(1), rx.recv()).await.unwrap().unwrap();
314 let Message::Text(text) = sent else { panic!() };
315 let v: Value = serde_json::from_str(&text).unwrap();
316 let id = v["id"].as_str().unwrap().to_string();
317
318 let resp_json = json!({
319 "id": id,
320 "status": 400,
321 "error": {
322 "code": -2010,
323 "msg": "Account has insufficient balance for requested action.",
324 },
325 "rateLimits": [
326 {
327 "rateLimitType": "ORDERS",
328 "interval": "SECOND",
329 "intervalNum": 10,
330 "limit": 50,
331 "count": 13
332 },
333 ],
334 });
335 WebsocketHandler::on_message(&*ws_api, resp_json.to_string(), conn.clone()).await;
336
337 let join = timeout(Duration::from_secs(1), handle).await.unwrap();
338 match join {
339 Ok(Err(e)) => {
340 let msg = e.to_string();
341 assert!(
342 msg.contains("Server‐side response error (code -2010): Account has insufficient balance for requested action."),
343 "Expected error msg to contain server error, got: {msg}"
344 );
345 }
346 Ok(Ok(_)) => panic!("Expected error"),
347 Err(_) => panic!("Task panicked"),
348 }
349 });
350 }
351
352 #[test]
353 fn session_logon_request_timeout() {
354 TOKIO_SHARED_RT.block_on(async {
355 let (ws_api, _conn, mut rx) = setup().await;
356 let client = AuthApiClient::new(ws_api.clone());
357
358 let handle = spawn(async move {
359 let params = SessionLogonParams::builder().build().unwrap();
360 client.session_logon(params).await
361 });
362
363 let sent = timeout(Duration::from_secs(1), rx.recv())
364 .await
365 .expect("send should occur")
366 .expect("channel closed");
367 let Message::Text(text) = sent else {
368 panic!("expected Message Text")
369 };
370
371 let _: Value = serde_json::from_str(&text).unwrap();
372
373 let result = handle.await.expect("task completed");
374 match result {
375 Err(e) => {
376 if let Some(inner) = e.downcast_ref::<WebsocketError>() {
377 assert!(matches!(inner, WebsocketError::Timeout));
378 } else {
379 panic!("Unexpected error type: {:?}", e);
380 }
381 }
382 Ok(_) => panic!("Expected timeout error"),
383 }
384 });
385 }
386
387 #[test]
388 fn session_logout_success() {
389 TOKIO_SHARED_RT.block_on(async {
390 let (ws_api, conn, mut rx) = setup().await;
391 let client = AuthApiClient::new(ws_api.clone());
392
393 let handle = spawn(async move {
394 let params = SessionLogoutParams::builder().build().unwrap();
395 client.session_logout(params).await
396 });
397
398 let sent = timeout(Duration::from_secs(1), rx.recv()).await.expect("send should occur").expect("channel closed");
399 let Message::Text(text) = sent else { panic!() };
400 let v: Value = serde_json::from_str(&text).unwrap();
401 let id = v["id"].as_str().unwrap();
402 assert_eq!(v["method"], "/session.logout".trim_start_matches('/'));
403 let mut resp_json: Value = serde_json::from_str(r#"{"id":"c174a2b1-3f51-4580-b200-8528bd237cb7","status":200,"result":{"apiKey":"apiKey","authorizedSince":1,"connectedSince":1649729873021,"returnRateLimits":false,"serverTime":1649730611671,"userDataStream":false}}"#).unwrap_or_else(|_| serde_json::json!({}));
404 resp_json["id"] = id.into();
405
406 let raw_data = resp_json.get("result").or_else(|| resp_json.get("response")).expect("no response in JSON");
407 let expected_data: Box<models::SessionLogoutResponseResult> = serde_json::from_value(raw_data.clone()).expect("should parse raw response");
408 let empty_array = Value::Array(vec![]);
409 let raw_rate_limits = resp_json.get("rateLimits").unwrap_or(&empty_array);
410 let expected_rate_limits: Option<Vec<WebsocketApiRateLimit>> =
411 match raw_rate_limits.as_array() {
412 Some(arr) if arr.is_empty() => None,
413 Some(_) => Some(serde_json::from_value(raw_rate_limits.clone()).expect("should parse rateLimits array")),
414 None => None,
415 };
416
417 WebsocketHandler::on_message(&*ws_api, resp_json.to_string(), conn.clone()).await;
418
419 let response = timeout(Duration::from_secs(1), handle).await.expect("task done").expect("no panic").expect("no error");
420let response = response.into_iter().next().expect("should have response");
421
422 let response_rate_limits = response.rate_limits.clone();
423 let response_data = response.data().expect("deserialize data");
424
425 assert_eq!(response_rate_limits, expected_rate_limits);
426 assert_eq!(response_data, expected_data);
427 });
428 }
429
430 #[test]
431 fn session_logout_error_response() {
432 TOKIO_SHARED_RT.block_on(async {
433 let (ws_api, conn, mut rx) = setup().await;
434 let client = AuthApiClient::new(ws_api.clone());
435
436 let handle = tokio::spawn(async move {
437 let params = SessionLogoutParams::builder().build().unwrap();
438 client.session_logout(params).await
439 });
440
441 let sent = timeout(Duration::from_secs(1), rx.recv()).await.unwrap().unwrap();
442 let Message::Text(text) = sent else { panic!() };
443 let v: Value = serde_json::from_str(&text).unwrap();
444 let id = v["id"].as_str().unwrap().to_string();
445
446 let resp_json = json!({
447 "id": id,
448 "status": 400,
449 "error": {
450 "code": -2010,
451 "msg": "Account has insufficient balance for requested action.",
452 },
453 "rateLimits": [
454 {
455 "rateLimitType": "ORDERS",
456 "interval": "SECOND",
457 "intervalNum": 10,
458 "limit": 50,
459 "count": 13
460 },
461 ],
462 });
463 WebsocketHandler::on_message(&*ws_api, resp_json.to_string(), conn.clone()).await;
464
465 let join = timeout(Duration::from_secs(1), handle).await.unwrap();
466 match join {
467 Ok(Err(e)) => {
468 let msg = e.to_string();
469 assert!(
470 msg.contains("Server‐side response error (code -2010): Account has insufficient balance for requested action."),
471 "Expected error msg to contain server error, got: {msg}"
472 );
473 }
474 Ok(Ok(_)) => panic!("Expected error"),
475 Err(_) => panic!("Task panicked"),
476 }
477 });
478 }
479
480 #[test]
481 fn session_logout_request_timeout() {
482 TOKIO_SHARED_RT.block_on(async {
483 let (ws_api, _conn, mut rx) = setup().await;
484 let client = AuthApiClient::new(ws_api.clone());
485
486 let handle = spawn(async move {
487 let params = SessionLogoutParams::builder().build().unwrap();
488 client.session_logout(params).await
489 });
490
491 let sent = timeout(Duration::from_secs(1), rx.recv())
492 .await
493 .expect("send should occur")
494 .expect("channel closed");
495 let Message::Text(text) = sent else {
496 panic!("expected Message Text")
497 };
498
499 let _: Value = serde_json::from_str(&text).unwrap();
500
501 let result = handle.await.expect("task completed");
502 match result {
503 Err(e) => {
504 if let Some(inner) = e.downcast_ref::<WebsocketError>() {
505 assert!(matches!(inner, WebsocketError::Timeout));
506 } else {
507 panic!("Unexpected error type: {:?}", e);
508 }
509 }
510 Ok(_) => panic!("Expected timeout error"),
511 }
512 });
513 }
514
515 #[test]
516 fn session_status_success() {
517 TOKIO_SHARED_RT.block_on(async {
518 let (ws_api, conn, mut rx) = setup().await;
519 let client = AuthApiClient::new(ws_api.clone());
520
521 let handle = spawn(async move {
522 let params = SessionStatusParams::builder().build().unwrap();
523 client.session_status(params).await
524 });
525
526 let sent = timeout(Duration::from_secs(1), rx.recv()).await.expect("send should occur").expect("channel closed");
527 let Message::Text(text) = sent else { panic!() };
528 let v: Value = serde_json::from_str(&text).unwrap();
529 let id = v["id"].as_str().unwrap();
530 assert_eq!(v["method"], "/session.status".trim_start_matches('/'));
531 let mut resp_json: Value = serde_json::from_str(r#"{"id":"b50c16cd-62c9-4e29-89e4-37f10111f5bf","status":200,"result":{"apiKey":"vmPUZE6mv9SD5VNHk4HlWFsOr6aKE2zvsw0MuIgwCIPy6utIco14y7Ju91duEh8A","authorizedSince":1649729878532,"connectedSince":1649729873021,"returnRateLimits":false,"serverTime":1649730611671,"userDataStream":true}}"#).unwrap_or_else(|_| serde_json::json!({}));
532 resp_json["id"] = id.into();
533
534 let raw_data = resp_json.get("result").or_else(|| resp_json.get("response")).expect("no response in JSON");
535 let expected_data: Box<models::SessionStatusResponseResult> = serde_json::from_value(raw_data.clone()).expect("should parse raw response");
536 let empty_array = Value::Array(vec![]);
537 let raw_rate_limits = resp_json.get("rateLimits").unwrap_or(&empty_array);
538 let expected_rate_limits: Option<Vec<WebsocketApiRateLimit>> =
539 match raw_rate_limits.as_array() {
540 Some(arr) if arr.is_empty() => None,
541 Some(_) => Some(serde_json::from_value(raw_rate_limits.clone()).expect("should parse rateLimits array")),
542 None => None,
543 };
544
545 WebsocketHandler::on_message(&*ws_api, resp_json.to_string(), conn.clone()).await;
546
547 let response = timeout(Duration::from_secs(1), handle).await.expect("task done").expect("no panic").expect("no error");
548
549
550 let response_rate_limits = response.rate_limits.clone();
551 let response_data = response.data().expect("deserialize data");
552
553 assert_eq!(response_rate_limits, expected_rate_limits);
554 assert_eq!(response_data, expected_data);
555 });
556 }
557
558 #[test]
559 fn session_status_error_response() {
560 TOKIO_SHARED_RT.block_on(async {
561 let (ws_api, conn, mut rx) = setup().await;
562 let client = AuthApiClient::new(ws_api.clone());
563
564 let handle = tokio::spawn(async move {
565 let params = SessionStatusParams::builder().build().unwrap();
566 client.session_status(params).await
567 });
568
569 let sent = timeout(Duration::from_secs(1), rx.recv()).await.unwrap().unwrap();
570 let Message::Text(text) = sent else { panic!() };
571 let v: Value = serde_json::from_str(&text).unwrap();
572 let id = v["id"].as_str().unwrap().to_string();
573
574 let resp_json = json!({
575 "id": id,
576 "status": 400,
577 "error": {
578 "code": -2010,
579 "msg": "Account has insufficient balance for requested action.",
580 },
581 "rateLimits": [
582 {
583 "rateLimitType": "ORDERS",
584 "interval": "SECOND",
585 "intervalNum": 10,
586 "limit": 50,
587 "count": 13
588 },
589 ],
590 });
591 WebsocketHandler::on_message(&*ws_api, resp_json.to_string(), conn.clone()).await;
592
593 let join = timeout(Duration::from_secs(1), handle).await.unwrap();
594 match join {
595 Ok(Err(e)) => {
596 let msg = e.to_string();
597 assert!(
598 msg.contains("Server‐side response error (code -2010): Account has insufficient balance for requested action."),
599 "Expected error msg to contain server error, got: {msg}"
600 );
601 }
602 Ok(Ok(_)) => panic!("Expected error"),
603 Err(_) => panic!("Task panicked"),
604 }
605 });
606 }
607
608 #[test]
609 fn session_status_request_timeout() {
610 TOKIO_SHARED_RT.block_on(async {
611 let (ws_api, _conn, mut rx) = setup().await;
612 let client = AuthApiClient::new(ws_api.clone());
613
614 let handle = spawn(async move {
615 let params = SessionStatusParams::builder().build().unwrap();
616 client.session_status(params).await
617 });
618
619 let sent = timeout(Duration::from_secs(1), rx.recv())
620 .await
621 .expect("send should occur")
622 .expect("channel closed");
623 let Message::Text(text) = sent else {
624 panic!("expected Message Text")
625 };
626
627 let _: Value = serde_json::from_str(&text).unwrap();
628
629 let result = handle.await.expect("task completed");
630 match result {
631 Err(e) => {
632 if let Some(inner) = e.downcast_ref::<WebsocketError>() {
633 assert!(matches!(inner, WebsocketError::Timeout));
634 } else {
635 panic!("Unexpected error type: {:?}", e);
636 }
637 }
638 Ok(_) => panic!("Expected timeout error"),
639 }
640 });
641 }
642}