Skip to main content

sa_token_core/
disable.rs

1// Author: 金书记
2//
3//! Account / Service Ban | 账号/服务封禁
4//!
5//! Account disable / ban checks, with one deliberate
6//! difference: every method is **account-system aware** (A3-18).
7
8use std::time::Duration;
9
10use crate::error::{SaTokenError, SaTokenResult};
11use crate::keys::LOGIN_TYPE_DEFAULT;
12use crate::manager::SaTokenManager;
13
14/// Default ban service identifier | 默认封禁服务标识
15pub const DEFAULT_DISABLE_SERVICE: &str = "login";
16
17/// Minimum ban level | 最低封禁等级
18pub const MIN_DISABLE_LEVEL: i32 = 1;
19
20/// Level returned when an account is not banned | 账号未被封禁时返回的等级
21pub const NOT_DISABLE_LEVEL: i32 = -2;
22
23/// Default level written by [`SaTokenManager::disable`] | 默认封禁等级
24pub const DEFAULT_DISABLE_LEVEL: i32 = 1;
25
26impl SaTokenManager {
27    #[inline]
28    fn disable_key_ns(&self, login_type: &str, login_id: &str, service: &str) -> String {
29        self.keys().disable(login_type, login_id, service)
30    }
31
32    /// Disable at a level for a login type | 按登录类型分级禁用
33    pub async fn disable_level_with_type(
34        &self,
35        login_type: &str,
36        login_id: &str,
37        service: &str,
38        level: i32,
39        time: i64,
40    ) -> SaTokenResult<()> {
41        if login_id.trim().is_empty() {
42            return Err(SaTokenError::ConfigError(
43                "login_id is required for disable".to_string(),
44            ));
45        }
46        if service.trim().is_empty() {
47            return Err(SaTokenError::ConfigError(
48                "service is required for disable".to_string(),
49            ));
50        }
51        if level < MIN_DISABLE_LEVEL && level != 0 {
52            return Err(SaTokenError::ConfigError(format!(
53                "disable level must be >= {MIN_DISABLE_LEVEL} (0 allowed)"
54            )));
55        }
56
57        let ttl = if time < 0 {
58            None
59        } else {
60            Some(Duration::from_secs(time as u64))
61        };
62
63        self.dao
64            .set_string(
65                &self.disable_key_ns(login_type, login_id, service),
66                &level.to_string(),
67                ttl,
68            )
69            .await?;
70
71        let ns = self.account_ns(login_type, login_id);
72        let event = crate::event::SaTokenEvent::banned(ns.as_str(), service, level)
73            .with_login_type(login_type);
74        self.event_bus.publish(event).await;
75
76        Ok(())
77    }
78
79    /// Disable for a login type | 按登录类型禁用
80    pub async fn disable_with_type(
81        &self,
82        login_type: &str,
83        login_id: &str,
84        time: i64,
85    ) -> SaTokenResult<()> {
86        self.disable_level_with_type(
87            login_type,
88            login_id,
89            DEFAULT_DISABLE_SERVICE,
90            DEFAULT_DISABLE_LEVEL,
91            time,
92        )
93        .await
94    }
95
96    /// Read disable level for a login type | 按登录类型读取禁用等级
97    pub async fn get_disable_level_with_type(
98        &self,
99        login_type: &str,
100        login_id: &str,
101        service: &str,
102    ) -> SaTokenResult<i32> {
103        let key = self.disable_key_ns(login_type, login_id, service);
104        let value = self.dao.get_string(&key).await?;
105
106        if let Some(v) = value {
107            return v.parse::<i32>().map_err(|_| {
108                SaTokenError::StorageError(format!("invalid disable level for key {key}"))
109            });
110        }
111
112        if let Some(level) = self.authz_service().is_disabled(login_id, service).await? {
113            return Ok(level);
114        }
115
116        Ok(NOT_DISABLE_LEVEL)
117    }
118
119    /// Whether disabled at/above level | 是否达到指定禁用等级
120    pub async fn is_disable_level_with_type(
121        &self,
122        login_type: &str,
123        login_id: &str,
124        service: &str,
125        level: i32,
126    ) -> SaTokenResult<bool> {
127        let disable_level = self
128            .get_disable_level_with_type(login_type, login_id, service)
129            .await?;
130        if disable_level == NOT_DISABLE_LEVEL {
131            return Ok(false);
132        }
133        Ok(disable_level >= level)
134    }
135
136    /// Fail if disabled at/above level | 达到禁用等级则报错
137    pub async fn check_disable_level_with_type(
138        &self,
139        login_type: &str,
140        login_id: &str,
141        service: &str,
142        level: i32,
143    ) -> SaTokenResult<()> {
144        let disable_level = self
145            .get_disable_level_with_type(login_type, login_id, service)
146            .await?;
147        if disable_level == NOT_DISABLE_LEVEL {
148            return Ok(());
149        }
150        if disable_level >= level {
151            return Err(SaTokenError::AccountBanned(format!(
152                "service={service} level={disable_level}"
153            )));
154        }
155        Ok(())
156    }
157
158    /// Fail if any listed service is disabled | 任一服务被禁用则报错
159    pub async fn check_disable_services_with_type(
160        &self,
161        login_type: &str,
162        login_id: &str,
163        services: &[&str],
164        level: i32,
165    ) -> SaTokenResult<()> {
166        for service in services {
167            self.check_disable_level_with_type(login_type, login_id, service, level)
168                .await?;
169        }
170        Ok(())
171    }
172
173    /// Clear disable for a login type | 按登录类型解除禁用
174    pub async fn untie_disable_with_type(
175        &self,
176        login_type: &str,
177        login_id: &str,
178        service: &str,
179    ) -> SaTokenResult<()> {
180        self.dao
181            .delete(&self.disable_key_ns(login_type, login_id, service))
182            .await?;
183
184        let ns = self.account_ns(login_type, login_id);
185        let event =
186            crate::event::SaTokenEvent::unbanned(ns.as_str(), service).with_login_type(login_type);
187        self.event_bus.publish(event).await;
188
189        Ok(())
190    }
191
192    /// Disable at a level (default login type) | 分级禁用(默认登录类型)
193    pub async fn disable_level(
194        &self,
195        login_id: &str,
196        service: &str,
197        level: i32,
198        time: i64,
199    ) -> SaTokenResult<()> {
200        self.disable_level_with_type(LOGIN_TYPE_DEFAULT, login_id, service, level, time)
201            .await
202    }
203
204    /// Disable account/service | 禁用账号或服务
205    pub async fn disable(&self, login_id: &str, time: i64) -> SaTokenResult<()> {
206        self.disable_with_type(LOGIN_TYPE_DEFAULT, login_id, time)
207            .await
208    }
209
210    /// Read disable level | 读取禁用等级
211    pub async fn get_disable_level(&self, login_id: &str, service: &str) -> SaTokenResult<i32> {
212        self.get_disable_level_with_type(LOGIN_TYPE_DEFAULT, login_id, service)
213            .await
214    }
215
216    /// Whether disabled at/above level | 是否达到指定禁用等级
217    pub async fn is_disable_level(
218        &self,
219        login_id: &str,
220        service: &str,
221        level: i32,
222    ) -> SaTokenResult<bool> {
223        self.is_disable_level_with_type(LOGIN_TYPE_DEFAULT, login_id, service, level)
224            .await
225    }
226
227    /// Fail if disabled at/above level | 达到禁用等级则报错
228    pub async fn check_disable_level(
229        &self,
230        login_id: &str,
231        service: &str,
232        level: i32,
233    ) -> SaTokenResult<()> {
234        self.check_disable_level_with_type(LOGIN_TYPE_DEFAULT, login_id, service, level)
235            .await
236    }
237
238    /// Fail if any listed service is disabled | 任一服务被禁用则报错
239    pub async fn check_disable_services(
240        &self,
241        login_id: &str,
242        services: &[&str],
243        level: i32,
244    ) -> SaTokenResult<()> {
245        self.check_disable_services_with_type(LOGIN_TYPE_DEFAULT, login_id, services, level)
246            .await
247    }
248
249    /// Clear disable flag | 解除禁用
250    pub async fn untie_disable(&self, login_id: &str, service: &str) -> SaTokenResult<()> {
251        self.untie_disable_with_type(LOGIN_TYPE_DEFAULT, login_id, service)
252            .await
253    }
254}
255
256#[cfg(test)]
257mod tests {
258    use super::*;
259    use crate::config::SaTokenConfig;
260    use sa_token_storage_memory::MemoryStorage;
261    use std::sync::Arc;
262
263    fn manager() -> SaTokenManager {
264        SaTokenManager::new(Arc::new(MemoryStorage::new()), SaTokenConfig::default())
265    }
266
267    #[tokio::test]
268    async fn disable_and_check_level() {
269        let mgr = manager();
270        mgr.disable_level("u1", "login", 2, 60).await.unwrap();
271        assert!(mgr.is_disable_level("u1", "login", 1).await.unwrap());
272        assert!(mgr.is_disable_level("u1", "login", 2).await.unwrap());
273        assert!(!mgr.is_disable_level("u1", "login", 3).await.unwrap());
274        assert!(mgr.check_disable_level("u1", "login", 2).await.is_err());
275        mgr.untie_disable("u1", "login").await.unwrap();
276        assert_eq!(
277            mgr.get_disable_level("u1", "login").await.unwrap(),
278            NOT_DISABLE_LEVEL
279        );
280    }
281}