33
44use axum:: http:: StatusCode ;
55use serde:: Deserialize ;
6- use std:: { sync:: Arc , time:: Duration } ;
6+ use std:: {
7+ collections:: HashMap ,
8+ sync:: { Arc , Mutex , OnceLock } ,
9+ time:: { Duration , SystemTime , UNIX_EPOCH } ,
10+ } ;
711
812pub ( crate ) const IDENTITY_ME_URL : & str = "https://auth.openbitfun.com/api/v1/me" ;
913
@@ -27,6 +31,63 @@ struct IdentityResponse {
2731 user : VerifiedIdentity ,
2832}
2933
34+ /// A completed poll response kept long enough to make the authority's
35+ /// single-use transaction safe to retry after a lost HTTP response.
36+ struct CompletedPoll {
37+ payload : serde_json:: Value ,
38+ expires_at : i64 ,
39+ }
40+
41+ /// Terminal sign-in outcomes are keyed by both transaction id and secret.
42+ /// A different secret must never overwrite the valid client's replay entry.
43+ fn completed_authorizations ( ) -> & ' static Mutex < HashMap < ( String , String ) , CompletedPoll > > {
44+ static COMPLETED : OnceLock < Mutex < HashMap < ( String , String ) , CompletedPoll > > > = OnceLock :: new ( ) ;
45+ COMPLETED . get_or_init ( || Mutex :: new ( HashMap :: new ( ) ) )
46+ }
47+
48+ const POLL_REPLAY_SECS : i64 = 900 ;
49+
50+ fn now_secs ( ) -> i64 {
51+ SystemTime :: now ( )
52+ . duration_since ( UNIX_EPOCH )
53+ . unwrap_or_default ( )
54+ . as_secs ( ) as i64
55+ }
56+
57+ fn replay_completed_poll (
58+ transaction_id : & str ,
59+ transaction_secret : & str ,
60+ ) -> Option < serde_json:: Value > {
61+ let key = ( transaction_id. to_string ( ) , transaction_secret. to_string ( ) ) ;
62+ let mut completed = completed_authorizations ( )
63+ . lock ( )
64+ . unwrap_or_else ( |poisoned| poisoned. into_inner ( ) ) ;
65+ completed. retain ( |_, poll| poll. expires_at > now_secs ( ) ) ;
66+ completed. get ( & key) . map ( |poll| poll. payload . clone ( ) )
67+ }
68+
69+ fn remember_completed_poll (
70+ transaction_id : & str ,
71+ transaction_secret : & str ,
72+ payload : & serde_json:: Value ,
73+ ) {
74+ if payload. get ( "status" ) . and_then ( |status| status. as_str ( ) ) == Some ( "pending" ) {
75+ return ;
76+ }
77+ let key = ( transaction_id. to_string ( ) , transaction_secret. to_string ( ) ) ;
78+ let mut completed = completed_authorizations ( )
79+ . lock ( )
80+ . unwrap_or_else ( |poisoned| poisoned. into_inner ( ) ) ;
81+ completed. retain ( |_, poll| poll. expires_at > now_secs ( ) ) ;
82+ completed. insert (
83+ key,
84+ CompletedPoll {
85+ payload : payload. clone ( ) ,
86+ expires_at : now_secs ( ) + POLL_REPLAY_SECS ,
87+ } ,
88+ ) ;
89+ }
90+
3091impl IdentityVerifier {
3192 pub ( crate ) async fn start_auth (
3293 & self ,
@@ -55,14 +116,20 @@ impl IdentityVerifier {
55116 {
56117 return Err ( StatusCode :: BAD_REQUEST ) ;
57118 }
58- self . auth_request (
59- "auth/desktop/poll" ,
60- serde_json:: json!( {
61- "transactionId" : transaction_id,
62- "transactionSecret" : transaction_secret,
63- } ) ,
64- )
65- . await
119+ if let Some ( payload) = replay_completed_poll ( transaction_id, transaction_secret) {
120+ return Ok ( payload) ;
121+ }
122+ let payload = self
123+ . auth_request (
124+ "auth/desktop/poll" ,
125+ serde_json:: json!( {
126+ "transactionId" : transaction_id,
127+ "transactionSecret" : transaction_secret,
128+ } ) ,
129+ )
130+ . await ?;
131+ remember_completed_poll ( transaction_id, transaction_secret, & payload) ;
132+ Ok ( payload)
66133 }
67134
68135 async fn auth_request (
@@ -279,6 +346,50 @@ mod tests {
279346 task. abort ( ) ;
280347 }
281348 }
349+
350+ #[ tokio:: test]
351+ async fn poll_replays_a_completed_transaction_after_a_mismatched_secret ( ) {
352+ use std:: sync:: atomic:: { AtomicUsize , Ordering } ;
353+
354+ let upstream_calls = Arc :: new ( AtomicUsize :: new ( 0 ) ) ;
355+ let listener = tokio:: net:: TcpListener :: bind ( "127.0.0.1:0" ) . await . unwrap ( ) ;
356+ let transaction_id = format ! ( "replayed-{}" , listener. local_addr( ) . unwrap( ) . port( ) ) ;
357+ let url = format ! ( "http://{}/me" , listener. local_addr( ) . unwrap( ) ) ;
358+ let counter = upstream_calls. clone ( ) ;
359+ let app = Router :: new ( ) . route (
360+ "/auth/desktop/poll" ,
361+ axum:: routing:: post ( move || {
362+ let counter = counter. clone ( ) ;
363+ async move {
364+ if counter. fetch_add ( 1 , Ordering :: SeqCst ) == 0 {
365+ axum:: Json ( serde_json:: json!( {
366+ "status" : "authorized" ,
367+ "tokens" : { "accessToken" : "token-1" } ,
368+ } ) )
369+ } else {
370+ axum:: Json ( serde_json:: json!( { "status" : "consumed" } ) )
371+ }
372+ }
373+ } ) ,
374+ ) ;
375+ let task = tokio:: spawn ( async move {
376+ axum:: serve ( listener, app) . await . unwrap ( ) ;
377+ } ) ;
378+ let client = IdentityVerifier :: with_url ( & url) . unwrap ( ) ;
379+
380+ let first = client. poll_auth ( & transaction_id, "secret-1" ) . await . unwrap ( ) ;
381+ let repeated = client. poll_auth ( & transaction_id, "secret-1" ) . await . unwrap ( ) ;
382+ let other_secret = client. poll_auth ( & transaction_id, "secret-2" ) . await . unwrap ( ) ;
383+ let repeated_after_other_secret =
384+ client. poll_auth ( & transaction_id, "secret-1" ) . await . unwrap ( ) ;
385+
386+ assert_eq ! ( first[ "status" ] , "authorized" ) ;
387+ assert_eq ! ( repeated, first) ;
388+ assert_eq ! ( other_secret[ "status" ] , "consumed" ) ;
389+ assert_eq ! ( repeated_after_other_secret, first) ;
390+ assert_eq ! ( upstream_calls. load( Ordering :: SeqCst ) , 2 ) ;
391+ task. abort ( ) ;
392+ }
282393}
283394
284395fn identity_request_permit ( ) -> Result < tokio:: sync:: OwnedSemaphorePermit , StatusCode > {
0 commit comments