Skip to main content

xapi_okx/common/
executor.rs

1use 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}