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,136 @@
1
+ """OAuth2 的 FastAPI HTTP 端点适配。
2
+
3
+ 协议引擎仍在 :mod:`sa_token.oauth2`;这里仅处理 HTTP 表单/JSON 和标准错误响应。
4
+ 授权页的 UI 与「用户是否同意」属于业务,本模块只提供授权成功后的跳转函数。
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from typing import Any
10
+ from urllib.parse import parse_qsl, urlencode
11
+
12
+ try:
13
+ from fastapi import APIRouter, Request
14
+ from fastapi.responses import JSONResponse, RedirectResponse
15
+ except ImportError as exc: # pragma: no cover
16
+ raise ImportError(
17
+ '该模块需要 fastapi,请执行:pip install "sa-token-python-core[fastapi]"'
18
+ ) from exc
19
+
20
+ from ..oauth2 import OAuth2Error, OAuth2Server
21
+
22
+ __all__ = ["create_oauth2_router", "authorization_redirect"]
23
+
24
+
25
+ async def _read_parameters(request: Request) -> dict[str, Any]:
26
+ content_type = request.headers.get("content-type", "")
27
+ if "application/json" in content_type:
28
+ payload = await request.json()
29
+ return payload if isinstance(payload, dict) else {}
30
+ body = (await request.body()).decode("utf-8")
31
+ return dict(parse_qsl(body, keep_blank_values=True))
32
+
33
+
34
+ def _scopes(value: Any) -> list[str]:
35
+ if isinstance(value, list):
36
+ return [str(item) for item in value]
37
+ return str(value or "").split()
38
+
39
+
40
+ def _oauth_error(error: OAuth2Error) -> JSONResponse:
41
+ status = 401 if error.error == "invalid_client" else error.http_status
42
+ headers = {"WWW-Authenticate": "Basic"} if status == 401 else None
43
+ return JSONResponse(
44
+ {
45
+ "error": error.error,
46
+ "error_description": error.description,
47
+ },
48
+ status_code=status,
49
+ headers=headers,
50
+ )
51
+
52
+
53
+ def create_oauth2_router(
54
+ server: OAuth2Server,
55
+ *,
56
+ prefix: str = "/oauth2",
57
+ ) -> APIRouter:
58
+ """创建 token / introspect / revoke 路由。
59
+
60
+ `/token` 同时接受 JSON 与 `application/x-www-form-urlencoded`,无需额外安装
61
+ `python-multipart`。
62
+ """
63
+ router = APIRouter(prefix=prefix, tags=["OAuth2"])
64
+
65
+ @router.post("/token")
66
+ async def token_endpoint(request: Request) -> JSONResponse:
67
+ parameters = await _read_parameters(request)
68
+ grant_type = str(parameters.get("grant_type", ""))
69
+ try:
70
+ if grant_type == "authorization_code":
71
+ response = await server.exchange_code_for_token(
72
+ code=str(parameters.get("code", "")),
73
+ client_id=str(parameters.get("client_id", "")),
74
+ client_secret=parameters.get("client_secret"),
75
+ redirect_uri=parameters.get("redirect_uri"),
76
+ code_verifier=parameters.get("code_verifier"),
77
+ )
78
+ elif grant_type == "refresh_token":
79
+ response = await server.refresh_access_token(
80
+ refresh_token=str(parameters.get("refresh_token", "")),
81
+ client_id=str(parameters.get("client_id", "")),
82
+ client_secret=parameters.get("client_secret"),
83
+ )
84
+ elif grant_type == "client_credentials":
85
+ response = await server.client_credentials_token(
86
+ client_id=str(parameters.get("client_id", "")),
87
+ client_secret=str(parameters.get("client_secret", "")),
88
+ scopes=_scopes(parameters.get("scope")),
89
+ )
90
+ else:
91
+ raise OAuth2Error("unsupported_grant_type", f"不支持的 grant_type: {grant_type}")
92
+ return JSONResponse(response.to_dict())
93
+ except OAuth2Error as error:
94
+ return _oauth_error(error)
95
+
96
+ @router.post("/introspect")
97
+ async def introspect_endpoint(request: Request) -> JSONResponse:
98
+ parameters = await _read_parameters(request)
99
+ return JSONResponse(await server.introspect(str(parameters.get("token", ""))))
100
+
101
+ @router.post("/revoke")
102
+ async def revoke_endpoint(request: Request) -> JSONResponse:
103
+ parameters = await _read_parameters(request)
104
+ await server.revoke_token(str(parameters.get("token", "")))
105
+ # RFC 7009:token 不存在时同样返回 200,避免泄露有效性。
106
+ return JSONResponse({})
107
+
108
+ return router
109
+
110
+
111
+ async def authorization_redirect(
112
+ server: OAuth2Server,
113
+ *,
114
+ login_id: str,
115
+ client_id: str,
116
+ redirect_uri: str,
117
+ scope: str = "",
118
+ state: str | None = None,
119
+ code_challenge: str | None = None,
120
+ code_challenge_method: str = "S256",
121
+ ) -> RedirectResponse:
122
+ """用户在业务授权页点击同意后,签发授权码并跳回客户端。"""
123
+ code = await server.create_authorization_code(
124
+ client_id=client_id,
125
+ login_id=login_id,
126
+ redirect_uri=redirect_uri,
127
+ scopes=_scopes(scope),
128
+ code_challenge=code_challenge,
129
+ code_challenge_method=code_challenge_method,
130
+ state=state,
131
+ )
132
+ parameters = {"code": code.code}
133
+ if state:
134
+ parameters["state"] = state
135
+ separator = "&" if "?" in redirect_uri else "?"
136
+ return RedirectResponse(f"{redirect_uri}{separator}{urlencode(parameters)}")
@@ -0,0 +1,191 @@
1
+ """Flask 适配。
2
+
3
+ Flask 是 WSGI 同步框架,所有调用经由 :mod:`sa_token.sync` 的同步桥接进入
4
+ 同一套 ``StpLogic``,不存在第二份鉴权实现。
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import functools
10
+ from collections.abc import Callable
11
+ from typing import TYPE_CHECKING, Any, TypeVar
12
+
13
+ from ..adapter.http import HttpContext
14
+ from ..adapter.path import PathAuthConfig
15
+ from ..adapter.pipeline import build_rule, resolve_token, run_auth_flow, run_path_auth
16
+ from ..exception import SaTokenException
17
+ from ..permission import MatchMode
18
+ from ..stp_util import get_manager
19
+ from ..sync import run_sync
20
+
21
+ if TYPE_CHECKING: # pragma: no cover - 仅供类型检查
22
+ from flask import Flask
23
+
24
+ __all__ = ["FlaskHttpContext", "SaTokenFlask"]
25
+
26
+ _F = TypeVar("_F", bound=Callable[..., Any])
27
+
28
+
29
+ class FlaskHttpContext(HttpContext):
30
+ """把 Flask ``request`` 包装成框架无关的上下文。"""
31
+
32
+ def __init__(self, request: Any) -> None:
33
+ self._request = request
34
+ self.state: dict[str, Any] = {}
35
+
36
+ def get_header(self, name: str) -> str | None:
37
+ return self._request.headers.get(name)
38
+
39
+ def get_cookie(self, name: str) -> str | None:
40
+ return self._request.cookies.get(name)
41
+
42
+ def get_query(self, name: str) -> str | None:
43
+ return self._request.args.get(name)
44
+
45
+ def get_path(self) -> str:
46
+ return self._request.path
47
+
48
+ def get_method(self) -> str:
49
+ return self._request.method
50
+
51
+
52
+ class SaTokenFlask:
53
+ """Flask 的标准注解鉴权,与 FastAPI / Django 同一套语义。
54
+
55
+ 用法::
56
+
57
+ sa = SaTokenFlask(app)
58
+
59
+ @app.get("/user")
60
+ @sa.check_login
61
+ def user_info():
62
+ return {"id": sa.login_id()}
63
+ """
64
+
65
+ def __init__(
66
+ self,
67
+ app: Flask | None = None,
68
+ *,
69
+ path_auth: PathAuthConfig | None = None,
70
+ login_type: str = "login",
71
+ ) -> None:
72
+ self.path_auth = path_auth
73
+ self.login_type = login_type
74
+ if app is not None:
75
+ self.init_app(app)
76
+
77
+ def init_app(self, app: Flask) -> None:
78
+ app.before_request(self._before_request)
79
+ app.register_error_handler(SaTokenException, self._handle_exception)
80
+
81
+ # 请求钩子 --------------------------------------------------------------
82
+
83
+ def _context(self) -> FlaskHttpContext:
84
+ from flask import request
85
+
86
+ return FlaskHttpContext(request)
87
+
88
+ def _before_request(self) -> Any:
89
+ from flask import g
90
+
91
+ ctx = self._context()
92
+ manager = get_manager()
93
+ if self.path_auth is None:
94
+ resolve_token(ctx, manager)
95
+ else:
96
+ run_sync(run_path_auth(ctx, manager, self.path_auth, login_type=self.login_type))
97
+ g.sa_token = ctx.state.get("stp_token")
98
+ g.sa_login_id = ctx.state.get("stp_login_id")
99
+ return None
100
+
101
+ def _handle_exception(self, exc: SaTokenException) -> Any:
102
+ from flask import jsonify
103
+
104
+ payload: dict[str, Any] = {
105
+ "code": exc.http_status,
106
+ "message": exc.message,
107
+ "error": type(exc).__name__,
108
+ }
109
+ detail_type = getattr(exc, "type", None)
110
+ if detail_type is not None:
111
+ payload["type"] = detail_type.value
112
+ return jsonify(payload), exc.http_status
113
+
114
+ # 取值 -----------------------------------------------------------------
115
+
116
+ def token(self) -> str | None:
117
+ return resolve_token(self._context(), get_manager())
118
+
119
+ def login_id(self) -> str:
120
+ """当前登录用户;未登录抛异常(会被错误处理器转成 401)。"""
121
+ from flask import g
122
+
123
+ cached = getattr(g, "sa_login_id", None)
124
+ if cached:
125
+ return cached
126
+ result = run_sync(run_auth_flow(self._context(), get_manager(), build_rule()))
127
+ assert result.login_id is not None
128
+ g.sa_login_id = result.login_id
129
+ return result.login_id
130
+
131
+ def login_id_or_none(self) -> str | None:
132
+ return run_sync(get_manager().stp(self.login_type).get_login_id_or_none(self.token()))
133
+
134
+ # 装饰器 ----------------------------------------------------------------
135
+
136
+ def _guard(self, rule_factory: Callable[[], Any]) -> Callable[[_F], _F]:
137
+ def decorator(view: _F) -> _F:
138
+ @functools.wraps(view)
139
+ def wrapper(*args: Any, **kwargs: Any) -> Any:
140
+ from flask import g
141
+
142
+ result = run_sync(
143
+ run_auth_flow(
144
+ self._context(),
145
+ get_manager(),
146
+ rule_factory(),
147
+ login_type=self.login_type,
148
+ )
149
+ )
150
+ g.sa_login_id = result.login_id
151
+ g.sa_token = result.token
152
+ return view(*args, **kwargs)
153
+
154
+ return wrapper # type: ignore[return-value]
155
+
156
+ return decorator
157
+
158
+ @property
159
+ def check_login(self) -> Callable[[_F], _F]:
160
+ """标准注解:``@sa.check_login``。"""
161
+ return self._guard(build_rule)
162
+
163
+ def check_permission(self, *permissions: str, mode: MatchMode = "OR") -> Callable[[_F], _F]:
164
+ return self._guard(lambda: build_rule(permissions=list(permissions), mode=mode))
165
+
166
+ def check_role(self, *roles: str, mode: MatchMode = "OR") -> Callable[[_F], _F]:
167
+ return self._guard(lambda: build_rule(roles=list(roles), mode=mode))
168
+
169
+ def check_safe(self, business: str) -> Callable[[_F], _F]:
170
+ def decorator(view: _F) -> _F:
171
+ @functools.wraps(view)
172
+ def wrapper(*args: Any, **kwargs: Any) -> Any:
173
+ logic = get_manager().stp(self.login_type)
174
+ run_sync(logic.check_safe(self.token(), business))
175
+ return view(*args, **kwargs)
176
+
177
+ return wrapper # type: ignore[return-value]
178
+
179
+ return decorator
180
+
181
+ def check_disable(self, service: str = "login", level: int = 1) -> Callable[[_F], _F]:
182
+ def decorator(view: _F) -> _F:
183
+ @functools.wraps(view)
184
+ def wrapper(*args: Any, **kwargs: Any) -> Any:
185
+ logic = get_manager().stp(self.login_type)
186
+ run_sync(logic.check_disable(self.login_id(), service=service, level=level))
187
+ return view(*args, **kwargs)
188
+
189
+ return wrapper # type: ignore[return-value]
190
+
191
+ return decorator
@@ -0,0 +1,227 @@
1
+ """Starlette / FastAPI 共用的适配层。
2
+
3
+ 职责严格限定在三件事:Request → HttpContext、调用统一管道、异常翻译。
4
+ 鉴权逻辑一行都不在这里实现。
5
+
6
+ 注意:这里必须在**运行时**导入 Starlette 的 ``Request``,不能放进
7
+ ``TYPE_CHECKING``。FastAPI 依靠运行时类型注解判断依赖项参数,
8
+ 拿不到真实类型时会把 ``request`` 当成查询参数,返回 422。
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from collections.abc import Awaitable, Callable
14
+ from typing import Any
15
+
16
+ try:
17
+ from starlette.applications import Starlette
18
+ from starlette.middleware.base import BaseHTTPMiddleware
19
+ from starlette.requests import Request
20
+ from starlette.responses import JSONResponse, Response
21
+ except ImportError as exc: # pragma: no cover - 依赖缺失路径
22
+ raise ImportError(
23
+ '该模块需要 starlette,请执行:pip install "sa-token-python-core[fastapi]"'
24
+ ) from exc
25
+
26
+ from ..adapter.http import HttpContext
27
+ from ..adapter.path import PathAuthConfig
28
+ from ..adapter.pipeline import build_rule, resolve_token, run_auth_flow, run_path_auth
29
+ from ..exception import SaTokenException
30
+ from ..permission import MatchMode
31
+ from ..stp_util import get_manager
32
+
33
+ __all__ = [
34
+ "StarletteHttpContext",
35
+ "SaTokenMiddleware",
36
+ "install_exception_handlers",
37
+ "sa_token_exception_handler",
38
+ "check_login",
39
+ "check_permission",
40
+ "check_role",
41
+ "check_disable",
42
+ "check_safe",
43
+ "current_login_id",
44
+ "current_login_id_or_none",
45
+ "current_token",
46
+ ]
47
+
48
+
49
+ class StarletteHttpContext(HttpContext):
50
+ """把 Starlette ``Request`` 包装成框架无关的上下文。"""
51
+
52
+ def __init__(self, request: Request) -> None:
53
+ self._request = request
54
+ self.state: dict[str, Any] = {}
55
+
56
+ def get_header(self, name: str) -> str | None:
57
+ return self._request.headers.get(name)
58
+
59
+ def get_cookie(self, name: str) -> str | None:
60
+ return self._request.cookies.get(name)
61
+
62
+ def get_query(self, name: str) -> str | None:
63
+ return self._request.query_params.get(name)
64
+
65
+ def get_path(self) -> str:
66
+ return self._request.url.path
67
+
68
+ def get_method(self) -> str:
69
+ return self._request.method
70
+
71
+
72
+ def _error_response(exc: SaTokenException) -> Response:
73
+ payload: dict[str, Any] = {
74
+ "code": exc.http_status,
75
+ "message": exc.message,
76
+ "error": type(exc).__name__,
77
+ }
78
+ detail_type = getattr(exc, "type", None)
79
+ if detail_type is not None:
80
+ payload["type"] = detail_type.value
81
+ return JSONResponse(payload, status_code=exc.http_status)
82
+
83
+
84
+ class SaTokenMiddleware(BaseHTTPMiddleware):
85
+ """请求级中间件。
86
+
87
+ 默认只解析 token 并绑定上下文,**不强制登录**:登录接口、健康检查这类
88
+ 公开路由永远存在,默认拦截会逼用户到处写例外。传入 ``path_auth``
89
+ 后才按规则表执行鉴权。
90
+ """
91
+
92
+ def __init__(
93
+ self,
94
+ app: Any,
95
+ *,
96
+ path_auth: PathAuthConfig | None = None,
97
+ login_type: str = "login",
98
+ ) -> None:
99
+ super().__init__(app)
100
+ self.path_auth = path_auth
101
+ self.login_type = login_type
102
+
103
+ async def dispatch(
104
+ self,
105
+ request: Request,
106
+ call_next: Callable[[Request], Awaitable[Response]],
107
+ ) -> Response:
108
+ ctx = StarletteHttpContext(request)
109
+ manager = get_manager()
110
+ try:
111
+ if self.path_auth is None:
112
+ resolve_token(ctx, manager)
113
+ else:
114
+ await run_path_auth(ctx, manager, self.path_auth, login_type=self.login_type)
115
+ except SaTokenException as exc:
116
+ return _error_response(exc)
117
+
118
+ request.state.sa_token = ctx.state.get("stp_token")
119
+ request.state.sa_login_id = ctx.state.get("stp_login_id")
120
+ return await call_next(request)
121
+
122
+
123
+ async def sa_token_exception_handler(request: Request, exc: Exception) -> Response:
124
+ """把核心异常翻译成 401 / 403 / 500。"""
125
+ assert isinstance(exc, SaTokenException)
126
+ return _error_response(exc)
127
+
128
+
129
+ def install_exception_handlers(app: Starlette) -> None:
130
+ """注册异常处理器,让路由里抛出的鉴权异常也能正确变成 HTTP 状态码。"""
131
+ app.add_exception_handler(SaTokenException, sa_token_exception_handler)
132
+
133
+
134
+ # FastAPI / Starlette 的 Depends 工厂(扩展写法)。
135
+ # 各框架标准鉴权是注解:FastAPI 用 @sa.check_login,见 SaTokenFastAPI。
136
+
137
+
138
+ async def _authorize(
139
+ request: Request,
140
+ *,
141
+ permissions: list[str] | None = None,
142
+ roles: list[str] | None = None,
143
+ mode: MatchMode = "OR",
144
+ ) -> tuple[str, str | None]:
145
+ ctx = StarletteHttpContext(request)
146
+ rule = build_rule(permissions=permissions, roles=roles, mode=mode)
147
+ result = await run_auth_flow(ctx, get_manager(), rule)
148
+ assert result.login_id is not None
149
+ request.state.sa_login_id = result.login_id
150
+ request.state.sa_token = result.token
151
+ return result.login_id, result.token
152
+
153
+
154
+ def check_login() -> Callable[[Request], Awaitable[str]]:
155
+ """Depends 扩展:``Depends(check_login())``。标准写法是 ``@sa.check_login``。"""
156
+
157
+ async def dependency(request: Request) -> str:
158
+ login_id, _ = await _authorize(request)
159
+ return login_id
160
+
161
+ return dependency
162
+
163
+
164
+ def check_permission(
165
+ *permissions: str,
166
+ mode: MatchMode = "OR",
167
+ ) -> Callable[[Request], Awaitable[str]]:
168
+ """Depends 扩展:``Depends(check_permission("order:delete"))``。"""
169
+
170
+ async def dependency(request: Request) -> str:
171
+ login_id, _ = await _authorize(request, permissions=list(permissions), mode=mode)
172
+ return login_id
173
+
174
+ return dependency
175
+
176
+
177
+ def check_role(*roles: str, mode: MatchMode = "OR") -> Callable[[Request], Awaitable[str]]:
178
+ """Depends 扩展:``Depends(check_role("admin"))``。"""
179
+
180
+ async def dependency(request: Request) -> str:
181
+ login_id, _ = await _authorize(request, roles=list(roles), mode=mode)
182
+ return login_id
183
+
184
+ return dependency
185
+
186
+
187
+ def check_disable(
188
+ service: str = "login",
189
+ level: int = 1,
190
+ ) -> Callable[[Request], Awaitable[str]]:
191
+ """Depends 扩展:``Depends(check_disable("comment"))``。"""
192
+
193
+ async def dependency(request: Request) -> str:
194
+ login_id, _ = await _authorize(request)
195
+ await get_manager().stp().check_disable(login_id, service=service, level=level)
196
+ return login_id
197
+
198
+ return dependency
199
+
200
+
201
+ def check_safe(business: str) -> Callable[[Request], Awaitable[str]]:
202
+ """Depends 扩展:``Depends(check_safe("pay"))``。"""
203
+
204
+ async def dependency(request: Request) -> str:
205
+ login_id, token = await _authorize(request)
206
+ await get_manager().stp().check_safe(token, business)
207
+ return login_id
208
+
209
+ return dependency
210
+
211
+
212
+ async def current_login_id(request: Request) -> str:
213
+ """Depends 扩展:``Depends(current_login_id)``。标准写法是 ``@sa.check_login``。"""
214
+ login_id, _ = await _authorize(request)
215
+ return login_id
216
+
217
+
218
+ async def current_login_id_or_none(request: Request) -> str | None:
219
+ """Depends 扩展:未登录时返回 None。"""
220
+ ctx = StarletteHttpContext(request)
221
+ token = resolve_token(ctx, get_manager())
222
+ return await get_manager().stp().get_login_id_or_none(token)
223
+
224
+
225
+ async def current_token(request: Request) -> str | None:
226
+ """Depends 扩展:只取 token,不做校验。"""
227
+ return resolve_token(StarletteHttpContext(request), get_manager())
sa_token/listener.py ADDED
@@ -0,0 +1,100 @@
1
+ """事件系统:用于审计日志、指标、消息推送。
2
+
3
+ 监听器异常不会影响主流程——审计失败不该导致用户登录失败。
4
+ 事件载荷只带 token 指纹,避免把原始 token 写进日志文件。
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import asyncio
10
+ import hashlib
11
+ import logging
12
+ from collections.abc import Awaitable, Callable
13
+ from dataclasses import dataclass, field
14
+ from enum import Enum
15
+ from typing import Any
16
+
17
+ from .model import now_ms
18
+
19
+ __all__ = ["Event", "EventData", "EventBus", "Listener"]
20
+
21
+ logger = logging.getLogger("sa_token.listener")
22
+
23
+
24
+ class Event(str, Enum):
25
+ LOGIN = "login"
26
+ LOGOUT = "logout"
27
+ KICKOUT = "kickout"
28
+ REPLACED = "replaced"
29
+ DISABLE = "disable"
30
+ UNTIE = "untie"
31
+ RENEW = "renew"
32
+ PERMISSION_CHECK = "permission_check"
33
+ ROLE_CHECK = "role_check"
34
+ ALL = "*"
35
+
36
+
37
+ @dataclass
38
+ class EventData:
39
+ event: Event
40
+ login_id: str | None = None
41
+ login_type: str = "login"
42
+ device: str | None = None
43
+ token_fingerprint: str | None = None
44
+ timestamp: int = field(default_factory=now_ms)
45
+ detail: dict[str, Any] = field(default_factory=dict)
46
+
47
+
48
+ Listener = Callable[[EventData], Any | Awaitable[Any]]
49
+
50
+
51
+ def fingerprint(token: str | None) -> str | None:
52
+ """token 的短指纹:可用于关联同一条会话,但无法还原出 token。"""
53
+ if not token:
54
+ return None
55
+ return hashlib.sha256(token.encode("utf-8")).hexdigest()[:16]
56
+
57
+
58
+ @dataclass(order=True)
59
+ class _Registration:
60
+ # 先按优先级倒序,再按注册序号,保证同优先级下注册顺序稳定。
61
+ sort_key: tuple[int, int]
62
+ listener: Listener = field(compare=False)
63
+ is_async: bool = field(compare=False, default=False)
64
+
65
+
66
+ class EventBus:
67
+ def __init__(self) -> None:
68
+ self._listeners: dict[Event, list[_Registration]] = {}
69
+ self._counter = 0
70
+
71
+ def on(self, event: Event, listener: Listener, *, priority: int = 0) -> None:
72
+ """注册监听器,``priority`` 越大越先执行。"""
73
+ self._counter += 1
74
+ registration = _Registration(
75
+ sort_key=(-priority, self._counter),
76
+ listener=listener,
77
+ is_async=asyncio.iscoroutinefunction(listener),
78
+ )
79
+ bucket = self._listeners.setdefault(event, [])
80
+ bucket.append(registration)
81
+ bucket.sort()
82
+
83
+ def off(self, event: Event, listener: Listener) -> None:
84
+ bucket = self._listeners.get(event)
85
+ if not bucket:
86
+ return
87
+ self._listeners[event] = [item for item in bucket if item.listener is not listener]
88
+
89
+ def clear(self) -> None:
90
+ self._listeners.clear()
91
+
92
+ async def emit(self, data: EventData) -> None:
93
+ matched = [*self._listeners.get(data.event, []), *self._listeners.get(Event.ALL, [])]
94
+ for registration in matched:
95
+ try:
96
+ result = registration.listener(data)
97
+ if registration.is_async or asyncio.iscoroutine(result):
98
+ await result
99
+ except Exception:
100
+ logger.exception("sa-token 事件监听器执行失败:event=%s", data.event.value)