1use std::time::Duration;
2
3use futures_util::{SinkExt, StreamExt};
4use tokio::net::TcpStream;
5use tokio_tungstenite::tungstenite::Message;
6use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
7
8use super::market_stream::platform_websocket_url;
9use super::*;
10
11pub const ACCOUNT_STREAM_AUTH_DOMAIN: &str = "strata:account-stream:v2";
12const AUTH_TIMEOUT: Duration = Duration::from_secs(10);
13
14type PlatformSocket = WebSocketStream<MaybeTlsStream<TcpStream>>;
15
16pub struct AccountStream {
19 market_id: String,
20 wallet_address: String,
21 stream_id: String,
22 sequence: u64,
23 initial_snapshot: Option<PlatformAccountEvent>,
24 socket: PlatformSocket,
25}
26
27impl std::fmt::Debug for AccountStream {
28 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
29 formatter
30 .debug_struct("AccountStream")
31 .field("market_id", &self.market_id)
32 .field("wallet_address", &self.wallet_address)
33 .field("stream_id", &self.stream_id)
34 .field("sequence", &self.sequence)
35 .finish_non_exhaustive()
36 }
37}
38
39impl AccountStream {
40 pub(crate) async fn connect<S: AccountSigner + ?Sized>(
41 client: &StrataClient,
42 market_id: &str,
43 signer: &S,
44 ) -> Result<Self, SdkError> {
45 let market_id = validate_platform_market_id(market_id)?;
46 let wallet_address =
47 canonical_public_key(signer.public_key(), "account signer public key")?;
48 let url = platform_websocket_url(
49 &client.base_url,
50 &format!("v2/markets/{market_id}/account/{wallet_address}/stream"),
51 )?;
52 let (mut socket, _) = tokio_tungstenite::connect_async(url.as_str())
53 .await
54 .map_err(|error| SdkError::Stream(error.to_string()))?;
55
56 let challenge_frame = next_text(&mut socket, "account authentication challenge").await?;
57 let challenge_event: PlatformAccountEvent = serde_json::from_str(&challenge_frame)
58 .map_err(|error| SdkError::InvalidResponse(error.to_string()))?;
59 let challenge = match challenge_event {
60 PlatformAccountEvent::AuthChallenge {
61 schema_version,
62 contract_version,
63 market_id: response_market,
64 wallet_address: response_wallet,
65 challenge,
66 server_time_ms,
67 expires_at_ms,
68 } => {
69 validate_platform_version(schema_version, &contract_version)?;
70 if response_market != market_id
71 || response_wallet != wallet_address
72 || expires_at_ms <= server_time_ms
73 || challenge.len() != 64
74 || !challenge
75 .bytes()
76 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
77 {
78 return Err(SdkError::InvalidResponse(
79 "account stream authentication bindings are invalid".to_owned(),
80 ));
81 }
82 challenge
83 }
84 _ => {
85 return Err(SdkError::InvalidResponse(
86 "account stream did not begin with authentication".to_owned(),
87 ))
88 }
89 };
90
91 let message = account_stream_auth_message(&market_id, &wallet_address, &challenge)?;
92 let signature = signer
93 .sign_message(&message)
94 .await
95 .map_err(SdkError::Signer)?;
96 if signature.len() != 64 {
97 return Err(SdkError::Signer(
98 "account signer must return a 64-byte Ed25519 signature".to_owned(),
99 ));
100 }
101 socket
102 .send(Message::Text(
103 serde_json::json!({
104 "type": "authenticate",
105 "signature": hex::encode(signature),
106 })
107 .to_string()
108 .into(),
109 ))
110 .await
111 .map_err(|error| SdkError::Stream(error.to_string()))?;
112
113 let snapshot_frame = next_text(&mut socket, "signed account snapshot").await?;
114 let snapshot: PlatformAccountEvent = serde_json::from_str(&snapshot_frame)
115 .map_err(|error| SdkError::InvalidResponse(error.to_string()))?;
116 let (stream_id, sequence) = match &snapshot {
117 PlatformAccountEvent::AccountSnapshot {
118 schema_version,
119 contract_version,
120 market_id: response_market,
121 wallet_address: response_wallet,
122 stream_id,
123 sequence,
124 orders,
125 fills,
126 ..
127 } => {
128 validate_account_identity(
129 schema_version,
130 contract_version,
131 response_market,
132 response_wallet,
133 &market_id,
134 &wallet_address,
135 )?;
136 if !valid_handle(stream_id, "account_stream_") {
137 return Err(SdkError::InvalidResponse(
138 "account stream identity is invalid".to_owned(),
139 ));
140 }
141 validate_account_state(orders, fills)?;
142 (
143 stream_id.clone(),
144 validate_response_atoms(sequence, "sequence", false)?,
145 )
146 }
147 _ => {
148 return Err(SdkError::InvalidResponse(
149 "account authentication did not return a signed snapshot".to_owned(),
150 ))
151 }
152 };
153
154 Ok(Self {
155 market_id,
156 wallet_address,
157 stream_id,
158 sequence,
159 initial_snapshot: Some(snapshot),
160 socket,
161 })
162 }
163
164 pub fn market_id(&self) -> &str {
165 &self.market_id
166 }
167
168 pub fn wallet_address(&self) -> &str {
169 &self.wallet_address
170 }
171
172 pub async fn next_event(&mut self) -> Result<Option<PlatformAccountEvent>, SdkError> {
176 if let Some(snapshot) = self.initial_snapshot.take() {
177 return Ok(Some(snapshot));
178 }
179 loop {
180 let Some(frame) = self.socket.next().await else {
181 return Ok(None);
182 };
183 let frame = frame.map_err(|error| SdkError::Stream(error.to_string()))?;
184 match frame {
185 Message::Text(text) => {
186 let event: PlatformAccountEvent = serde_json::from_str(&text)
187 .map_err(|error| SdkError::InvalidResponse(error.to_string()))?;
188 if let Err(error) = self.validate_event(&event) {
189 let _ = self.socket.close(None).await;
190 return Err(error);
191 }
192 return Ok(Some(event));
193 }
194 Message::Ping(payload) => {
195 self.socket
196 .send(Message::Pong(payload))
197 .await
198 .map_err(|error| SdkError::Stream(error.to_string()))?;
199 }
200 Message::Pong(_) => {}
201 Message::Close(_) => return Ok(None),
202 _ => {
203 let _ = self.socket.close(None).await;
204 return Err(SdkError::InvalidResponse(
205 "account stream sent a non-text data frame".to_owned(),
206 ));
207 }
208 }
209 }
210 }
211
212 pub async fn close(&mut self) -> Result<(), SdkError> {
213 self.socket
214 .close(None)
215 .await
216 .map_err(|error| SdkError::Stream(error.to_string()))
217 }
218
219 fn validate_event(&mut self, event: &PlatformAccountEvent) -> Result<(), SdkError> {
220 match event {
221 PlatformAccountEvent::AuthChallenge { .. } => {
222 return Err(SdkError::InvalidResponse(
223 "account stream challenged after state delivery".to_owned(),
224 ))
225 }
226 PlatformAccountEvent::AccountSnapshot {
227 schema_version,
228 contract_version,
229 market_id,
230 wallet_address,
231 stream_id,
232 sequence,
233 orders,
234 fills,
235 ..
236 } => {
237 self.validate_identity(
238 *schema_version,
239 contract_version,
240 market_id,
241 wallet_address,
242 )?;
243 let next = validate_response_atoms(sequence, "sequence", false)?;
244 if stream_id != &self.stream_id || next <= self.sequence {
245 return Err(SdkError::InvalidResponse(
246 "account recovery snapshot did not advance its sequence".to_owned(),
247 ));
248 }
249 validate_account_state(orders, fills)?;
250 self.sequence = next;
251 }
252 PlatformAccountEvent::OrdersSnapshot {
253 schema_version,
254 contract_version,
255 market_id,
256 wallet_address,
257 stream_id,
258 sequence,
259 previous_sequence,
260 orders,
261 ..
262 } => {
263 self.validate_identity(
264 *schema_version,
265 contract_version,
266 market_id,
267 wallet_address,
268 )?;
269 self.validate_sequence(stream_id, sequence, previous_sequence)?;
270 validate_account_state(orders, &[])?;
271 }
272 PlatformAccountEvent::Fill {
273 schema_version,
274 contract_version,
275 market_id,
276 wallet_address,
277 stream_id,
278 sequence,
279 previous_sequence,
280 fill,
281 ..
282 } => {
283 self.validate_identity(
284 *schema_version,
285 contract_version,
286 market_id,
287 wallet_address,
288 )?;
289 self.validate_sequence(stream_id, sequence, previous_sequence)?;
290 validate_account_state(&[], std::slice::from_ref(fill))?;
291 }
292 PlatformAccountEvent::Heartbeat {
293 schema_version,
294 contract_version,
295 market_id,
296 wallet_address,
297 stream_id,
298 sequence,
299 previous_sequence,
300 ..
301 } => {
302 self.validate_identity(
303 *schema_version,
304 contract_version,
305 market_id,
306 wallet_address,
307 )?;
308 self.validate_sequence(stream_id, sequence, previous_sequence)?;
309 }
310 }
311 Ok(())
312 }
313
314 fn validate_identity(
315 &self,
316 schema_version: u16,
317 contract_version: &str,
318 market_id: &str,
319 wallet_address: &str,
320 ) -> Result<(), SdkError> {
321 validate_account_identity(
322 &schema_version,
323 contract_version,
324 market_id,
325 wallet_address,
326 &self.market_id,
327 &self.wallet_address,
328 )
329 }
330
331 fn validate_sequence(
332 &mut self,
333 stream_id: &str,
334 sequence: &str,
335 previous_sequence: &str,
336 ) -> Result<(), SdkError> {
337 let next = validate_response_atoms(sequence, "sequence", false)?;
338 let previous = validate_response_atoms(previous_sequence, "previous_sequence", false)?;
339 if stream_id != self.stream_id
340 || previous != self.sequence
341 || next != previous.saturating_add(1)
342 {
343 return Err(SdkError::InvalidResponse(
344 "account stream sequence gap detected".to_owned(),
345 ));
346 }
347 self.sequence = next;
348 Ok(())
349 }
350}
351
352async fn next_text(socket: &mut PlatformSocket, expected: &str) -> Result<String, SdkError> {
353 let frame = tokio::time::timeout(AUTH_TIMEOUT, socket.next())
354 .await
355 .map_err(|_| SdkError::Stream(format!("{expected} timed out")))?
356 .ok_or_else(|| SdkError::Stream(format!("socket closed before {expected}")))?
357 .map_err(|error| SdkError::Stream(error.to_string()))?;
358 let Message::Text(text) = frame else {
359 return Err(SdkError::InvalidResponse(format!(
360 "expected a text {expected}"
361 )));
362 };
363 Ok(text.to_string())
364}
365
366fn validate_account_identity(
367 schema_version: &u16,
368 contract_version: &str,
369 actual_market: &str,
370 actual_wallet: &str,
371 expected_market: &str,
372 expected_wallet: &str,
373) -> Result<(), SdkError> {
374 validate_platform_market_response(
375 *schema_version,
376 contract_version,
377 actual_market,
378 expected_market,
379 )?;
380 if actual_wallet != expected_wallet {
381 return Err(SdkError::InvalidResponse(
382 "account stream wallet does not match signed request".to_owned(),
383 ));
384 }
385 Ok(())
386}
387
388pub(crate) fn validate_account_state(
389 orders: &[PlatformAccountOrder],
390 fills: &[PlatformAccountFill],
391) -> Result<(), SdkError> {
392 let mut order_ids = HashSet::new();
393 if orders.iter().any(|order| {
394 let original =
395 validate_response_atoms(&order.original_size_atoms, "original_size_atoms", false);
396 let remaining =
397 validate_response_atoms(&order.remaining_size_atoms, "remaining_size_atoms", false);
398 let state_matches = match (original, remaining) {
399 (Ok(original), Ok(remaining)) if remaining <= original => {
400 order.state
401 == if remaining == original {
402 PlatformOrderState::Open
403 } else {
404 PlatformOrderState::PartiallyFilled
405 }
406 }
407 _ => false,
408 };
409 !valid_handle(&order.order_id, "order_")
410 || !order_ids.insert(order.order_id.as_str())
411 || validate_response_atoms(&order.limit_price_atoms, "limit_price_atoms", false)
412 .is_err()
413 || !state_matches
414 }) {
415 return Err(SdkError::InvalidResponse(
416 "account stream contains an invalid order".to_owned(),
417 ));
418 }
419 let mut fill_ids = HashSet::new();
420 if fills.iter().any(|fill| {
421 let valid_confirmation = fill
422 .confirmed_at_ms
423 .is_none_or(|confirmed| confirmed >= fill.executed_at_ms)
424 && (fill.settlement != PlatformSettlementState::Confirmed
425 || fill.confirmed_at_ms.is_some());
426 let valid_transaction = fill.transaction_id.as_deref().is_none_or(|transaction| {
427 (32..=100).contains(&transaction.len())
428 && transaction.bytes().all(|byte| {
429 byte.is_ascii_alphanumeric() && !matches!(byte, b'0' | b'O' | b'I' | b'l')
430 })
431 });
432 !valid_handle(&fill.fill_id, "fill_")
433 || !fill_ids.insert(fill.fill_id.as_str())
434 || validate_response_atoms(&fill.price_atoms, "price_atoms", false).is_err()
435 || validate_response_atoms(&fill.size_atoms, "size_atoms", false).is_err()
436 || validate_response_atoms(&fill.fee_quote_atoms, "fee_quote_atoms", true).is_err()
437 || validate_signed_response_atoms(
438 &fill.realized_pnl_quote_atoms,
439 "realized_pnl_quote_atoms",
440 )
441 .is_err()
442 || !valid_confirmation
443 || !valid_transaction
444 }) {
445 return Err(SdkError::InvalidResponse(
446 "account stream contains an invalid fill".to_owned(),
447 ));
448 }
449 Ok(())
450}
451
452fn validate_signed_response_atoms(value: &str, field: &str) -> Result<(), SdkError> {
453 let digits = value.strip_prefix('-').unwrap_or(value);
454 if digits.is_empty()
455 || (digits.len() > 1 && digits.starts_with('0'))
456 || (value.starts_with('-') && digits == "0")
457 || !digits.bytes().all(|byte| byte.is_ascii_digit())
458 || digits.parse::<u64>().is_err()
459 {
460 return Err(SdkError::InvalidResponse(format!(
461 "{field} must be a canonical signed atomic decimal string"
462 )));
463 }
464 Ok(())
465}
466
467pub fn account_stream_auth_message(
468 market_id: &str,
469 wallet_address: &str,
470 challenge: &str,
471) -> Result<Vec<u8>, SdkError> {
472 let market_id = validate_platform_market_id(market_id)?;
473 let wallet_address = canonical_public_key(wallet_address, "wallet_address")?;
474 if challenge.len() != 64
475 || !challenge
476 .bytes()
477 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
478 {
479 return Err(SdkError::InvalidRequest(
480 "account stream challenge must be 32-byte lowercase hexadecimal".to_owned(),
481 ));
482 }
483 Ok(
484 format!("{ACCOUNT_STREAM_AUTH_DOMAIN}\n{market_id}\n{wallet_address}\n{challenge}")
485 .into_bytes(),
486 )
487}