xapi_okx/common/
executor.rs1use crate::common::{
2 endpoint::OkxEndpoint,
3 ratelimiter::{OkxRateLimitKey, OkxRatelimiter},
4 signer::OkxSigner,
5 ws::stream::{OkxWsStream, OkxWsStreamCall},
6};
7use serde::de::DeserializeOwned;
8use std::sync::Arc;
9use tokio::sync::{Mutex, mpsc, oneshot};
10use typed_builder::TypedBuilder;
11use xapi_shared::{
12 ratelimiter::SharedRatelimiterTrait, rest::SharedRestClientTrait, signer::SharedSignerTrait,
13 ws::error::SharedWsError,
14};
15
16#[derive(TypedBuilder)]
17pub struct OkxExecutor {
18 endpoint: OkxEndpoint,
19 #[builder(default = reqwest::Client::new())]
20 rest_client: reqwest::Client,
21 #[builder(default = None, setter(strip_option))]
22 signer: Option<OkxSigner>,
23 #[builder(default = Arc::new(OkxRatelimiter::default()))]
24 ratelimiter: Arc<OkxRatelimiter>,
25 #[builder(default)]
26 streams: Mutex<Vec<ezsockets::Client<OkxWsStream>>>,
27}
28
29impl SharedRestClientTrait<OkxRateLimitKey> for OkxExecutor {
30 fn get_client(&self) -> &reqwest::Client {
31 &self.rest_client
32 }
33
34 fn get_signer(&self) -> &dyn SharedSignerTrait {
35 if let Some(signer) = &self.signer {
36 signer
37 } else {
38 tracing::error!("signer is not set for OkxExecutor");
39 panic!("signer is not set for OkxExecutor");
40 }
41 }
42
43 fn get_ratelimiter(&self) -> Arc<dyn SharedRatelimiterTrait<OkxRateLimitKey> + Sync + Send> {
44 self.ratelimiter.clone()
45 }
46}
47
48impl OkxExecutor {
49 pub fn get_endpoint(&self) -> &OkxEndpoint {
50 &self.endpoint
51 }
52
53 pub async fn subscribe_stream<T: DeserializeOwned + Send + 'static>(
54 &self,
55 url: &str,
56 arg: serde_json::Value,
57 ) -> Result<mpsc::Receiver<Result<T, SharedWsError>>, SharedWsError> {
58 let client = OkxWsStream::connect(ezsockets::ClientConfig::new(url)).await;
59
60 let (raw_tx, mut raw_rx) = mpsc::channel(128);
61
62 let (oneshot_tx, oneshot_rx) = oneshot::channel();
63
64 let message = OkxWsStreamCall::SubscribeStream {
65 args: vec![(arg, raw_tx)],
66 tx: oneshot_tx,
67 };
68
69 client
70 .call(message)
71 .inspect_err(|err| tracing::error!(?err, "failed to call ws stream"))
72 .map_err(|err| SharedWsError::AppError(err.to_string()))?;
73
74 self.streams.lock().await.push(client);
75
76 oneshot_rx
77 .await
78 .map_err(|err| SharedWsError::ChannelClosedError(err.to_string()))??;
79
80 let (tx, rx) = mpsc::channel(128);
81
82 tokio::spawn(async move {
83 while let Some(result) = raw_rx.recv().await {
84 let msg = match result {
85 Ok(resp) => match serde_json::from_value::<T>(resp.data) {
86 Ok(data) => Ok(data),
87 Err(err) => {
88 tracing::error!(?err, "failed to parse message");
89 Err(SharedWsError::SerdeError(err.to_string()))
90 }
91 },
92 Err(err) => Err(err),
93 };
94
95 if let Err(err) = tx.send(msg).await {
96 tracing::error!(?err, "failed to send message");
97 }
98 }
99 });
100
101 Ok(rx)
102 }
103}