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.
- sa_token/__init__.py +89 -0
- sa_token/adapter/__init__.py +24 -0
- sa_token/adapter/http.py +71 -0
- sa_token/adapter/path.py +163 -0
- sa_token/adapter/pipeline.py +97 -0
- sa_token/config.py +130 -0
- sa_token/context.py +63 -0
- sa_token/exception.py +143 -0
- sa_token/integration/__init__.py +10 -0
- sa_token/integration/django.py +131 -0
- sa_token/integration/fastapi.py +315 -0
- sa_token/integration/fastapi_oauth2.py +136 -0
- sa_token/integration/flask.py +191 -0
- sa_token/integration/starlette.py +227 -0
- sa_token/listener.py +100 -0
- sa_token/manager.py +244 -0
- sa_token/model.py +145 -0
- sa_token/oauth2/__init__.py +19 -0
- sa_token/oauth2/model.py +122 -0
- sa_token/oauth2/server.py +361 -0
- sa_token/online/__init__.py +292 -0
- sa_token/permission.py +67 -0
- sa_token/py.typed +0 -0
- sa_token/security/__init__.py +14 -0
- sa_token/security/nonce.py +93 -0
- sa_token/security/refresh.py +300 -0
- sa_token/security/temp_token.py +114 -0
- sa_token/session.py +96 -0
- sa_token/sso/__init__.py +217 -0
- sa_token/storage/__init__.py +22 -0
- sa_token/storage/base.py +66 -0
- sa_token/storage/memory.py +154 -0
- sa_token/storage/redis.py +136 -0
- sa_token/stp_interface.py +20 -0
- sa_token/stp_logic.py +911 -0
- sa_token/stp_util.py +367 -0
- sa_token/strategy/__init__.py +77 -0
- sa_token/strategy/base.py +22 -0
- sa_token/strategy/builtin.py +99 -0
- sa_token/strategy/jwt.py +72 -0
- sa_token/sync.py +268 -0
- sa_token/token_io.py +66 -0
- sa_token_python_core-0.1.1.dist-info/METADATA +756 -0
- sa_token_python_core-0.1.1.dist-info/RECORD +46 -0
- sa_token_python_core-0.1.1.dist-info/WHEEL +4 -0
- sa_token_python_core-0.1.1.dist-info/licenses/LICENSE +201 -0
sa_token/sso/__init__.py
ADDED
|
@@ -0,0 +1,217 @@
|
|
|
1
|
+
"""SSO 单点登录:ticket 换登录态 + 统一登出。
|
|
2
|
+
|
|
3
|
+
Server 与 Client 都只依赖核心存储与 ``StpLogic``,HTTP 跳转由使用方实现,
|
|
4
|
+
因此整条流程可以在单测里不起服务地跑通。
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import hashlib
|
|
10
|
+
import hmac
|
|
11
|
+
import secrets
|
|
12
|
+
from dataclasses import dataclass, field
|
|
13
|
+
from typing import TYPE_CHECKING
|
|
14
|
+
from urllib.parse import urlencode
|
|
15
|
+
|
|
16
|
+
from ..exception import SaTokenException
|
|
17
|
+
|
|
18
|
+
if TYPE_CHECKING: # pragma: no cover - 仅供类型检查
|
|
19
|
+
from ..manager import SaTokenManager
|
|
20
|
+
|
|
21
|
+
__all__ = ["SsoConfig", "SsoError", "SsoTicket", "SsoServer", "SsoClient"]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class SsoError(SaTokenException):
|
|
25
|
+
http_status = 400
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass
|
|
29
|
+
class SsoConfig:
|
|
30
|
+
server_url: str = ""
|
|
31
|
+
"""认证中心的登录页地址。"""
|
|
32
|
+
|
|
33
|
+
ticket_timeout: int = 300
|
|
34
|
+
"""ticket 有效期,应当很短:它只用于一次跳转。"""
|
|
35
|
+
|
|
36
|
+
allowed_services: list[str] = field(default_factory=list)
|
|
37
|
+
"""允许接入的 service 白名单,空列表表示不限制(仅建议内网使用)。"""
|
|
38
|
+
|
|
39
|
+
secret_key: str | None = None
|
|
40
|
+
"""跨进程 ticket 校验的 HMAC 密钥;生产环境建议配置至少 32 字节。"""
|
|
41
|
+
|
|
42
|
+
def is_allowed_service(self, service: str) -> bool:
|
|
43
|
+
return not self.allowed_services or service in self.allowed_services
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@dataclass(frozen=True)
|
|
47
|
+
class SsoTicket:
|
|
48
|
+
ticket: str
|
|
49
|
+
service: str
|
|
50
|
+
signature: str | None = None
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class SsoServer:
|
|
54
|
+
"""认证中心。"""
|
|
55
|
+
|
|
56
|
+
def __init__(self, manager: SaTokenManager, config: SsoConfig | None = None) -> None:
|
|
57
|
+
self._manager = manager
|
|
58
|
+
self.config = config or SsoConfig()
|
|
59
|
+
|
|
60
|
+
def _key(self, suffix: str, identifier: str) -> str:
|
|
61
|
+
return self._manager.config.make_key("sso", suffix, identifier)
|
|
62
|
+
|
|
63
|
+
@property
|
|
64
|
+
def _storage(self):
|
|
65
|
+
return self._manager.storage
|
|
66
|
+
|
|
67
|
+
async def create_ticket(self, login_id: str, service: str) -> str:
|
|
68
|
+
"""用户在认证中心登录后,为目标应用签发一次性 ticket。"""
|
|
69
|
+
if not self.config.is_allowed_service(service):
|
|
70
|
+
raise SsoError(f"service 未在白名单中:{service}")
|
|
71
|
+
ticket = secrets.token_urlsafe(24)
|
|
72
|
+
# ticket 绑定 service,避免 A 应用拿到的 ticket 被拿去登录 B 应用。
|
|
73
|
+
await self._storage.set(
|
|
74
|
+
self._key("ticket", ticket),
|
|
75
|
+
f"{login_id}\n{service}",
|
|
76
|
+
self.config.ticket_timeout,
|
|
77
|
+
)
|
|
78
|
+
await self._register_service(login_id, service)
|
|
79
|
+
return ticket
|
|
80
|
+
|
|
81
|
+
async def create_signed_ticket(self, login_id: str, service: str) -> SsoTicket:
|
|
82
|
+
"""签发 ticket,并在配置密钥时附带 HMAC 签名。"""
|
|
83
|
+
ticket = await self.create_ticket(login_id, service)
|
|
84
|
+
return SsoTicket(
|
|
85
|
+
ticket=ticket,
|
|
86
|
+
service=service,
|
|
87
|
+
signature=self.sign_ticket(ticket, service),
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
def sign_ticket(self, ticket: str, service: str) -> str | None:
|
|
91
|
+
if not self.config.secret_key:
|
|
92
|
+
return None
|
|
93
|
+
payload = f"{ticket}\n{service}".encode()
|
|
94
|
+
return hmac.new(
|
|
95
|
+
self.config.secret_key.encode(),
|
|
96
|
+
payload,
|
|
97
|
+
hashlib.sha256,
|
|
98
|
+
).hexdigest()
|
|
99
|
+
|
|
100
|
+
def verify_ticket_signature(
|
|
101
|
+
self, ticket: str, service: str, signature: str | None
|
|
102
|
+
) -> None:
|
|
103
|
+
expected = self.sign_ticket(ticket, service)
|
|
104
|
+
if expected is None:
|
|
105
|
+
return
|
|
106
|
+
if signature is None or not hmac.compare_digest(expected, signature):
|
|
107
|
+
raise SsoError("ticket 签名无效")
|
|
108
|
+
|
|
109
|
+
async def validate_ticket(
|
|
110
|
+
self,
|
|
111
|
+
ticket: str,
|
|
112
|
+
service: str,
|
|
113
|
+
*,
|
|
114
|
+
signature: str | None = None,
|
|
115
|
+
) -> str:
|
|
116
|
+
"""校验并消费 ticket,返回 ``login_id``。
|
|
117
|
+
|
|
118
|
+
读取即删除:ticket 必须一次性,否则中间人可以重放它换取会话。
|
|
119
|
+
"""
|
|
120
|
+
self.verify_ticket_signature(ticket, service, signature)
|
|
121
|
+
key = self._key("ticket", ticket)
|
|
122
|
+
raw = await self._storage.get(key)
|
|
123
|
+
if raw is None:
|
|
124
|
+
raise SsoError("ticket 无效或已过期")
|
|
125
|
+
if not await self._storage.compare_and_delete(key, raw):
|
|
126
|
+
raise SsoError("ticket 已被使用")
|
|
127
|
+
login_id, bound_service = raw.split("\n", 1)
|
|
128
|
+
if bound_service != service:
|
|
129
|
+
raise SsoError("ticket 与 service 不匹配")
|
|
130
|
+
return login_id
|
|
131
|
+
|
|
132
|
+
async def _register_service(self, login_id: str, service: str) -> None:
|
|
133
|
+
"""记录该用户在哪些应用登录过,统一登出时需要逐个通知。"""
|
|
134
|
+
key = self._key("services", login_id)
|
|
135
|
+
for _ in range(12):
|
|
136
|
+
raw = await self._storage.get(key)
|
|
137
|
+
services = raw.split("\n") if raw else []
|
|
138
|
+
if service in services:
|
|
139
|
+
return
|
|
140
|
+
services.append(service)
|
|
141
|
+
new_raw = "\n".join(services)
|
|
142
|
+
if raw is None:
|
|
143
|
+
if await self._storage.set_if_absent(key, new_raw):
|
|
144
|
+
return
|
|
145
|
+
elif await self._storage.compare_and_set(key, raw, new_raw):
|
|
146
|
+
return
|
|
147
|
+
raise SsoError("并发登记 SSO service 失败,请稍后重试")
|
|
148
|
+
|
|
149
|
+
async def get_registered_services(self, login_id: str) -> list[str]:
|
|
150
|
+
raw = await self._storage.get(self._key("services", login_id))
|
|
151
|
+
return raw.split("\n") if raw else []
|
|
152
|
+
|
|
153
|
+
async def logout(self, login_id: str) -> list[str]:
|
|
154
|
+
"""统一登出:清掉认证中心会话,返回需要通知的客户端列表。
|
|
155
|
+
|
|
156
|
+
实际的 HTTP 回调由使用方发起——通知方式(同步 HTTP、消息队列)
|
|
157
|
+
取决于部署形态,不应该由库来替用户决定。
|
|
158
|
+
"""
|
|
159
|
+
services = await self.get_registered_services(login_id)
|
|
160
|
+
await self._manager.stp().logout(login_id)
|
|
161
|
+
await self._storage.delete(self._key("services", login_id))
|
|
162
|
+
return services
|
|
163
|
+
|
|
164
|
+
def build_login_url(self, service: str, *, redirect: str | None = None) -> str:
|
|
165
|
+
params = {"service": service}
|
|
166
|
+
if redirect:
|
|
167
|
+
params["redirect"] = redirect
|
|
168
|
+
separator = "&" if "?" in self.config.server_url else "?"
|
|
169
|
+
return f"{self.config.server_url}{separator}{urlencode(params)}"
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
class SsoClient:
|
|
173
|
+
"""接入方应用。"""
|
|
174
|
+
|
|
175
|
+
def __init__(
|
|
176
|
+
self,
|
|
177
|
+
manager: SaTokenManager,
|
|
178
|
+
*,
|
|
179
|
+
server_url: str,
|
|
180
|
+
service: str,
|
|
181
|
+
login_type: str = "login",
|
|
182
|
+
) -> None:
|
|
183
|
+
self._manager = manager
|
|
184
|
+
self.server_url = server_url
|
|
185
|
+
self.service = service
|
|
186
|
+
self.login_type = login_type
|
|
187
|
+
|
|
188
|
+
def get_login_url(self, redirect: str | None = None) -> str:
|
|
189
|
+
params = {"service": self.service}
|
|
190
|
+
if redirect:
|
|
191
|
+
params["redirect"] = redirect
|
|
192
|
+
separator = "&" if "?" in self.server_url else "?"
|
|
193
|
+
return f"{self.server_url}{separator}{urlencode(params)}"
|
|
194
|
+
|
|
195
|
+
async def login_by_ticket(
|
|
196
|
+
self,
|
|
197
|
+
server: SsoServer,
|
|
198
|
+
ticket: str,
|
|
199
|
+
*,
|
|
200
|
+
device: str | None = None,
|
|
201
|
+
signature: str | None = None,
|
|
202
|
+
) -> str:
|
|
203
|
+
"""同进程部署时直接校验 ticket 并建立本地登录态。"""
|
|
204
|
+
login_id = await server.validate_ticket(
|
|
205
|
+
ticket,
|
|
206
|
+
self.service,
|
|
207
|
+
signature=signature,
|
|
208
|
+
)
|
|
209
|
+
return await self._manager.stp(self.login_type).login(login_id, device=device)
|
|
210
|
+
|
|
211
|
+
async def login_by_login_id(self, login_id: str, *, device: str | None = None) -> str:
|
|
212
|
+
"""跨进程部署时,ticket 已由 HTTP 调用认证中心换成 login_id。"""
|
|
213
|
+
return await self._manager.stp(self.login_type).login(login_id, device=device)
|
|
214
|
+
|
|
215
|
+
async def handle_logout(self, login_id: str) -> None:
|
|
216
|
+
"""收到认证中心的登出通知后,清掉本地会话。"""
|
|
217
|
+
await self._manager.stp(self.login_type).logout(login_id)
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
"""存储层:契约 + 内置实现。
|
|
2
|
+
|
|
3
|
+
``RedisStorage`` 依赖可选的 redis 包,因此这里按需惰性导入,
|
|
4
|
+
保证只装了核心包的用户 ``import sa_token.storage`` 不会报错。
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from .base import TTL_NEVER_EXPIRE, SaStorage
|
|
12
|
+
from .memory import MemoryStorage
|
|
13
|
+
|
|
14
|
+
__all__ = ["SaStorage", "MemoryStorage", "RedisStorage", "TTL_NEVER_EXPIRE"]
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def __getattr__(name: str) -> Any:
|
|
18
|
+
if name == "RedisStorage":
|
|
19
|
+
from .redis import RedisStorage
|
|
20
|
+
|
|
21
|
+
return RedisStorage
|
|
22
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
sa_token/storage/base.py
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
"""存储契约。
|
|
2
|
+
|
|
3
|
+
核心层只依赖这套接口,因此内存、Redis、数据库甚至自研存储都能平替。
|
|
4
|
+
所有值都是字符串(JSON),这样 Redis 里的数据可以被其它语言的 sa-token 读写。
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Protocol, runtime_checkable
|
|
10
|
+
|
|
11
|
+
__all__ = ["SaStorage", "TTL_NEVER_EXPIRE"]
|
|
12
|
+
|
|
13
|
+
#: ttl 传入该值(或 None)表示永不过期。
|
|
14
|
+
TTL_NEVER_EXPIRE = -1
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@runtime_checkable
|
|
18
|
+
class SaStorage(Protocol):
|
|
19
|
+
"""键值存储接口。
|
|
20
|
+
|
|
21
|
+
``ttl`` 单位为秒;``None`` 与 ``-1`` 等价,均表示永不过期。
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
async def get(self, key: str) -> str | None: ...
|
|
25
|
+
|
|
26
|
+
async def set(self, key: str, value: str, ttl: int | None = None) -> None: ...
|
|
27
|
+
|
|
28
|
+
async def delete(self, key: str) -> None: ...
|
|
29
|
+
|
|
30
|
+
async def exists(self, key: str) -> bool: ...
|
|
31
|
+
|
|
32
|
+
async def expire(self, key: str, ttl: int | None) -> bool:
|
|
33
|
+
"""更新 TTL,键不存在时返回 False。"""
|
|
34
|
+
...
|
|
35
|
+
|
|
36
|
+
async def ttl(self, key: str) -> int:
|
|
37
|
+
"""返回剩余秒数;-1 表示永不过期,-2 表示键不存在。"""
|
|
38
|
+
...
|
|
39
|
+
|
|
40
|
+
async def set_if_absent(self, key: str, value: str, ttl: int | None = None) -> bool:
|
|
41
|
+
"""键不存在时写入,返回是否写入成功。用于一次性 token 的原子占位。"""
|
|
42
|
+
...
|
|
43
|
+
|
|
44
|
+
async def compare_and_set(
|
|
45
|
+
self,
|
|
46
|
+
key: str,
|
|
47
|
+
expected: str,
|
|
48
|
+
new_value: str,
|
|
49
|
+
ttl: int | None = None,
|
|
50
|
+
) -> bool:
|
|
51
|
+
"""当前值等于 expected 时才写入,用于并发下的安全更新。"""
|
|
52
|
+
...
|
|
53
|
+
|
|
54
|
+
async def compare_and_delete(self, key: str, expected: str) -> bool: ...
|
|
55
|
+
|
|
56
|
+
async def scan(self, pattern: str, cursor: str | None = None, count: int = 100) -> tuple[
|
|
57
|
+
str | None, list[str]
|
|
58
|
+
]:
|
|
59
|
+
"""按 glob 模式游标扫描键,返回 ``(下一个游标, 键列表)``;游标为 None 表示结束。"""
|
|
60
|
+
...
|
|
61
|
+
|
|
62
|
+
async def clear(self) -> None:
|
|
63
|
+
"""清空全部数据,仅用于测试与开发。"""
|
|
64
|
+
...
|
|
65
|
+
|
|
66
|
+
async def close(self) -> None: ...
|
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
"""内存存储:开发、测试与单机场景的默认实现。
|
|
2
|
+
|
|
3
|
+
过期采用惰性清理 + 可选的后台清扫,避免为了精确过期而引入常驻线程。
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
import asyncio
|
|
9
|
+
import fnmatch
|
|
10
|
+
import time
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
|
|
13
|
+
from .base import TTL_NEVER_EXPIRE
|
|
14
|
+
|
|
15
|
+
__all__ = ["MemoryStorage"]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass
|
|
19
|
+
class _Entry:
|
|
20
|
+
value: str
|
|
21
|
+
expire_at: float | None
|
|
22
|
+
|
|
23
|
+
def is_expired(self, now: float) -> bool:
|
|
24
|
+
return self.expire_at is not None and self.expire_at <= now
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _to_expire_at(ttl: int | None) -> float | None:
|
|
28
|
+
if ttl is None or ttl == TTL_NEVER_EXPIRE:
|
|
29
|
+
return None
|
|
30
|
+
if ttl <= 0:
|
|
31
|
+
# 传入 0 或负数(-1 以外)视为立即过期,避免写入永不清理的脏数据。
|
|
32
|
+
return time.monotonic()
|
|
33
|
+
return time.monotonic() + ttl
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class MemoryStorage:
|
|
37
|
+
"""进程内存储。
|
|
38
|
+
|
|
39
|
+
注意:数据随进程退出而丢失,多进程部署(gunicorn 多 worker)下各进程互相看不见
|
|
40
|
+
对方的登录态,生产环境请改用 ``RedisStorage``。
|
|
41
|
+
"""
|
|
42
|
+
|
|
43
|
+
def __init__(self, *, cleanup_interval: float = 60.0) -> None:
|
|
44
|
+
self._data: dict[str, _Entry] = {}
|
|
45
|
+
self._lock = asyncio.Lock()
|
|
46
|
+
self._cleanup_interval = cleanup_interval
|
|
47
|
+
self._last_cleanup = time.monotonic()
|
|
48
|
+
|
|
49
|
+
async def get(self, key: str) -> str | None:
|
|
50
|
+
async with self._lock:
|
|
51
|
+
return self._get_unlocked(key)
|
|
52
|
+
|
|
53
|
+
async def set(self, key: str, value: str, ttl: int | None = None) -> None:
|
|
54
|
+
async with self._lock:
|
|
55
|
+
self._data[key] = _Entry(value, _to_expire_at(ttl))
|
|
56
|
+
self._maybe_cleanup()
|
|
57
|
+
|
|
58
|
+
async def delete(self, key: str) -> None:
|
|
59
|
+
async with self._lock:
|
|
60
|
+
self._data.pop(key, None)
|
|
61
|
+
|
|
62
|
+
async def exists(self, key: str) -> bool:
|
|
63
|
+
async with self._lock:
|
|
64
|
+
return self._get_unlocked(key) is not None
|
|
65
|
+
|
|
66
|
+
async def expire(self, key: str, ttl: int | None) -> bool:
|
|
67
|
+
async with self._lock:
|
|
68
|
+
entry = self._data.get(key)
|
|
69
|
+
if entry is None or entry.is_expired(time.monotonic()):
|
|
70
|
+
self._data.pop(key, None)
|
|
71
|
+
return False
|
|
72
|
+
entry.expire_at = _to_expire_at(ttl)
|
|
73
|
+
return True
|
|
74
|
+
|
|
75
|
+
async def ttl(self, key: str) -> int:
|
|
76
|
+
async with self._lock:
|
|
77
|
+
now = time.monotonic()
|
|
78
|
+
entry = self._data.get(key)
|
|
79
|
+
if entry is None or entry.is_expired(now):
|
|
80
|
+
self._data.pop(key, None)
|
|
81
|
+
return -2
|
|
82
|
+
if entry.expire_at is None:
|
|
83
|
+
return TTL_NEVER_EXPIRE
|
|
84
|
+
return max(0, int(round(entry.expire_at - now)))
|
|
85
|
+
|
|
86
|
+
async def set_if_absent(self, key: str, value: str, ttl: int | None = None) -> bool:
|
|
87
|
+
async with self._lock:
|
|
88
|
+
if self._get_unlocked(key) is not None:
|
|
89
|
+
return False
|
|
90
|
+
self._data[key] = _Entry(value, _to_expire_at(ttl))
|
|
91
|
+
return True
|
|
92
|
+
|
|
93
|
+
async def compare_and_set(
|
|
94
|
+
self,
|
|
95
|
+
key: str,
|
|
96
|
+
expected: str,
|
|
97
|
+
new_value: str,
|
|
98
|
+
ttl: int | None = None,
|
|
99
|
+
) -> bool:
|
|
100
|
+
async with self._lock:
|
|
101
|
+
if self._get_unlocked(key) != expected:
|
|
102
|
+
return False
|
|
103
|
+
self._data[key] = _Entry(new_value, _to_expire_at(ttl))
|
|
104
|
+
return True
|
|
105
|
+
|
|
106
|
+
async def compare_and_delete(self, key: str, expected: str) -> bool:
|
|
107
|
+
async with self._lock:
|
|
108
|
+
if self._get_unlocked(key) != expected:
|
|
109
|
+
return False
|
|
110
|
+
self._data.pop(key, None)
|
|
111
|
+
return True
|
|
112
|
+
|
|
113
|
+
async def scan(
|
|
114
|
+
self,
|
|
115
|
+
pattern: str,
|
|
116
|
+
cursor: str | None = None,
|
|
117
|
+
count: int = 100,
|
|
118
|
+
) -> tuple[str | None, list[str]]:
|
|
119
|
+
async with self._lock:
|
|
120
|
+
now = time.monotonic()
|
|
121
|
+
keys = sorted(
|
|
122
|
+
key
|
|
123
|
+
for key, entry in self._data.items()
|
|
124
|
+
if not entry.is_expired(now) and fnmatch.fnmatchcase(key, pattern)
|
|
125
|
+
)
|
|
126
|
+
start = int(cursor) if cursor else 0
|
|
127
|
+
page = keys[start : start + count]
|
|
128
|
+
next_cursor = str(start + count) if start + count < len(keys) else None
|
|
129
|
+
return next_cursor, page
|
|
130
|
+
|
|
131
|
+
async def clear(self) -> None:
|
|
132
|
+
async with self._lock:
|
|
133
|
+
self._data.clear()
|
|
134
|
+
|
|
135
|
+
async def close(self) -> None:
|
|
136
|
+
await self.clear()
|
|
137
|
+
|
|
138
|
+
def _get_unlocked(self, key: str) -> str | None:
|
|
139
|
+
entry = self._data.get(key)
|
|
140
|
+
if entry is None:
|
|
141
|
+
return None
|
|
142
|
+
if entry.is_expired(time.monotonic()):
|
|
143
|
+
self._data.pop(key, None)
|
|
144
|
+
return None
|
|
145
|
+
return entry.value
|
|
146
|
+
|
|
147
|
+
def _maybe_cleanup(self) -> None:
|
|
148
|
+
now = time.monotonic()
|
|
149
|
+
if now - self._last_cleanup < self._cleanup_interval:
|
|
150
|
+
return
|
|
151
|
+
self._last_cleanup = now
|
|
152
|
+
expired = [key for key, entry in self._data.items() if entry.is_expired(now)]
|
|
153
|
+
for key in expired:
|
|
154
|
+
self._data.pop(key, None)
|
|
@@ -0,0 +1,136 @@
|
|
|
1
|
+
"""Redis 存储:生产环境实现。
|
|
2
|
+
|
|
3
|
+
需要额外安装:``pip install "sa-token-python-core[redis]"``。
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
from typing import TYPE_CHECKING, Any
|
|
9
|
+
|
|
10
|
+
from .base import TTL_NEVER_EXPIRE
|
|
11
|
+
|
|
12
|
+
if TYPE_CHECKING: # pragma: no cover - 仅供类型检查
|
|
13
|
+
from redis.asyncio import Redis
|
|
14
|
+
|
|
15
|
+
__all__ = ["RedisStorage"]
|
|
16
|
+
|
|
17
|
+
#: 值相等才删除,避免误删被其它请求刷新过的键。
|
|
18
|
+
_COMPARE_AND_DELETE = """
|
|
19
|
+
if redis.call('GET', KEYS[1]) == ARGV[1] then
|
|
20
|
+
return redis.call('DEL', KEYS[1])
|
|
21
|
+
end
|
|
22
|
+
return 0
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
#: 值相等才更新;ARGV[3] 为 -1 时保持 key 永不过期。
|
|
26
|
+
_COMPARE_AND_SET = """
|
|
27
|
+
if redis.call('GET', KEYS[1]) ~= ARGV[1] then
|
|
28
|
+
return 0
|
|
29
|
+
end
|
|
30
|
+
if tonumber(ARGV[3]) < 0 then
|
|
31
|
+
redis.call('SET', KEYS[1], ARGV[2])
|
|
32
|
+
else
|
|
33
|
+
redis.call('SET', KEYS[1], ARGV[2], 'EX', tonumber(ARGV[3]))
|
|
34
|
+
end
|
|
35
|
+
return 1
|
|
36
|
+
"""
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _normalize_ttl(ttl: int | None) -> int | None:
|
|
40
|
+
"""把 ``None`` / ``-1`` 统一成「不设置过期」。"""
|
|
41
|
+
if ttl is None or ttl == TTL_NEVER_EXPIRE:
|
|
42
|
+
return None
|
|
43
|
+
return max(1, ttl)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class RedisStorage:
|
|
47
|
+
"""基于 ``redis.asyncio`` 的存储实现。
|
|
48
|
+
|
|
49
|
+
要求客户端以字符串模式解码(``decode_responses=True``);使用 :meth:`from_url`
|
|
50
|
+
构造时会自动设置。
|
|
51
|
+
"""
|
|
52
|
+
|
|
53
|
+
def __init__(self, client: Redis) -> None:
|
|
54
|
+
self._redis = client
|
|
55
|
+
self._compare_and_delete = client.register_script(_COMPARE_AND_DELETE)
|
|
56
|
+
self._compare_and_set = client.register_script(_COMPARE_AND_SET)
|
|
57
|
+
|
|
58
|
+
@classmethod
|
|
59
|
+
def from_url(cls, url: str, **kwargs: Any) -> RedisStorage:
|
|
60
|
+
"""``RedisStorage.from_url("redis://localhost:6379/0")``"""
|
|
61
|
+
try:
|
|
62
|
+
from redis.asyncio import Redis
|
|
63
|
+
except ImportError as exc: # pragma: no cover - 依赖缺失路径
|
|
64
|
+
raise ImportError(
|
|
65
|
+
'RedisStorage 需要 redis 依赖,请执行:pip install "sa-token-python-core[redis]"'
|
|
66
|
+
) from exc
|
|
67
|
+
kwargs.setdefault("decode_responses", True)
|
|
68
|
+
return cls(Redis.from_url(url, **kwargs))
|
|
69
|
+
|
|
70
|
+
async def get(self, key: str) -> str | None:
|
|
71
|
+
return await self._redis.get(key)
|
|
72
|
+
|
|
73
|
+
async def set(self, key: str, value: str, ttl: int | None = None) -> None:
|
|
74
|
+
seconds = _normalize_ttl(ttl)
|
|
75
|
+
if seconds is None:
|
|
76
|
+
await self._redis.set(key, value)
|
|
77
|
+
else:
|
|
78
|
+
await self._redis.set(key, value, ex=seconds)
|
|
79
|
+
|
|
80
|
+
async def delete(self, key: str) -> None:
|
|
81
|
+
await self._redis.delete(key)
|
|
82
|
+
|
|
83
|
+
async def exists(self, key: str) -> bool:
|
|
84
|
+
return bool(await self._redis.exists(key))
|
|
85
|
+
|
|
86
|
+
async def expire(self, key: str, ttl: int | None) -> bool:
|
|
87
|
+
seconds = _normalize_ttl(ttl)
|
|
88
|
+
if seconds is None:
|
|
89
|
+
return bool(await self._redis.persist(key))
|
|
90
|
+
return bool(await self._redis.expire(key, seconds))
|
|
91
|
+
|
|
92
|
+
async def ttl(self, key: str) -> int:
|
|
93
|
+
return int(await self._redis.ttl(key))
|
|
94
|
+
|
|
95
|
+
async def set_if_absent(self, key: str, value: str, ttl: int | None = None) -> bool:
|
|
96
|
+
seconds = _normalize_ttl(ttl)
|
|
97
|
+
if seconds is None:
|
|
98
|
+
return bool(await self._redis.set(key, value, nx=True))
|
|
99
|
+
return bool(await self._redis.set(key, value, ex=seconds, nx=True))
|
|
100
|
+
|
|
101
|
+
async def compare_and_set(
|
|
102
|
+
self,
|
|
103
|
+
key: str,
|
|
104
|
+
expected: str,
|
|
105
|
+
new_value: str,
|
|
106
|
+
ttl: int | None = None,
|
|
107
|
+
) -> bool:
|
|
108
|
+
seconds = _normalize_ttl(ttl)
|
|
109
|
+
result = await self._compare_and_set(
|
|
110
|
+
keys=[key],
|
|
111
|
+
args=[expected, new_value, TTL_NEVER_EXPIRE if seconds is None else seconds],
|
|
112
|
+
)
|
|
113
|
+
return bool(result)
|
|
114
|
+
|
|
115
|
+
async def compare_and_delete(self, key: str, expected: str) -> bool:
|
|
116
|
+
return bool(await self._compare_and_delete(keys=[key], args=[expected]))
|
|
117
|
+
|
|
118
|
+
async def scan(
|
|
119
|
+
self,
|
|
120
|
+
pattern: str,
|
|
121
|
+
cursor: str | None = None,
|
|
122
|
+
count: int = 100,
|
|
123
|
+
) -> tuple[str | None, list[str]]:
|
|
124
|
+
next_cursor, keys = await self._redis.scan(
|
|
125
|
+
cursor=int(cursor) if cursor else 0,
|
|
126
|
+
match=pattern,
|
|
127
|
+
count=count,
|
|
128
|
+
)
|
|
129
|
+
return (str(next_cursor) if next_cursor else None), list(keys)
|
|
130
|
+
|
|
131
|
+
async def clear(self) -> None:
|
|
132
|
+
"""清空当前库,仅供测试使用。"""
|
|
133
|
+
await self._redis.flushdb()
|
|
134
|
+
|
|
135
|
+
async def close(self) -> None:
|
|
136
|
+
await self._redis.aclose()
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""权限数据源。
|
|
2
|
+
|
|
3
|
+
默认实现从存储里读 ``set_permissions`` 写入的数据,适合中小项目与网关;
|
|
4
|
+
已有 RBAC 表的项目应实现本接口,直接从数据库读,避免两份权限数据不同步。
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Protocol, runtime_checkable
|
|
10
|
+
|
|
11
|
+
__all__ = ["StpInterface"]
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@runtime_checkable
|
|
15
|
+
class StpInterface(Protocol):
|
|
16
|
+
"""业务侧权限 / 角色数据源。"""
|
|
17
|
+
|
|
18
|
+
async def get_permission_list(self, login_id: str, login_type: str) -> list[str]: ...
|
|
19
|
+
|
|
20
|
+
async def get_role_list(self, login_id: str, login_type: str) -> list[str]: ...
|