sa-token-python-core 0.1.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (46) hide show
  1. sa_token/__init__.py +89 -0
  2. sa_token/adapter/__init__.py +24 -0
  3. sa_token/adapter/http.py +71 -0
  4. sa_token/adapter/path.py +163 -0
  5. sa_token/adapter/pipeline.py +97 -0
  6. sa_token/config.py +130 -0
  7. sa_token/context.py +63 -0
  8. sa_token/exception.py +143 -0
  9. sa_token/integration/__init__.py +10 -0
  10. sa_token/integration/django.py +131 -0
  11. sa_token/integration/fastapi.py +315 -0
  12. sa_token/integration/fastapi_oauth2.py +136 -0
  13. sa_token/integration/flask.py +191 -0
  14. sa_token/integration/starlette.py +227 -0
  15. sa_token/listener.py +100 -0
  16. sa_token/manager.py +244 -0
  17. sa_token/model.py +145 -0
  18. sa_token/oauth2/__init__.py +19 -0
  19. sa_token/oauth2/model.py +122 -0
  20. sa_token/oauth2/server.py +361 -0
  21. sa_token/online/__init__.py +292 -0
  22. sa_token/permission.py +67 -0
  23. sa_token/py.typed +0 -0
  24. sa_token/security/__init__.py +14 -0
  25. sa_token/security/nonce.py +93 -0
  26. sa_token/security/refresh.py +300 -0
  27. sa_token/security/temp_token.py +114 -0
  28. sa_token/session.py +96 -0
  29. sa_token/sso/__init__.py +217 -0
  30. sa_token/storage/__init__.py +22 -0
  31. sa_token/storage/base.py +66 -0
  32. sa_token/storage/memory.py +154 -0
  33. sa_token/storage/redis.py +136 -0
  34. sa_token/stp_interface.py +20 -0
  35. sa_token/stp_logic.py +911 -0
  36. sa_token/stp_util.py +367 -0
  37. sa_token/strategy/__init__.py +77 -0
  38. sa_token/strategy/base.py +22 -0
  39. sa_token/strategy/builtin.py +99 -0
  40. sa_token/strategy/jwt.py +72 -0
  41. sa_token/sync.py +268 -0
  42. sa_token/token_io.py +66 -0
  43. sa_token_python_core-0.1.1.dist-info/METADATA +756 -0
  44. sa_token_python_core-0.1.1.dist-info/RECORD +46 -0
  45. sa_token_python_core-0.1.1.dist-info/WHEEL +4 -0
  46. sa_token_python_core-0.1.1.dist-info/licenses/LICENSE +201 -0
