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
@@ -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}")
@@ -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]: ...