11use super :: * ;
22use parking_lot:: RwLock ;
3- use std:: time:: Instant ;
3+ use std:: { fmt :: Display , time:: Instant } ;
44
55use tracing:: { error, warn} ;
66
@@ -11,6 +11,23 @@ pub struct Ban {
1111 pool : Pool ,
1212}
1313
14+ #[ derive( Debug , Copy , Clone ) ]
15+ pub enum UnbanReason {
16+ AllTargetsBanned ,
17+ Expired ,
18+ Manual ,
19+ }
20+
21+ impl Display for UnbanReason {
22+ fn fmt ( & self , f : & mut std:: fmt:: Formatter < ' _ > ) -> std:: fmt:: Result {
23+ match self {
24+ Self :: AllTargetsBanned => write ! ( f, "all targets banned" ) ,
25+ Self :: Expired => write ! ( f, "expired" ) ,
26+ Self :: Manual => write ! ( f, "manual" ) ,
27+ }
28+ }
29+ }
30+
1431impl Ban {
1532 /// Create new ban handler.
1633 pub ( super ) fn new ( pool : & Pool ) -> Self {
@@ -48,15 +65,25 @@ impl Ban {
4865 }
4966
5067 /// Unban the database.
51- pub fn unban ( & self , manual_check : bool ) {
68+ ///
69+ /// FIXME(lev): `reason` seems like it should be
70+ /// used as an operand but it's only used for logging.
71+ /// We should unify methods and provide one public interface to this.
72+ ///
73+ pub fn unban ( & self , manual_check : bool , reason : UnbanReason ) {
5274 let mut guard = self . inner . upgradable_read ( ) ;
5375 if let Some ( ref ban) = guard. ban {
76+ let mut unbanned = false ;
5477 if ban. error != Error :: ManualBan || !manual_check {
5578 guard. with_upgraded ( |guard| {
5679 guard. ban = None ;
5780 } ) ;
81+ unbanned = true ;
82+ }
83+
84+ if unbanned {
85+ warn ! ( "resuming read queries: {} [{}]" , reason, self . pool. addr( ) ) ;
5886 }
59- warn ! ( "resuming read queries [{}]" , self . pool. addr( ) ) ;
6087 }
6188 }
6289
@@ -114,7 +141,11 @@ impl Ban {
114141 } ;
115142 drop ( guard) ;
116143 if unbanned {
117- warn ! ( "resuming read queries [{}]" , self . pool. addr( ) ) ;
144+ warn ! (
145+ "resuming read queries: {} [{}]" ,
146+ UnbanReason :: Expired ,
147+ self . pool. addr( )
148+ ) ;
118149 }
119150 unbanned
120151 }
@@ -185,7 +216,7 @@ mod tests {
185216 let pool = Pool :: new_test ( ) ;
186217 let ban = Ban :: new ( & pool) ;
187218 ban. ban ( Error :: ServerError , Duration :: from_secs ( 1 ) ) ;
188- ban. unban ( false ) ;
219+ ban. unban ( false , UnbanReason :: Expired ) ;
189220 assert ! ( !ban. banned( ) ) ;
190221 assert ! ( ban. error( ) . is_none( ) ) ;
191222 }
@@ -195,7 +226,7 @@ mod tests {
195226 let pool = Pool :: new_test ( ) ;
196227 let ban = Ban :: new ( & pool) ;
197228 ban. ban ( Error :: ManualBan , Duration :: from_secs ( 1 ) ) ;
198- ban. unban ( true ) ;
229+ ban. unban ( true , UnbanReason :: Expired ) ;
199230 assert ! ( ban. banned( ) ) ;
200231 assert_eq ! ( ban. error( ) , Some ( Error :: ManualBan ) ) ;
201232 }
@@ -205,7 +236,7 @@ mod tests {
205236 let pool = Pool :: new_test ( ) ;
206237 let ban = Ban :: new ( & pool) ;
207238 ban. ban ( Error :: ServerError , Duration :: from_secs ( 1 ) ) ;
208- ban. unban ( true ) ;
239+ ban. unban ( true , UnbanReason :: Manual ) ;
209240 assert ! ( !ban. banned( ) ) ;
210241 assert ! ( ban. error( ) . is_none( ) ) ;
211242 }
@@ -303,7 +334,7 @@ mod tests {
303334
304335 for _ in 0 ..100 {
305336 ban. ban ( Error :: ServerError , Duration :: from_secs ( 1 ) ) ;
306- ban. unban ( false ) ;
337+ ban. unban ( false , UnbanReason :: Expired ) ;
307338 }
308339
309340 h1. join ( ) . unwrap ( ) ;
@@ -374,7 +405,7 @@ mod tests {
374405 for _ in 0 ..10 {
375406 assert ! ( ban. ban( Error :: ServerError , Duration :: from_secs( 1 ) ) ) ;
376407 assert ! ( ban. banned( ) ) ;
377- ban. unban ( false ) ;
408+ ban. unban ( false , UnbanReason :: AllTargetsBanned ) ;
378409 assert ! ( !ban. banned( ) ) ;
379410 }
380411 }
0 commit comments