sa_token/stp_logic.py ADDED
@@ -0,0 +1,911 @@
1
+ """StpLogic:认证鉴权的唯一实现。
2
+
3
+ 这是整个项目的核心。框架适配层、OAuth2、SSO、在线用户全部复用这里的方法,
4
+ 禁止在上层重新实现登录校验或权限匹配,否则语义一定会分叉。
5
+
6
+ 所有方法都不感知 HTTP:参数要么是 ``login_id``,要么是 ``token`` 字符串。
7
+ 因此脚本、定时任务、RPC 与 Web 接口用的是同一套逻辑。
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import json
13
+ from typing import TYPE_CHECKING, Any
14
+
15
+ from .context import get_current_token, set_current
16
+ from .exception import (
17
+ DisableException,
18
+ NotLoginException,
19
+ NotLoginType,
20
+ NotPermissionException,
21
+ NotRoleException,
22
+ NotSafeException,
23
+ SaTokenException,
24
+ )
25
+ from .listener import Event, EventData, fingerprint
26
+ from .model import (
27
+ DEFAULT_DEVICE,
28
+ SessionData,
29
+ TerminalInfo,
30
+ TokenInfo,
31
+ now_ms,
32
+ )
33
+ from .permission import MatchMode, has_element, match_all, match_any
34
+ from .session import SaSession
35
+
36
+ if TYPE_CHECKING: # pragma: no cover - 仅供类型检查
37
+ from .config import SaTokenConfig
38
+ from .manager import SaTokenManager
39
+ from .security import LoginTokenPair
40
+ from .storage.base import SaStorage
41
+
42
+ __all__ = ["StpLogic"]
43
+
44
+ #: 默认的封禁服务名,对应「整个账号被封」。
45
+ DEFAULT_DISABLE_SERVICE = "login"
46
+
47
+
48
+ class StpLogic:
49
+ """单个账号体系(``login_type``)的认证逻辑。
50
+
51
+ 同一个 Manager 下可以有多个 ``StpLogic``(如 ``user`` 与 ``admin``),
52
+ 它们的存储键互相隔离,同一个 ``login_id`` 在两套体系里是两个独立身份。
53
+ """
54
+
55
+ def __init__(self, manager: SaTokenManager, login_type: str = "login") -> None:
56
+ self._manager = manager
57
+ self.login_type = login_type
58
+
59
+ # ------------------------------------------------------------------ 基础设施
60
+
61
+ @property
62
+ def config(self) -> SaTokenConfig:
63
+ return self._manager.config
64
+
65
+ @property
66
+ def storage(self) -> SaStorage:
67
+ return self._manager.storage
68
+
69
+ def _token_key(self, token: str) -> str:
70
+ return self.config.make_key(self.login_type, "token", token)
71
+
72
+ def _session_key(self, login_id: str) -> str:
73
+ return self.config.make_key(self.login_type, "session", login_id)
74
+
75
+ def _token_session_key(self, token: str) -> str:
76
+ return self.config.make_key(self.login_type, "token-session", token)
77
+
78
+ def _last_active_key(self, token: str) -> str:
79
+ return self.config.make_key(self.login_type, "last-active", token)
80
+
81
+ def _permission_key(self, login_id: str) -> str:
82
+ return self.config.make_key(self.login_type, "permission", login_id)
83
+
84
+ def _role_key(self, login_id: str) -> str:
85
+ return self.config.make_key(self.login_type, "role", login_id)
86
+
87
+ def _permission_cache_key(self, login_id: str, kind: str) -> str:
88
+ return self.config.make_key(self.login_type, f"{kind}-cache", login_id)
89
+
90
+ def _disable_key(self, login_id: str, service: str) -> str:
91
+ return self.config.make_key(self.login_type, "disable", f"{login_id}:{service}")
92
+
93
+ def _safe_key(self, token: str, business: str) -> str:
94
+ return self.config.make_key(self.login_type, "safe", f"{token}:{business}")
95
+
96
+ def _session_ttl(self) -> int | None:
97
+ return None if self.config.timeout < 0 else self.config.timeout
98
+
99
+ @staticmethod
100
+ def normalize_login_id(login_id: Any) -> str:
101
+ """统一成字符串并拒绝会破坏存储键的取值。"""
102
+ if login_id is None:
103
+ raise SaTokenException("login_id 不能为空")
104
+ normalized = str(login_id).strip()
105
+ if not normalized:
106
+ raise SaTokenException("login_id 不能为空")
107
+ if ":" in normalized:
108
+ raise SaTokenException("login_id 不能包含冒号,它与存储键分隔符冲突")
109
+ return normalized
110
+
111
+ async def _emit(
112
+ self,
113
+ event: Event,
114
+ *,
115
+ login_id: str | None = None,
116
+ device: str | None = None,
117
+ token: str | None = None,
118
+ **detail: Any,
119
+ ) -> None:
120
+ await self._manager.events.emit(
121
+ EventData(
122
+ event=event,
123
+ login_id=login_id,
124
+ login_type=self.login_type,
125
+ device=device,
126
+ token_fingerprint=fingerprint(token),
127
+ detail=detail,
128
+ )
129
+ )
130
+
131
+ # ------------------------------------------------------------------ 登录
132
+
133
+ async def login(
134
+ self,
135
+ login_id: Any,
136
+ *,
137
+ device: str | None = None,
138
+ timeout: int | None = None,
139
+ tag: str | None = None,
140
+ extra: dict[str, Any] | None = None,
141
+ token_value: str | None = None,
142
+ ) -> str:
143
+ """登录并返回 token。
144
+
145
+ 业务侧负责校验账号密码,本方法只负责签发与记录登录态。
146
+ """
147
+ normalized_id = self.normalize_login_id(login_id)
148
+ device_name = device or DEFAULT_DEVICE
149
+ await self.check_disable(normalized_id)
150
+
151
+ effective_timeout = self.config.timeout if timeout is None else timeout
152
+ session = await self.get_session(normalized_id, create=True)
153
+ assert session is not None
154
+
155
+ reused_token = await self._apply_concurrent_policy(session, normalized_id, device_name)
156
+ if reused_token is not None:
157
+ return reused_token
158
+
159
+ token, info = await self._allocate_token(
160
+ normalized_id,
161
+ device_name,
162
+ effective_timeout,
163
+ tag,
164
+ extra,
165
+ token_value,
166
+ )
167
+ await self._touch_active(token, effective_timeout)
168
+
169
+ session.raw.history_terminal_count += 1
170
+ session.raw.terminal_list.append(
171
+ TerminalInfo(
172
+ token=token,
173
+ device=device_name,
174
+ index=session.raw.history_terminal_count,
175
+ )
176
+ )
177
+ await self._enforce_max_login_count(session)
178
+ await session.save()
179
+
180
+ await self._emit(Event.LOGIN, login_id=normalized_id, device=device_name, token=token)
181
+ return token
182
+
183
+ async def login_with_refresh(
184
+ self,
185
+ login_id: Any,
186
+ *,
187
+ device: str | None = None,
188
+ timeout: int | None = None,
189
+ tag: str | None = None,
190
+ extra: dict[str, Any] | None = None,
191
+ ) -> LoginTokenPair:
192
+ """登录并签发 access/refresh token 对。"""
193
+ normalized_id = self.normalize_login_id(login_id)
194
+ device_name = device or DEFAULT_DEVICE
195
+ access_token = await self.login(
196
+ normalized_id,
197
+ device=device_name,
198
+ timeout=timeout,
199
+ tag=tag,
200
+ extra=extra,
201
+ )
202
+ return await self._manager.refresh_tokens.issue(
203
+ access_token,
204
+ normalized_id,
205
+ login_type=self.login_type,
206
+ device=device_name,
207
+ )
208
+
209
+ async def _allocate_token(
210
+ self,
211
+ login_id: str,
212
+ device: str,
213
+ timeout: int | None,
214
+ tag: str | None,
215
+ extra: dict[str, Any] | None,
216
+ token_value: str | None,
217
+ ) -> tuple[str, TokenInfo]:
218
+ """原子占用 token;自定义值冲突时报错,随机值冲突则重试。"""
219
+ if token_value is not None and not token_value.strip():
220
+ raise SaTokenException("指定的 token_value 不能为空")
221
+ attempts = 1 if token_value is not None else 12
222
+ for _ in range(attempts):
223
+ token = token_value or self._manager.strategy.generate(login_id, extra)
224
+ info = TokenInfo(
225
+ login_id=login_id,
226
+ device=device,
227
+ login_type=self.login_type,
228
+ timeout=timeout,
229
+ tag=tag,
230
+ )
231
+ if await self.storage.set_if_absent(self._token_key(token), info.to_json(), timeout):
232
+ return token, info
233
+ if token_value is not None:
234
+ raise SaTokenException("指定的 token_value 已被占用")
235
+ raise SaTokenException("无法分配唯一 token,请检查自定义 TokenStrategy")
236
+
237
+ async def _apply_concurrent_policy(
238
+ self,
239
+ session: SaSession,
240
+ login_id: str,
241
+ device: str,
242
+ ) -> str | None:
243
+ """按 is_concurrent / is_share 处理已有登录,返回可复用的 token。"""
244
+ if not self.config.is_concurrent:
245
+ # 不允许并发在线:旧登录一律顶下线。
246
+ scope_all = self.config.replaced_range == "all_device"
247
+ await self._offline_terminals(
248
+ session,
249
+ login_id,
250
+ device=None if scope_all else device,
251
+ state=NotLoginType.BE_REPLACED,
252
+ )
253
+ return None
254
+
255
+ if not self.config.is_share:
256
+ return None
257
+
258
+ # 共享模式:同设备类型复用同一个 token,前提是它仍然有效。
259
+ for terminal in session.raw.terminal_list:
260
+ if terminal.device != device:
261
+ continue
262
+ info = await self._read_token_info(terminal.token)
263
+ if info is not None and not info.is_offline:
264
+ return terminal.token
265
+ return None
266
+
267
+ async def _enforce_max_login_count(self, session: SaSession) -> None:
268
+ limit = self.config.max_login_count
269
+ if limit is None or limit <= 0:
270
+ return
271
+ overflow = len(session.raw.terminal_list) - limit
272
+ if overflow <= 0:
273
+ return
274
+
275
+ mode = self.config.overflow_logout_mode
276
+ state = {
277
+ "kickout": NotLoginType.KICK_OUT,
278
+ "replaced": NotLoginType.BE_REPLACED,
279
+ }.get(mode)
280
+ # 终端列表按登录时间追加,队首即最旧的登录。
281
+ for terminal in session.raw.terminal_list[:overflow]:
282
+ if state is None:
283
+ await self._destroy_token(terminal.token)
284
+ else:
285
+ await self._mark_token_offline(terminal.token, state)
286
+ session.raw.terminal_list = session.raw.terminal_list[overflow:]
287
+
288
+ # ------------------------------------------------------------------ 登出 / 踢人
289
+
290
+ async def logout(self, login_id: Any, *, device: str | None = None) -> None:
291
+ """主动登出:token 记录直接删除,不留下线原因。"""
292
+ normalized_id = self.normalize_login_id(login_id)
293
+ session = await self.get_session(normalized_id, create=False)
294
+ if session is None:
295
+ return
296
+ removed = await self._remove_terminals(session, device, destroy=True)
297
+ await self._save_or_drop_session(session, normalized_id)
298
+ if device is None:
299
+ await self._manager.refresh_tokens.revoke_all_for_login(
300
+ self.login_type, normalized_id
301
+ )
302
+ if removed:
303
+ await self._emit(Event.LOGOUT, login_id=normalized_id, device=device)
304
+
305
+ async def logout_by_token(
306
+ self, token: str | None = None, *, revoke_refresh: bool = True
307
+ ) -> None:
308
+ """登出单个 token,常用于「只退出当前设备」。"""
309
+ token = token or get_current_token()
310
+ if not token:
311
+ return
312
+ info = await self._read_token_info(token)
313
+ if revoke_refresh:
314
+ await self._manager.refresh_tokens.revoke_for_access(token)
315
+ await self._destroy_token(token)
316
+ if info is None:
317
+ return
318
+ session = await self.get_session(info.login_id, create=False)
319
+ if session is not None:
320
+ session.raw.terminal_list = [
321
+ terminal for terminal in session.raw.terminal_list if terminal.token != token
322
+ ]
323
+ await self._save_or_drop_session(session, info.login_id)
324
+ await self._emit(
325
+ Event.LOGOUT, login_id=info.login_id, device=info.device, token=token
326
+ )
327
+
328
+ async def kickout(self, login_id: Any, *, device: str | None = None) -> None:
329
+ """踢人下线:保留下线原因,用户下次请求会收到 ``KICK_OUT``。"""
330
+ await self._offline_by_login_id(login_id, device, NotLoginType.KICK_OUT, Event.KICKOUT)
331
+
332
+ async def kickout_by_token(self, token: str | None = None) -> None:
333
+ """按 token 踢人,保留 ``KICK_OUT`` 原因并移除对应终端。"""
334
+ token = token or get_current_token()
335
+ if not token:
336
+ return
337
+ info = await self._read_token_info(token)
338
+ if info is None:
339
+ return
340
+ await self._manager.refresh_tokens.revoke_for_access(token)
341
+ await self._mark_token_offline(token, NotLoginType.KICK_OUT)
342
+ session = await self.get_session(info.login_id, create=False)
343
+ if session is not None:
344
+ session.raw.terminal_list = [
345
+ terminal for terminal in session.raw.terminal_list if terminal.token != token
346
+ ]
347
+ await self._save_or_drop_session(session, info.login_id)
348
+ await self._emit(
349
+ Event.KICKOUT,
350
+ login_id=info.login_id,
351
+ device=info.device,
352
+ token=token,
353
+ )
354
+
355
+ async def replaced(self, login_id: Any, *, device: str | None = None) -> None:
356
+ """顶号下线:语义上「被新登录挤掉」,与踢人区分开便于前端提示。"""
357
+ await self._offline_by_login_id(login_id, device, NotLoginType.BE_REPLACED, Event.REPLACED)
358
+
359
+ async def _offline_by_login_id(
360
+ self,
361
+ login_id: Any,
362
+ device: str | None,
363
+ state: NotLoginType,
364
+ event: Event,
365
+ ) -> None:
366
+ normalized_id = self.normalize_login_id(login_id)
367
+ session = await self.get_session(normalized_id, create=False)
368
+ if session is None:
369
+ return
370
+ offline_count = await self._offline_terminals(session, normalized_id, device, state)
371
+ await self._save_or_drop_session(session, normalized_id)
372
+ if device is None:
373
+ await self._manager.refresh_tokens.revoke_all_for_login(
374
+ self.login_type, normalized_id
375
+ )
376
+ if offline_count:
377
+ await self._emit(event, login_id=normalized_id, device=device)
378
+
379
+ async def _offline_terminals(
380
+ self,
381
+ session: SaSession,
382
+ login_id: str,
383
+ device: str | None,
384
+ state: NotLoginType,
385
+ ) -> int:
386
+ """把匹配的终端标记为下线状态,并从终端列表移除。"""
387
+ matched = [
388
+ terminal
389
+ for terminal in session.raw.terminal_list
390
+ if device is None or terminal.device == device
391
+ ]
392
+ for terminal in matched:
393
+ await self._mark_token_offline(terminal.token, state)
394
+ if matched:
395
+ removed = {terminal.token for terminal in matched}
396
+ session.raw.terminal_list = [
397
+ terminal
398
+ for terminal in session.raw.terminal_list
399
+ if terminal.token not in removed
400
+ ]
401
+ return len(matched)
402
+
403
+ async def _remove_terminals(
404
+ self,
405
+ session: SaSession,
406
+ device: str | None,
407
+ *,
408
+ destroy: bool,
409
+ ) -> int:
410
+ matched = [
411
+ terminal
412
+ for terminal in session.raw.terminal_list
413
+ if device is None or terminal.device == device
414
+ ]
415
+ for terminal in matched:
416
+ if destroy:
417
+ await self._destroy_token(terminal.token)
418
+ if matched:
419
+ removed = {terminal.token for terminal in matched}
420
+ session.raw.terminal_list = [
421
+ terminal
422
+ for terminal in session.raw.terminal_list
423
+ if terminal.token not in removed
424
+ ]
425
+ return len(matched)
426
+
427
+ async def _save_or_drop_session(self, session: SaSession, login_id: str) -> None:
428
+ """没有终端在线且未配置保留时,顺手清掉 Account-Session。"""
429
+ if not session.raw.terminal_list and not self.config.is_logout_keep_session:
430
+ await self.storage.delete(self._session_key(login_id))
431
+ return
432
+ await session.save()
433
+
434
+ async def _destroy_token(self, token: str) -> None:
435
+ await self.storage.delete(self._token_key(token))
436
+ await self.storage.delete(self._token_session_key(token))
437
+ await self.storage.delete(self._last_active_key(token))
438
+
439
+ async def _mark_token_offline(self, token: str, state: NotLoginType) -> None:
440
+ if not self.config.offline_record_enabled:
441
+ await self._destroy_token(token)
442
+ return
443
+ info = await self._read_token_info(token)
444
+ if info is None:
445
+ return
446
+ info.state = state.value
447
+ info.offline_time = now_ms()
448
+ await self.storage.set(
449
+ self._token_key(token), info.to_json(), self.config.offline_record_timeout
450
+ )
451
+ await self.storage.delete(self._token_session_key(token))
452
+ await self.storage.delete(self._last_active_key(token))
453
+
454
+ # ------------------------------------------------------------------ 校验
455
+
456
+ async def check_login(self, token: str | None = None) -> str:
457
+ """校验登录态,返回 ``login_id``;失败抛 :class:`NotLoginException`。"""
458
+ token = token or get_current_token()
459
+ if not token:
460
+ raise NotLoginException(NotLoginType.NOT_TOKEN, login_type=self.login_type)
461
+
462
+ info = await self._read_token_info(token)
463
+ if info is None:
464
+ raise NotLoginException(
465
+ NotLoginType.INVALID_TOKEN, login_type=self.login_type, token=token
466
+ )
467
+ if info.is_offline:
468
+ offline_type = info.offline_type
469
+ assert offline_type is not None
470
+ raise NotLoginException(offline_type, login_type=self.login_type, token=token)
471
+
472
+ await self._check_active_timeout(token, info)
473
+ await self.check_disable(info.login_id)
474
+ await self._renew(token, info)
475
+
476
+ set_current(token, info.login_id)
477
+ return info.login_id
478
+
479
+ async def is_login(self, token: str | None = None) -> bool:
480
+ """:meth:`check_login` 的布尔包装,不抛异常。"""
481
+ try:
482
+ await self.check_login(token)
483
+ return True
484
+ except (NotLoginException, DisableException):
485
+ return False
486
+
487
+ async def get_login_id(self, token: str | None = None) -> str:
488
+ return await self.check_login(token)
489
+
490
+ async def get_login_id_or_none(self, token: str | None = None) -> str | None:
491
+ try:
492
+ return await self.check_login(token)
493
+ except (NotLoginException, DisableException):
494
+ return None
495
+
496
+ async def get_token_info(self, token: str | None = None) -> TokenInfo | None:
497
+ token = token or get_current_token()
498
+ if not token:
499
+ return None
500
+ return await self._read_token_info(token)
501
+
502
+ async def get_offline_reason(self, token: str) -> dict[str, Any] | None:
503
+ """查询被踢 / 被顶的下线原因与时间。"""
504
+ info = await self._read_token_info(token)
505
+ if info is None or not info.is_offline:
506
+ return None
507
+ return {"reason": info.state, "time": info.offline_time}
508
+
509
+ async def _read_token_info(self, token: str) -> TokenInfo | None:
510
+ raw = await self.storage.get(self._token_key(token))
511
+ if raw is None:
512
+ return None
513
+ return TokenInfo.from_json(raw)
514
+
515
+ async def _write_token_info(self, token: str, info: TokenInfo, timeout: int | None) -> None:
516
+ await self.storage.set(self._token_key(token), info.to_json(), timeout)
517
+
518
+ async def _check_active_timeout(self, token: str, info: TokenInfo) -> None:
519
+ active_timeout = (
520
+ info.active_timeout
521
+ if self.config.dynamic_active_timeout and info.active_timeout is not None
522
+ else self.config.active_timeout
523
+ )
524
+ if active_timeout is None or active_timeout < 0:
525
+ return
526
+ raw = await self.storage.get(self._last_active_key(token))
527
+ last_active = int(raw) if raw and raw.isdigit() else info.active_time
528
+ if now_ms() - last_active > active_timeout * 1000:
529
+ await self._mark_token_offline(token, NotLoginType.TOKEN_FREEZE)
530
+ raise NotLoginException(
531
+ NotLoginType.TOKEN_FREEZE, login_type=self.login_type, token=token
532
+ )
533
+
534
+ async def _touch_active(self, token: str, timeout: int | None) -> None:
535
+ if self.config.active_timeout < 0 and not self.config.dynamic_active_timeout:
536
+ return
537
+ await self.storage.set(self._last_active_key(token), str(now_ms()), timeout)
538
+
539
+ async def _renew(self, token: str, info: TokenInfo) -> None:
540
+ """续期在校验通过后同步执行。
541
+
542
+ 这里刻意不用后台任务:``asyncio.create_task`` 产生的孤儿任务在请求
543
+ 结束、事件循环关闭时可能被丢弃,导致续期静默失败,排查成本远高于
544
+ 它省下的那点延迟。
545
+ """
546
+ timeout = info.timeout if info.timeout is not None else self.config.timeout
547
+ if timeout is not None and timeout >= 0:
548
+ await self._touch_active(token, timeout)
549
+ if not self.config.auto_renew or timeout is None or timeout < 0:
550
+ return
551
+ await self.storage.expire(self._token_key(token), timeout)
552
+ await self.storage.expire(self._session_key(info.login_id), timeout)
553
+ await self.storage.expire(self._token_session_key(token), timeout)
554
+ await self._emit(Event.RENEW, login_id=info.login_id, token=token, timeout=timeout)
555
+
556
+ async def renew_timeout(self, token: str, timeout: int) -> bool:
557
+ """手动续签指定 token。"""
558
+ info = await self._read_token_info(token)
559
+ if info is None or info.is_offline:
560
+ return False
561
+ info.timeout = timeout
562
+ await self._write_token_info(token, info, timeout)
563
+ await self.storage.expire(self._session_key(info.login_id), timeout)
564
+ return True
565
+
566
+ # ------------------------------------------------------------------ Session
567
+
568
+ async def get_session(self, login_id: Any, *, create: bool = True) -> SaSession | None:
569
+ normalized_id = self.normalize_login_id(login_id)
570
+ key = self._session_key(normalized_id)
571
+ raw = await self.storage.get(key)
572
+ data = SessionData.from_json(raw) if raw else None
573
+ if data is None:
574
+ if not create:
575
+ return None
576
+ data = SessionData(id=normalized_id)
577
+ await self.storage.set(key, data.to_json(), self._session_ttl())
578
+ return SaSession(self.storage, key, data, self._session_ttl())
579
+
580
+ async def get_token_session(self, token: str | None = None) -> SaSession | None:
581
+ token = token or get_current_token()
582
+ if not token:
583
+ return None
584
+ info = await self._read_token_info(token)
585
+ if info is None or info.is_offline:
586
+ return None
587
+ key = self._token_session_key(token)
588
+ raw = await self.storage.get(key)
589
+ data = SessionData.from_json(raw) if raw else None
590
+ if data is None:
591
+ data = SessionData(id=token)
592
+ await self.storage.set(key, data.to_json(), info.timeout)
593
+ return SaSession(self.storage, key, data, info.timeout)
594
+
595
+ async def delete_session(self, login_id: Any) -> None:
596
+ normalized_id = self.normalize_login_id(login_id)
597
+ await self.storage.delete(self._session_key(normalized_id))
598
+
599
+ # ------------------------------------------------------------------ 权限 / 角色
600
+
601
+ async def get_permissions(self, login_id: Any) -> list[str]:
602
+ return await self._load_auth_list(login_id, "permission")
603
+
604
+ async def get_roles(self, login_id: Any) -> list[str]:
605
+ return await self._load_auth_list(login_id, "role")
606
+
607
+ async def _load_auth_list(self, login_id: Any, kind: str) -> list[str]:
608
+ normalized_id = self.normalize_login_id(login_id)
609
+ interface = self._manager.stp_interface
610
+ if interface is None:
611
+ key = self._permission_key(normalized_id) if kind == "permission" else self._role_key(
612
+ normalized_id
613
+ )
614
+ return self._decode_list(await self.storage.get(key))
615
+
616
+ cache_timeout = self.config.perm_cache_timeout
617
+ cache_key = self._permission_cache_key(normalized_id, kind)
618
+ if cache_timeout > 0:
619
+ cached = await self.storage.get(cache_key)
620
+ if cached is not None:
621
+ return self._decode_list(cached)
622
+
623
+ if kind == "permission":
624
+ values = await interface.get_permission_list(normalized_id, self.login_type)
625
+ else:
626
+ values = await interface.get_role_list(normalized_id, self.login_type)
627
+ values = [str(item) for item in values]
628
+ if cache_timeout > 0:
629
+ await self.storage.set(cache_key, json.dumps(values), cache_timeout)
630
+ return values
631
+
632
+ @staticmethod
633
+ def _decode_list(raw: str | None) -> list[str]:
634
+ if not raw:
635
+ return []
636
+ try:
637
+ values = json.loads(raw)
638
+ except ValueError:
639
+ return []
640
+ return [str(item) for item in values] if isinstance(values, list) else []
641
+
642
+ async def set_permissions(self, login_id: Any, permissions: list[str]) -> None:
643
+ await self._store_auth_list(login_id, "permission", permissions)
644
+
645
+ async def set_roles(self, login_id: Any, roles: list[str]) -> None:
646
+ await self._store_auth_list(login_id, "role", roles)
647
+
648
+ async def _store_auth_list(self, login_id: Any, kind: str, values: list[str]) -> None:
649
+ normalized_id = self.normalize_login_id(login_id)
650
+ key = (
651
+ self._permission_key(normalized_id)
652
+ if kind == "permission"
653
+ else self._role_key(normalized_id)
654
+ )
655
+ await self.storage.set(key, json.dumps([str(item) for item in values]))
656
+ # 配置了 StpInterface 时,写入必须让缓存立刻失效,否则改权限不会即时生效。
657
+ await self.storage.delete(self._permission_cache_key(normalized_id, kind))
658
+
659
+ async def add_permission(self, login_id: Any, permission: str) -> None:
660
+ current = await self.get_permissions(login_id)
661
+ if permission not in current:
662
+ await self.set_permissions(login_id, [*current, permission])
663
+
664
+ async def remove_permission(self, login_id: Any, permission: str) -> None:
665
+ current = await self.get_permissions(login_id)
666
+ await self.set_permissions(login_id, [item for item in current if item != permission])
667
+
668
+ async def clear_permissions(self, login_id: Any) -> None:
669
+ await self.set_permissions(login_id, [])
670
+
671
+ async def add_role(self, login_id: Any, role: str) -> None:
672
+ current = await self.get_roles(login_id)
673
+ if role not in current:
674
+ await self.set_roles(login_id, [*current, role])
675
+
676
+ async def remove_role(self, login_id: Any, role: str) -> None:
677
+ current = await self.get_roles(login_id)
678
+ await self.set_roles(login_id, [item for item in current if item != role])
679
+
680
+ async def clear_roles(self, login_id: Any) -> None:
681
+ await self.set_roles(login_id, [])
682
+
683
+ async def has_permission(self, login_id: Any, permission: str) -> bool:
684
+ return has_element(await self.get_permissions(login_id), permission)
685
+
686
+ async def has_permissions_and(self, login_id: Any, permissions: list[str]) -> bool:
687
+ return match_all(await self.get_permissions(login_id), permissions) is None
688
+
689
+ async def has_permissions_or(self, login_id: Any, permissions: list[str]) -> bool:
690
+ return match_any(await self.get_permissions(login_id), permissions)
691
+
692
+ async def has_role(self, login_id: Any, role: str) -> bool:
693
+ return has_element(await self.get_roles(login_id), role)
694
+
695
+ async def has_roles_and(self, login_id: Any, roles: list[str]) -> bool:
696
+ return match_all(await self.get_roles(login_id), roles) is None
697
+
698
+ async def has_roles_or(self, login_id: Any, roles: list[str]) -> bool:
699
+ return match_any(await self.get_roles(login_id), roles)
700
+
701
+ async def check_permission(
702
+ self,
703
+ login_id: Any,
704
+ permissions: str | list[str],
705
+ *,
706
+ mode: MatchMode = "OR",
707
+ ) -> None:
708
+ required = [permissions] if isinstance(permissions, str) else list(permissions)
709
+ if not required:
710
+ return
711
+ granted = await self.get_permissions(login_id)
712
+ await self._emit(
713
+ Event.PERMISSION_CHECK,
714
+ login_id=str(login_id),
715
+ required=required,
716
+ mode=mode,
717
+ )
718
+ if mode == "AND":
719
+ missing = match_all(granted, required)
720
+ if missing is not None:
721
+ raise NotPermissionException(missing, login_type=self.login_type)
722
+ return
723
+ if not match_any(granted, required):
724
+ raise NotPermissionException(" | ".join(required), login_type=self.login_type)
725
+
726
+ async def check_role(
727
+ self,
728
+ login_id: Any,
729
+ roles: str | list[str],
730
+ *,
731
+ mode: MatchMode = "OR",
732
+ ) -> None:
733
+ required = [roles] if isinstance(roles, str) else list(roles)
734
+ if not required:
735
+ return
736
+ granted = await self.get_roles(login_id)
737
+ await self._emit(Event.ROLE_CHECK, login_id=str(login_id), required=required, mode=mode)
738
+ if mode == "AND":
739
+ missing = match_all(granted, required)
740
+ if missing is not None:
741
+ raise NotRoleException(missing, login_type=self.login_type)
742
+ return
743
+ if not match_any(granted, required):
744
+ raise NotRoleException(" | ".join(required), login_type=self.login_type)
745
+
746
+ # ------------------------------------------------------------------ 封禁
747
+
748
+ async def disable(
749
+ self,
750
+ login_id: Any,
751
+ seconds: int,
752
+ *,
753
+ service: str = DEFAULT_DISABLE_SERVICE,
754
+ level: int = 1,
755
+ ) -> None:
756
+ """封禁账号的某项服务,``seconds=-1`` 表示永久封禁。"""
757
+ if level < 1:
758
+ raise SaTokenException("封禁等级必须大于等于 1")
759
+ normalized_id = self.normalize_login_id(login_id)
760
+ await self.storage.set(
761
+ self._disable_key(normalized_id, service),
762
+ str(level),
763
+ None if seconds < 0 else seconds,
764
+ )
765
+ await self._emit(
766
+ Event.DISABLE, login_id=normalized_id, service=service, level=level, seconds=seconds
767
+ )
768
+
769
+ async def untie(self, login_id: Any, *, service: str = DEFAULT_DISABLE_SERVICE) -> None:
770
+ normalized_id = self.normalize_login_id(login_id)
771
+ await self.storage.delete(self._disable_key(normalized_id, service))
772
+ await self._emit(Event.UNTIE, login_id=normalized_id, service=service)
773
+
774
+ async def get_disable_level(
775
+ self,
776
+ login_id: Any,
777
+ *,
778
+ service: str = DEFAULT_DISABLE_SERVICE,
779
+ ) -> int:
780
+ """返回封禁等级,未封禁时返回 0。"""
781
+ normalized_id = self.normalize_login_id(login_id)
782
+ raw = await self.storage.get(self._disable_key(normalized_id, service))
783
+ if raw is None:
784
+ return 0
785
+ try:
786
+ return int(raw)
787
+ except ValueError:
788
+ return 0
789
+
790
+ async def is_disable(
791
+ self,
792
+ login_id: Any,
793
+ *,
794
+ service: str = DEFAULT_DISABLE_SERVICE,
795
+ level: int = 1,
796
+ ) -> bool:
797
+ return await self.get_disable_level(login_id, service=service) >= level
798
+
799
+ async def get_disable_time(
800
+ self,
801
+ login_id: Any,
802
+ *,
803
+ service: str = DEFAULT_DISABLE_SERVICE,
804
+ ) -> int:
805
+ """剩余封禁秒数;-1 永久,-2 未封禁。"""
806
+ normalized_id = self.normalize_login_id(login_id)
807
+ return await self.storage.ttl(self._disable_key(normalized_id, service))
808
+
809
+ async def check_disable(
810
+ self,
811
+ login_id: Any,
812
+ *,
813
+ service: str = DEFAULT_DISABLE_SERVICE,
814
+ level: int = 1,
815
+ ) -> None:
816
+ normalized_id = self.normalize_login_id(login_id)
817
+ current_level = await self.get_disable_level(normalized_id, service=service)
818
+ if current_level >= level:
819
+ remaining = await self.get_disable_time(normalized_id, service=service)
820
+ raise DisableException(
821
+ normalized_id, service, current_level, remaining, login_type=self.login_type
822
+ )
823
+
824
+ # ------------------------------------------------------------------ 二级认证
825
+
826
+ async def open_safe(self, token: str, business: str, seconds: int) -> None:
827
+ """打开二级认证窗口,用于支付、改密等敏感操作。"""
828
+ await self.storage.set(self._safe_key(token, business), str(now_ms()), seconds)
829
+
830
+ async def is_safe(self, token: str | None, business: str) -> bool:
831
+ token = token or get_current_token()
832
+ if not token:
833
+ return False
834
+ return await self.storage.exists(self._safe_key(token, business))
835
+
836
+ async def check_safe(self, token: str | None, business: str) -> None:
837
+ if not await self.is_safe(token, business):
838
+ raise NotSafeException(business, login_type=self.login_type)
839
+
840
+ async def close_safe(self, token: str, business: str) -> None:
841
+ await self.storage.delete(self._safe_key(token, business))
842
+
843
+ # ------------------------------------------------------------------ 查询
844
+
845
+ async def get_terminal_list(
846
+ self,
847
+ login_id: Any,
848
+ *,
849
+ device: str | None = None,
850
+ ) -> list[TerminalInfo]:
851
+ session = await self.get_session(login_id, create=False)
852
+ if session is None:
853
+ return []
854
+ return [
855
+ terminal
856
+ for terminal in session.raw.terminal_list
857
+ if device is None or terminal.device == device
858
+ ]
859
+
860
+ async def get_token_value_list_by_login_id(
861
+ self,
862
+ login_id: Any,
863
+ *,
864
+ device: str | None = None,
865
+ ) -> list[str]:
866
+ return [
867
+ terminal.token for terminal in await self.get_terminal_list(login_id, device=device)
868
+ ]
869
+
870
+ async def search_token_value(
871
+ self,
872
+ keyword: str = "",
873
+ *,
874
+ start: int = 0,
875
+ size: int = 100,
876
+ ) -> list[str]:
877
+ """管理端用:按关键字扫描在线 token。
878
+
879
+ 依赖存储的 ``scan``,在大规模 Redis 上属于重操作,不要放进请求热路径。
880
+ """
881
+ prefix = f"{self.config.key_prefix(self.login_type)}token:"
882
+ pattern = f"{prefix}*{keyword}*" if keyword else f"{prefix}*"
883
+ found: list[str] = []
884
+ cursor: str | None = None
885
+ while True:
886
+ cursor, keys = await self.storage.scan(pattern, cursor, 200)
887
+ found.extend(key[len(prefix) :] for key in keys)
888
+ if cursor is None:
889
+ break
890
+ found.sort()
891
+ return found[start : start + size] if size > 0 else found[start:]
892
+
893
+ async def search_session(
894
+ self,
895
+ keyword: str = "",
896
+ *,
897
+ start: int = 0,
898
+ size: int = 100,
899
+ ) -> list[str]:
900
+ """管理端用:扫描 Account-Session ID。"""
901
+ prefix = f"{self.config.key_prefix(self.login_type)}session:"
902
+ pattern = f"{prefix}*{keyword}*" if keyword else f"{prefix}*"
903
+ found: list[str] = []
904
+ cursor: str | None = None
905
+ while True:
906
+ cursor, keys = await self.storage.scan(pattern, cursor, 200)
907
+ found.extend(key[len(prefix) :] for key in keys)
908
+ if cursor is None:
909
+ break
910
+ found.sort()
911
+ return found[start : start + size] if size > 0 else found[start:]