1use std::collections::HashMap;
19
20use nautilus_core::{string::secret::SecretString, time::get_atomic_clock_realtime};
21use nautilus_network::http::{
22 HttpClient, HttpRedirectPolicy, Method, create_standard_nautilus_headers,
23};
24use serde::Deserialize;
25use zeroize::{Zeroize, ZeroizeOnDrop};
26
27use crate::{
28 common::{credential::EvmPrivateKey, urls::clob_http_url},
29 http::error::{Error, Result, decode_response},
30 signing::eip712::sign_clob_auth,
31};
32
33#[derive(Debug, Deserialize, Zeroize, ZeroizeOnDrop)]
35#[serde(rename_all = "camelCase")]
36pub struct ApiCredentials {
37 pub api_key: SecretString,
38 pub secret: SecretString,
39 pub passphrase: SecretString,
40}
41
42pub async fn create_api_key(
48 private_key: &EvmPrivateKey,
49 nonce: u64,
50 base_url: Option<&str>,
51) -> Result<ApiCredentials> {
52 let (client, headers, base) = prepare_l1_request(private_key, nonce, base_url)?;
53 let url = format!("{base}/auth/api-key");
54 let response = client
55 .request(Method::POST, url, None, Some(headers), None, None, None)
56 .await
57 .map_err(Error::from_http_client)?;
58
59 decode_response(&response)
60}
61
62pub async fn derive_api_key(
68 private_key: &EvmPrivateKey,
69 nonce: u64,
70 base_url: Option<&str>,
71) -> Result<ApiCredentials> {
72 let (client, headers, base) = prepare_l1_request(private_key, nonce, base_url)?;
73 let url = format!("{base}/auth/derive-api-key");
74 let response = client
75 .request(Method::GET, url, None, Some(headers), None, None, None)
76 .await
77 .map_err(Error::from_http_client)?;
78
79 decode_response(&response)
80}
81
82pub async fn create_or_derive_api_key(
89 private_key: &EvmPrivateKey,
90 nonce: u64,
91 base_url: Option<&str>,
92) -> Result<ApiCredentials> {
93 match create_api_key(private_key, nonce, base_url).await {
94 Ok(creds) => Ok(creds),
95 Err(e) if e.is_http_status_error() => derive_api_key(private_key, nonce, base_url).await,
96 Err(e) => Err(e),
97 }
98}
99
100fn prepare_l1_request(
101 private_key: &EvmPrivateKey,
102 nonce: u64,
103 base_url: Option<&str>,
104) -> Result<(HttpClient, HashMap<String, String>, String)> {
105 let base = base_url
106 .unwrap_or_else(|| clob_http_url())
107 .trim_end_matches('/')
108 .to_string();
109 let timestamp =
110 (get_atomic_clock_realtime().get_time_ns().as_u64() / 1_000_000_000).to_string();
111 let (address, signature) = sign_clob_auth(private_key, ×tamp, nonce)?;
112 let headers = l1_headers(&address, &signature, ×tamp, nonce);
113 let client = HttpClient::builder()
114 .redirect_policy(HttpRedirectPolicy::Reject)
115 .headers(create_standard_nautilus_headers().into_iter().collect())
116 .build()
117 .map_err(Error::from_http_client)?;
118 Ok((client, headers, base))
119}
120
121fn l1_headers(
122 address: &str,
123 signature: &str,
124 timestamp: &str,
125 nonce: u64,
126) -> HashMap<String, String> {
127 HashMap::from([
128 ("POLY_ADDRESS".to_string(), address.to_string()),
129 ("POLY_SIGNATURE".to_string(), signature.to_string()),
130 ("POLY_TIMESTAMP".to_string(), timestamp.to_string()),
131 ("POLY_NONCE".to_string(), nonce.to_string()),
132 ])
133}
134
135#[cfg(test)]
136mod tests {
137 use std::sync::Arc;
138
139 use axum::{
140 Router,
141 extract::State,
142 http::{HeaderMap, Method as AxumMethod, StatusCode, Uri},
143 response::{IntoResponse, Response},
144 routing::{get, post},
145 };
146 use rstest::rstest;
147 use tokio::{
148 io::{AsyncReadExt, AsyncWriteExt},
149 net::TcpStream,
150 sync::Mutex,
151 task::JoinHandle,
152 };
153
154 use super::*;
155
156 const TEST_PRIVATE_KEY: &str =
157 "0xac0974bec39a17e36ba4a6b4d238ff944bacb478cbed5efcae784d7bf4f2ff80";
158 const TEST_ADDRESS: &str = "0xf39fd6e51aad88f6f4ce6ab8827279cfffb92266";
159 const API_CREDENTIALS_RESPONSE: &str =
160 include_str!("../../test_data/http_api_credentials.json");
161 const ERROR_RESPONSE: &str = include_str!("../../test_data/http_order_response_error_500.json");
162 const WRONG_SCHEMA_RESPONSE: &str = include_str!("../../test_data/http_empty_page.json");
163
164 type RequestLog = Arc<Mutex<Vec<RecordedRequest>>>;
165
166 #[derive(Clone, Copy)]
167 struct TestResponse {
168 status: StatusCode,
169 body: &'static str,
170 }
171
172 #[derive(Clone)]
173 struct TestServerState {
174 create_response: TestResponse,
175 derive_response: TestResponse,
176 requests: RequestLog,
177 }
178
179 #[derive(Debug)]
180 struct RecordedRequest {
181 method: AxumMethod,
182 path: String,
183 headers: HeaderMap,
184 }
185
186 #[rstest]
187 fn test_api_credentials_zeroize_and_redact_debug() {
188 fn assert_zeroize_on_drop<T: ZeroizeOnDrop>() {}
189
190 let mut credentials = ApiCredentials {
191 api_key: SecretString::from("api-key-sentinel"),
192 secret: SecretString::from("secret-sentinel"),
193 passphrase: SecretString::from("passphrase-sentinel"),
194 };
195
196 let debug = format!("{credentials:?}");
197 credentials.zeroize();
198
199 assert_zeroize_on_drop::<ApiCredentials>();
200 assert!(!debug.contains("api-key-sentinel"));
201 assert!(!debug.contains("secret-sentinel"));
202 assert!(!debug.contains("passphrase-sentinel"));
203 assert_eq!(credentials.api_key.expose_secret(), "");
204 assert_eq!(credentials.secret.expose_secret(), "");
205 assert_eq!(credentials.passphrase.expose_secret(), "");
206 }
207
208 #[rstest]
209 #[tokio::test]
210 async fn test_create_or_derive_api_key_returns_created_credentials_without_fallback() {
211 let (base_url, requests, server) = start_test_server(
212 TestResponse {
213 status: StatusCode::OK,
214 body: API_CREDENTIALS_RESPONSE,
215 },
216 TestResponse {
217 status: StatusCode::INTERNAL_SERVER_ERROR,
218 body: ERROR_RESPONSE,
219 },
220 )
221 .await;
222 let private_key = EvmPrivateKey::new(TEST_PRIVATE_KEY).unwrap();
223 let timestamp_before = unix_seconds();
224
225 let credentials = create_or_derive_api_key(&private_key, 7, Some(&base_url))
226 .await
227 .unwrap();
228 let timestamp_after = unix_seconds();
229
230 let requests = requests.lock().await;
231 assert_credentials(&credentials);
232 assert_eq!(requests.len(), 1);
233 assert_eq!(requests[0].method, AxumMethod::POST);
234 assert_eq!(requests[0].path, "/auth/api-key");
235 assert_l1_headers(
236 &requests[0].headers,
237 &private_key,
238 7,
239 timestamp_before,
240 timestamp_after,
241 );
242 server.abort();
243 }
244
245 #[rstest]
246 #[tokio::test]
247 async fn test_derive_api_key_sends_l1_get_and_decodes_credentials() {
248 let (base_url, requests, server) = start_test_server(
249 TestResponse {
250 status: StatusCode::INTERNAL_SERVER_ERROR,
251 body: ERROR_RESPONSE,
252 },
253 TestResponse {
254 status: StatusCode::OK,
255 body: API_CREDENTIALS_RESPONSE,
256 },
257 )
258 .await;
259 let private_key = EvmPrivateKey::new(TEST_PRIVATE_KEY).unwrap();
260 let timestamp_before = unix_seconds();
261
262 let credentials = derive_api_key(&private_key, 11, Some(&base_url))
263 .await
264 .unwrap();
265 let timestamp_after = unix_seconds();
266
267 let requests = requests.lock().await;
268 assert_credentials(&credentials);
269 assert_eq!(requests.len(), 1);
270 assert_eq!(requests[0].method, AxumMethod::GET);
271 assert_eq!(requests[0].path, "/auth/derive-api-key");
272 assert_l1_headers(
273 &requests[0].headers,
274 &private_key,
275 11,
276 timestamp_before,
277 timestamp_after,
278 );
279 server.abort();
280 }
281
282 #[rstest]
283 #[case(StatusCode::CONFLICT)]
284 #[case(StatusCode::TOO_MANY_REQUESTS)]
285 #[tokio::test]
286 async fn test_create_or_derive_api_key_falls_back_after_http_error(
287 #[case] create_status: StatusCode,
288 ) {
289 let (base_url, requests, server) = start_test_server(
290 TestResponse {
291 status: create_status,
292 body: ERROR_RESPONSE,
293 },
294 TestResponse {
295 status: StatusCode::OK,
296 body: API_CREDENTIALS_RESPONSE,
297 },
298 )
299 .await;
300 let private_key = EvmPrivateKey::new(TEST_PRIVATE_KEY).unwrap();
301 let timestamp_before = unix_seconds();
302
303 let credentials = create_or_derive_api_key(&private_key, 13, Some(&base_url))
304 .await
305 .unwrap();
306 let timestamp_after = unix_seconds();
307
308 let requests = requests.lock().await;
309 assert_credentials(&credentials);
310 assert_eq!(requests.len(), 2);
311 assert_eq!(requests[0].method, AxumMethod::POST);
312 assert_eq!(requests[0].path, "/auth/api-key");
313 assert_eq!(requests[1].method, AxumMethod::GET);
314 assert_eq!(requests[1].path, "/auth/derive-api-key");
315 assert_l1_headers(
316 &requests[0].headers,
317 &private_key,
318 13,
319 timestamp_before,
320 timestamp_after,
321 );
322 assert_l1_headers(
323 &requests[1].headers,
324 &private_key,
325 13,
326 timestamp_before,
327 timestamp_after,
328 );
329 server.abort();
330 }
331
332 #[rstest]
333 #[tokio::test]
334 async fn test_create_or_derive_api_key_does_not_fall_back_after_decode_error() {
335 let (base_url, requests, server) = start_test_server(
336 TestResponse {
337 status: StatusCode::OK,
338 body: WRONG_SCHEMA_RESPONSE,
339 },
340 TestResponse {
341 status: StatusCode::OK,
342 body: API_CREDENTIALS_RESPONSE,
343 },
344 )
345 .await;
346 let private_key = EvmPrivateKey::new(TEST_PRIVATE_KEY).unwrap();
347 let timestamp_before = unix_seconds();
348
349 let error = create_or_derive_api_key(&private_key, 17, Some(&base_url))
350 .await
351 .unwrap_err();
352 let timestamp_after = unix_seconds();
353
354 let requests = requests.lock().await;
355 assert!(matches!(error, Error::Serde(_)));
356 assert_eq!(requests.len(), 1);
357 assert_eq!(requests[0].method, AxumMethod::POST);
358 assert_eq!(requests[0].path, "/auth/api-key");
359 assert_l1_headers(
360 &requests[0].headers,
361 &private_key,
362 17,
363 timestamp_before,
364 timestamp_after,
365 );
366 server.abort();
367 }
368
369 #[rstest]
370 #[tokio::test]
371 async fn test_create_or_derive_api_key_does_not_fall_back_after_transport_error() {
372 let (base_url, requests, server) = start_connection_drop_server().await;
373 let private_key = EvmPrivateKey::new(TEST_PRIVATE_KEY).unwrap();
374
375 let error = create_or_derive_api_key(&private_key, 19, Some(&base_url))
376 .await
377 .unwrap_err();
378
379 let requests = requests.lock().await;
380 assert!(matches!(error, Error::Transport(_)));
381 assert_eq!(requests.as_slice(), ["POST /auth/api-key HTTP/1.1"]);
382 server.abort();
383 }
384
385 async fn start_test_server(
386 create_response: TestResponse,
387 derive_response: TestResponse,
388 ) -> (String, RequestLog, JoinHandle<()>) {
389 let requests = Arc::new(Mutex::new(Vec::new()));
390 let state = TestServerState {
391 create_response,
392 derive_response,
393 requests: Arc::clone(&requests),
394 };
395 let app = Router::new()
396 .route("/auth/api-key", post(handle_create))
397 .route("/auth/derive-api-key", get(handle_derive))
398 .with_state(state);
399 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
400 let address = listener.local_addr().unwrap();
401 let server = tokio::spawn(async move {
402 axum::serve(listener, app).await.unwrap();
403 });
404
405 (format!("http://{address}/"), requests, server)
406 }
407
408 async fn start_connection_drop_server() -> (String, Arc<Mutex<Vec<String>>>, JoinHandle<()>) {
409 let requests = Arc::new(Mutex::new(Vec::new()));
410 let server_requests = Arc::clone(&requests);
411 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
412 let address = listener.local_addr().unwrap();
413 let server = tokio::spawn(async move {
414 let (mut create_stream, _) = listener.accept().await.unwrap();
415 let create_request = read_request_line(&mut create_stream).await;
416 server_requests.lock().await.push(create_request);
417 drop(create_stream);
418
419 let (mut derive_stream, _) = listener.accept().await.unwrap();
420 let derive_request = read_request_line(&mut derive_stream).await;
421 server_requests.lock().await.push(derive_request);
422 let response = format!(
423 "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
424 API_CREDENTIALS_RESPONSE.len(),
425 API_CREDENTIALS_RESPONSE,
426 );
427 derive_stream.write_all(response.as_bytes()).await.unwrap();
428 });
429
430 (format!("http://{address}/"), requests, server)
431 }
432
433 async fn read_request_line(stream: &mut TcpStream) -> String {
434 let mut request = Vec::new();
435
436 loop {
437 let mut chunk = [0; 1024];
438 let count = stream.read(&mut chunk).await.unwrap();
439 assert!(count > 0);
440 request.extend_from_slice(&chunk[..count]);
441 if request.windows(4).any(|window| window == b"\r\n\r\n") {
442 break;
443 }
444 }
445
446 String::from_utf8(request)
447 .unwrap()
448 .lines()
449 .next()
450 .unwrap()
451 .to_string()
452 }
453
454 async fn handle_create(
455 State(state): State<TestServerState>,
456 method: AxumMethod,
457 uri: Uri,
458 headers: HeaderMap,
459 ) -> Response {
460 let response = state.create_response;
461 record_and_respond(state, method, uri, headers, response).await
462 }
463
464 async fn handle_derive(
465 State(state): State<TestServerState>,
466 method: AxumMethod,
467 uri: Uri,
468 headers: HeaderMap,
469 ) -> Response {
470 let response = state.derive_response;
471 record_and_respond(state, method, uri, headers, response).await
472 }
473
474 async fn record_and_respond(
475 state: TestServerState,
476 method: AxumMethod,
477 uri: Uri,
478 headers: HeaderMap,
479 response: TestResponse,
480 ) -> Response {
481 state.requests.lock().await.push(RecordedRequest {
482 method,
483 path: uri.path().to_string(),
484 headers,
485 });
486 (
487 response.status,
488 [("content-type", "application/json")],
489 response.body,
490 )
491 .into_response()
492 }
493
494 fn assert_credentials(credentials: &ApiCredentials) {
495 assert_eq!(credentials.api_key.expose_secret(), "test-api-key");
496 assert_eq!(credentials.secret.expose_secret(), "test-secret");
497 assert_eq!(credentials.passphrase.expose_secret(), "test-passphrase");
498 }
499
500 fn assert_l1_headers(
501 headers: &HeaderMap,
502 private_key: &EvmPrivateKey,
503 nonce: u64,
504 timestamp_before: u64,
505 timestamp_after: u64,
506 ) {
507 let address = headers.get("poly_address").unwrap().to_str().unwrap();
508 let signature = headers.get("poly_signature").unwrap().to_str().unwrap();
509 let timestamp = headers
510 .get("poly_timestamp")
511 .unwrap()
512 .to_str()
513 .unwrap()
514 .parse::<u64>()
515 .unwrap();
516 let header_nonce = headers.get("poly_nonce").unwrap().to_str().unwrap();
517 let (_, expected_signature) =
518 sign_clob_auth(private_key, ×tamp.to_string(), nonce).unwrap();
519
520 assert_eq!(address, TEST_ADDRESS);
521 assert_eq!(signature, expected_signature);
522 assert!(timestamp >= timestamp_before);
523 assert!(timestamp <= timestamp_after);
524 assert_eq!(header_nonce, nonce.to_string());
525 }
526
527 fn unix_seconds() -> u64 {
528 get_atomic_clock_realtime().get_time_ns().as_u64() / 1_000_000_000
529 }
530}