Skip to main content

strata_sdk/
account_stream.rs

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
16/// One externally authenticated private account stream. The SDK retains no
17/// signer and fails closed on any identity or sequence mismatch.
18pub 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    /// The first call returns the authenticated snapshot received during
173    /// connection. Later calls return sequenced order, fill, heartbeat, or
174    /// recovery-snapshot events.
175    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}