1use bytes::Bytes;
4use http::Method;
5use http::header::{CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue};
6use serde::Serialize;
7use serde::de::DeserializeOwned;
8
9use crate::OkxRegion;
10use crate::api::account::Account;
11use crate::api::convert::Convert;
12use crate::api::finance::Finance;
13use crate::api::funding::Funding;
14use crate::api::market::Market;
15use crate::api::public_data::PublicData;
16use crate::api::sub_account::SubAccount;
17use crate::api::trade::Trade;
18use crate::credentials::Credentials;
19use crate::error::{Error, RestError};
20use crate::model::OkxResponse;
21use crate::signing;
22use crate::transport::{DefaultTransport, Transport};
23
24pub struct OkxClient<T = DefaultTransport> {
36 transport: T,
37 credentials: Option<Credentials>,
38 base_url: String,
39 demo: bool,
40}
41
42#[cfg(feature = "reqwest")]
43impl OkxClient {
44 pub fn builder() -> OkxClientBuilder<DefaultTransport> {
47 OkxClientBuilder::from_transport(crate::transport::ReqwestTransport::new())
48 }
49}
50
51impl<T> OkxClient<T> {
52 pub fn with_transport(transport: T) -> OkxClientBuilder<T> {
54 OkxClientBuilder::from_transport(transport)
55 }
56}
57
58impl<T: Transport> OkxClient<T> {
59 pub fn market(&self) -> Market<'_, T> {
61 Market::new(self)
62 }
63
64 pub fn public_data(&self) -> PublicData<'_, T> {
66 PublicData::new(self)
67 }
68
69 pub fn account(&self) -> Account<'_, T> {
71 Account::new(self)
72 }
73
74 pub fn funding(&self) -> Funding<'_, T> {
76 Funding::new(self)
77 }
78
79 pub fn convert(&self) -> Convert<'_, T> {
81 Convert::new(self)
82 }
83
84 pub fn finance(&self) -> Finance<'_, T> {
86 Finance::new(self)
87 }
88
89 pub fn sub_account(&self) -> SubAccount<'_, T> {
91 SubAccount::new(self)
92 }
93
94 pub fn trade(&self) -> Trade<'_, T> {
96 Trade::new(self)
97 }
98
99 pub(crate) async fn get<Q, D>(
102 &self,
103 endpoint: &'static str,
104 query: &Q,
105 authenticated: bool,
106 ) -> Result<D, Error>
107 where
108 Q: Serialize,
109 D: DeserializeOwned,
110 {
111 let qs = serde_urlencoded::to_string(query)
112 .map_err(|e| RestError::Encode { source: e.into() })?;
113 let request_path = if qs.is_empty() {
114 endpoint.to_owned()
115 } else {
116 format!("{endpoint}?{qs}")
117 };
118 self.send(
119 endpoint,
120 Method::GET,
121 &request_path,
122 Bytes::new(),
123 authenticated,
124 )
125 .await
126 }
127
128 pub(crate) async fn post<B, D>(
131 &self,
132 endpoint: &'static str,
133 body: &B,
134 authenticated: bool,
135 ) -> Result<D, Error>
136 where
137 B: Serialize,
138 D: DeserializeOwned,
139 {
140 let body = serde_json::to_vec(body).map_err(|e| RestError::Encode { source: e.into() })?;
141 self.send(
142 endpoint,
143 Method::POST,
144 endpoint,
145 Bytes::from(body),
146 authenticated,
147 )
148 .await
149 }
150
151 async fn send<D>(
152 &self,
153 endpoint: &'static str,
154 method: Method,
155 request_path: &str,
156 body: Bytes,
157 authenticated: bool,
158 ) -> Result<D, Error>
159 where
160 D: DeserializeOwned,
161 {
162 let url = format!("{}{}", self.base_url, request_path);
163 let mut builder = http::Request::builder().method(method.clone()).uri(url);
164
165 let headers = builder
166 .headers_mut()
167 .expect("a freshly constructed request builder has no error");
168 headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
169 if self.demo {
170 headers.insert(
171 HeaderName::from_static("x-simulated-trading"),
172 HeaderValue::from_static("1"),
173 );
174 }
175 if authenticated {
176 let credentials = self.credentials.as_ref().ok_or_else(|| {
177 RestError::Configuration("authenticated endpoint requires credentials".to_owned())
178 })?;
179 let timestamp = signing::timestamp();
180 let body_str = std::str::from_utf8(&body).unwrap_or_default();
181 let prehash = signing::pre_hash(×tamp, method.as_str(), request_path, body_str);
182 let signature = signing::sign(&prehash, credentials.secret_key());
183 insert_header(headers, "ok-access-key", credentials.api_key())?;
184 insert_header(headers, "ok-access-sign", &signature)?;
185 insert_header(headers, "ok-access-timestamp", ×tamp)?;
186 insert_header(headers, "ok-access-passphrase", credentials.passphrase())?;
187 }
188
189 let request = builder
190 .body(body)
191 .map_err(|e| RestError::Encode { source: e.into() })?;
192 let response = self
193 .transport
194 .send(request)
195 .await
196 .map_err(RestError::from)?;
197 let status = response.status();
198 let bytes = response.into_body();
199 if !status.is_success() {
200 return Err(RestError::HttpStatus {
201 endpoint,
202 status,
203 body: String::from_utf8_lossy(&bytes).into_owned(),
204 }
205 .into());
206 }
207 let envelope: OkxResponse<D> =
208 serde_json::from_slice(&bytes).map_err(|e| RestError::Decode {
209 endpoint,
210 source: e,
211 })?;
212 if envelope.code != "0" {
213 return Err(RestError::Okx {
214 endpoint,
215 code: envelope.code,
216 message: envelope.msg,
217 }
218 .into());
219 }
220 Ok(envelope.data)
221 }
222}
223
224fn insert_header(headers: &mut HeaderMap, name: &'static str, value: &str) -> Result<(), Error> {
225 let value = HeaderValue::from_str(value)
226 .map_err(|e| RestError::Configuration(format!("invalid header value for {name}: {e}")))?;
227 headers.insert(HeaderName::from_static(name), value);
228 Ok(())
229}
230
231pub struct OkxClientBuilder<T = DefaultTransport> {
235 transport: T,
236 credentials: Option<Credentials>,
237 base_url: String,
238 demo: bool,
239}
240
241impl<T> OkxClientBuilder<T> {
242 fn from_transport(transport: T) -> Self {
243 Self {
244 transport,
245 credentials: None,
246 base_url: crate::API_URL.to_owned(),
247 demo: false,
248 }
249 }
250
251 pub fn credentials(mut self, credentials: Credentials) -> Self {
253 self.credentials = Some(credentials);
254 self
255 }
256
257 pub fn region(mut self, region: OkxRegion) -> Self {
263 self.base_url = region.api_url().to_owned();
264 self
265 }
266
267 pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
272 self.base_url = base_url.into();
273 self
274 }
275
276 pub fn demo_trading(mut self, demo: bool) -> Self {
278 self.demo = demo;
279 self
280 }
281
282 pub fn transport<U>(self, transport: U) -> OkxClientBuilder<U> {
284 OkxClientBuilder {
285 transport,
286 credentials: self.credentials,
287 base_url: self.base_url,
288 demo: self.demo,
289 }
290 }
291
292 pub fn build(self) -> OkxClient<T> {
294 OkxClient {
295 transport: self.transport,
296 credentials: self.credentials,
297 base_url: self.base_url,
298 demo: self.demo,
299 }
300 }
301}
302
303#[cfg(test)]
304mod tests {
305 use crate::error::{Error, RestError};
306 use crate::test_util::MockTransport;
307 use crate::{OkxClient, OkxRegion};
308
309 #[tokio::test]
313 async fn non_zero_code_is_api_error() {
314 let mock = MockTransport::new(r#"{"code":"51000","msg":"Parameter error","data":[]}"#);
315 let client = OkxClient::with_transport(mock).build();
316 let err = client.market().get_ticker("BAD").await.unwrap_err();
317 match err {
318 Error::Rest(RestError::Okx { code, message, .. }) => {
319 assert_eq!(code, "51000");
320 assert_eq!(message, "Parameter error");
321 }
322 other => panic!("expected Error::Rest(RestError::Okx), got {other:?}"),
323 }
324 }
325
326 #[test]
327 fn okx_region_returns_expected_api_urls() {
328 assert_eq!(OkxRegion::Global.api_url(), "https://www.okx.com");
329 assert_eq!(OkxRegion::Us.api_url(), "https://us.okx.com");
330 assert_eq!(OkxRegion::Eea.api_url(), "https://eea.okx.com");
331 }
332
333 #[tokio::test]
334 async fn region_sets_request_base_url() {
335 let mock = MockTransport::new(r#"{"code":"0","msg":"","data":[]}"#);
336 let client = OkxClient::with_transport(mock.clone())
337 .region(OkxRegion::Us)
338 .build();
339
340 client.market().get_ticker("BTC-USDT").await.unwrap();
341
342 let req = mock.captured();
343 assert!(
344 req.uri
345 .starts_with("https://us.okx.com/api/v5/market/ticker"),
346 "unexpected URI: {}",
347 req.uri
348 );
349 }
350
351 #[tokio::test]
352 async fn base_url_overrides_selected_region() {
353 let mock = MockTransport::new(r#"{"code":"0","msg":"","data":[]}"#);
354 let client = OkxClient::with_transport(mock.clone())
355 .region(OkxRegion::Eea)
356 .base_url("https://example.test")
357 .build();
358
359 client.market().get_ticker("BTC-USDT").await.unwrap();
360
361 let req = mock.captured();
362 assert!(
363 req.uri
364 .starts_with("https://example.test/api/v5/market/ticker"),
365 "unexpected URI: {}",
366 req.uri
367 );
368 }
369}