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
|
@@ -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)
|