@@ -8,8 +8,7 @@ use crate::{
88 alias:: { ReadStdbConnectedMessage , ReadStdbDisconnectedMessage } ,
99 channel_bridge:: { channel_sender, register_channel} ,
1010 message:: {
11- StdbConnectErrorMessage , StdbConnectedMessage , StdbDisconnectRequestedMessage ,
12- StdbDisconnectedMessage ,
11+ DisconnectIntent , StdbConnectErrorMessage , StdbConnectedMessage , StdbDisconnectedMessage ,
1312 } ,
1413 set:: StdbSet ,
1514} ;
@@ -23,12 +22,15 @@ use spacetimedb_sdk::{
2322 __codegen:: { DbConnection , SpacetimeModule } ,
2423 Compression , ConnectionId , DbConnectionBuilder , DbContext , Identity , Result ,
2524} ;
26- use std:: sync:: Arc ;
25+ use std:: sync:: {
26+ Arc ,
27+ atomic:: { AtomicBool , Ordering } ,
28+ } ;
2729
2830/// Stores the in-flight task for a pending connection attempt.
2931#[ derive( Resource ) ]
3032pub ( crate ) struct PendingConnection < C : DbContext + Send + Sync + ' static > (
31- pub ( crate ) Task < Result < Arc < C > > > ,
33+ pub ( crate ) Task < Result < ( Arc < C > , Arc < AtomicBool > ) > > ,
3234) ;
3335
3436/// Internal connection driver configuration.
@@ -100,7 +102,7 @@ where
100102 M : SpacetimeModule < DbConnection = C > ,
101103{
102104 /// Produces a configured [`DbConnectionBuilder`] for this connection.
103- fn connection_builder ( & self ) -> DbConnectionBuilder < M > {
105+ fn connection_builder ( & self , disconnect_requested : Arc < AtomicBool > ) -> DbConnectionBuilder < M > {
104106 let connected_tx = self . connected_tx . clone ( ) ;
105107 let disconnected_tx = self . disconnected_tx . clone ( ) ;
106108 let connect_error_tx = self . connect_error_tx . clone ( ) ;
@@ -117,7 +119,14 @@ where
117119 } ) ;
118120 } )
119121 . on_disconnect ( move |_ctx, err| {
120- let _ = disconnected_tx. send ( StdbDisconnectedMessage { err } ) ;
122+ let result = if disconnect_requested. swap ( false , Ordering :: AcqRel ) {
123+ Ok ( DisconnectIntent :: Requested )
124+ } else if let Some ( err) = err {
125+ Err ( err)
126+ } else {
127+ Ok ( DisconnectIntent :: Lost )
128+ } ;
129+ let _ = disconnected_tx. send ( StdbDisconnectedMessage { result } ) ;
121130 } )
122131 . on_connect_error ( move |_ctx, err| {
123132 // TODO: waiting for STDB release with fix for this to function properly.
@@ -128,11 +137,20 @@ where
128137 /// Builds a SpacetimeDB connection from this config.
129138 ///
130139 /// The returned connection is not started automatically.
131- pub ( crate ) async fn build_connection ( & self ) -> Result < Arc < C > > {
140+ pub ( crate ) async fn build_connection ( & self ) -> Result < ( Arc < C > , Arc < AtomicBool > ) > {
141+ let disconnect_requested = Arc :: new ( AtomicBool :: new ( false ) ) ;
132142 #[ cfg( not( feature = "browser" ) ) ]
133- return self . connection_builder ( ) . build ( ) . map ( Arc :: new) ;
143+ let connection = self
144+ . connection_builder ( Arc :: clone ( & disconnect_requested) )
145+ . build ( )
146+ . map ( Arc :: new) ?;
134147 #[ cfg( feature = "browser" ) ]
135- return self . connection_builder ( ) . build ( ) . await . map ( Arc :: new) ;
148+ let connection = self
149+ . connection_builder ( Arc :: clone ( & disconnect_requested) )
150+ . build ( )
151+ . await
152+ . map ( Arc :: new) ?;
153+ Ok ( ( connection, disconnect_requested) )
136154 }
137155}
138156
@@ -144,12 +162,16 @@ where
144162pub struct StdbConnection < T : DbContext + ' static > {
145163 /// The underlying connection context.
146164 conn : Arc < T > ,
165+ disconnect_requested : Arc < AtomicBool > ,
147166}
148167
149168impl < T : DbContext > StdbConnection < T > {
150169 /// Wraps an existing shared connection.
151- fn new ( conn : Arc < T > ) -> Self {
152- Self { conn }
170+ fn new ( conn : Arc < T > , disconnect_requested : Arc < AtomicBool > ) -> Self {
171+ Self {
172+ conn,
173+ disconnect_requested,
174+ }
153175 }
154176}
155177
@@ -176,7 +198,12 @@ impl<T: DbContext> StdbConnection<T> {
176198
177199 /// Closes the connection to the SpacetimeDB server.
178200 pub fn disconnect ( & self ) -> Result < ( ) > {
179- self . conn . disconnect ( )
201+ self . disconnect_requested . store ( true , Ordering :: Release ) ;
202+ let result = self . conn . disconnect ( ) ;
203+ if result. is_err ( ) {
204+ self . disconnect_requested . store ( false , Ordering :: Release ) ;
205+ }
206+ result
180207 }
181208
182209 /// Returns a builder for database subscriptions.
@@ -237,7 +264,6 @@ impl<
237264 register_channel :: < StdbConnectedMessage > ( app) ;
238265 register_channel :: < StdbDisconnectedMessage > ( app) ;
239266 register_channel :: < StdbConnectErrorMessage > ( app) ;
240- app. add_message :: < StdbDisconnectRequestedMessage > ( ) ;
241267
242268 let world = app. world ( ) ;
243269 app. insert_resource ( StdbConnectionConfig :: < C , M > {
@@ -303,7 +329,7 @@ fn poll_pending_connection<
303329 } ;
304330
305331 match result {
306- Ok ( conn) => {
332+ Ok ( ( conn, disconnect_requested ) ) => {
307333 let driver = world
308334 . get_resource :: < StdbConnectionConfig < C , M > > ( )
309335 . expect ( "StdbConnectionConfig should exist when activating a connection" )
@@ -317,7 +343,7 @@ fn poll_pending_connection<
317343 if let Some ( prev_conn) = world. get_resource :: < StdbConnection < C > > ( ) {
318344 let _ = prev_conn. disconnect ( ) ;
319345 }
320- world. insert_resource ( StdbConnection :: new ( conn) ) ;
346+ world. insert_resource ( StdbConnection :: new ( conn, disconnect_requested ) ) ;
321347 }
322348 Err ( err) => {
323349 world. write_message ( StdbConnectErrorMessage { err } ) ;
0 commit comments