165 lines
5.2 KiB
Markdown
165 lines
5.2 KiB
Markdown
# 一、基础概念
|
||
## 1.1 OAuth2 Password Bearer 模式
|
||
  用于**用户名+密码**登录,获取**access_token**。
|
||
|
||
## 1.2 FastAPI 的 OAuth2PasswordBearer
|
||
  从请求头**Authorization**中提取Token,Token格式为 **Bearer Token**,必须是这个格式,如果不是则会提示**401 Unauthorized**错误。
|
||
|
||
# 二、核心流程
|
||
## 2.1 获取Token
|
||
```python
|
||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="session")
|
||
|
||
@app.post("/session")
|
||
async def login(form_data: OAuth2PasswordRequestForm = Depends()):
|
||
# 1. 验证用户名和密码
|
||
# 2. 返回 Token
|
||
return Token
|
||
```
|
||
|
||
   这里的**tokenUrl="session"** 对应的是fastapi中的url路径(@app.post("/session")),用于Swagger文档中的Token认证。
|
||
   前端必须要以FormData表单的形式传递username和password,且必须是username和password字段。
|
||
```ts
|
||
const data = new FormData()
|
||
data.append('username', username)
|
||
data.append('password', password)
|
||
```
|
||
   在获取Token前,一般还需要进行验证用户名和密码是否和数据库中的信息一致。
|
||
   可以采用JWT格式封装TokenValue:
|
||
```python
|
||
def create_token(payload: dict, expires_delta: Optional[timedelta] = None):
|
||
# 复制一份
|
||
payload_copy = payload.copy()
|
||
|
||
# 加上有效时间
|
||
if expires_delta:
|
||
expire = datetime.now() + expires_delta
|
||
else:
|
||
expire = datetime.now() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
|
||
payload_copy.update({"exp": expire})
|
||
|
||
# 生成jwt Token
|
||
return jwt.encode(payload_copy, SECRET_KEY, algorithm=ALGORITHM)
|
||
```
|
||
|
||
## 2.2 访问接口
|
||
```python
|
||
@router.delete("/{blog_id}", summary="删除博客内容", response_model=bool)
|
||
def delete_blog(blog_id: int, db: Session = Depends(get_db), _ = Depends(verify_token)):
|
||
return blog_service.delete_blog(db, blog_id)
|
||
```
|
||
   例如访问这个删除接口,在参数中添加验证token的依赖:
|
||
```python
|
||
async def verify_token(token: str = Depends(oauth2_scheme)):
|
||
try:
|
||
# 1. 检验token信息
|
||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||
# 2. 校验payload信息
|
||
if not verify_payload(payload):
|
||
raise get_credentials_exception()
|
||
# 3. 校验数据库中是否存在 payload中的账户信息
|
||
sub = get_payload_sub(payload)
|
||
|
||
# 4. 存储账户信息
|
||
context_sub.set(sub)
|
||
except jwt.exceptions.InvalidTokenError:
|
||
raise get_credentials_exception()
|
||
```
|
||
   该依赖又依赖于子依赖oauth2_scheme,通过调用OAuth2PasswordBearer方法从请求头Authorization中获取token值。
|
||
   校验Token通常包含校验格式是否正确和Token包含的账户信息是否正确。
|
||
|
||
# 三、参考代码
|
||
```python
|
||
from contextvars import ContextVar
|
||
from datetime import datetime, timedelta
|
||
from typing import Optional
|
||
|
||
from fastapi import Depends
|
||
from fastapi.security import OAuth2PasswordBearer
|
||
import jwt
|
||
from passlib.context import CryptContext
|
||
|
||
from middleware.exceptions import get_credentials_exception
|
||
|
||
# 密钥和算法配置
|
||
SECRET_KEY = "sjdi@!#3ksj2780se1283"
|
||
ALGORITHM = "HS256"
|
||
ACCESS_TOKEN_EXPIRE_MINUTES = 30
|
||
|
||
# 密码哈希上下文
|
||
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
||
|
||
# OAuth2 方案
|
||
# 设置默认登录接口
|
||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="session")
|
||
|
||
# 请求上下文
|
||
context_sub: ContextVar[str] = ContextVar('sub')
|
||
|
||
|
||
# 获取数据库 加密密码
|
||
def get_password_hash(password: str):
|
||
return pwd_context.hash(password)
|
||
|
||
|
||
# 验证数据库密码
|
||
def verify_password(plain_password: str, hashed_password: str):
|
||
return pwd_context.verify(plain_password, hashed_password)
|
||
|
||
|
||
# 生成payload
|
||
def create_payload(sub: str) -> dict:
|
||
return {
|
||
"sub": sub
|
||
}
|
||
|
||
|
||
# 验证payload
|
||
def verify_payload(payload: dict) -> bool:
|
||
return "sub" in payload
|
||
|
||
|
||
# 从payload中获取用户
|
||
def get_payload_sub(payload: dict) -> str:
|
||
return payload["sub"]
|
||
|
||
|
||
# 创建token
|
||
def create_token(payload: dict, expires_delta: Optional[timedelta] = None):
|
||
# 复制一份
|
||
payload_copy = payload.copy()
|
||
|
||
# 加上有效时间
|
||
if expires_delta:
|
||
expire = datetime.now() + expires_delta
|
||
else:
|
||
expire = datetime.now() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
|
||
payload_copy.update({"exp": expire})
|
||
|
||
# 生成jwt Token
|
||
return jwt.encode(payload_copy, SECRET_KEY, algorithm=ALGORITHM)
|
||
|
||
|
||
# 验证token
|
||
async def verify_token(token: str = Depends(oauth2_scheme)):
|
||
try:
|
||
# 1. 检验token信息
|
||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||
# 2. 校验payload信息
|
||
if not verify_payload(payload):
|
||
raise get_credentials_exception()
|
||
# 3. 校验数据库中是否存在 payload中的账户信息
|
||
sub = get_payload_sub(payload)
|
||
|
||
# 4. 存储账户信息
|
||
context_sub.set(sub)
|
||
except jwt.exceptions.InvalidTokenError:
|
||
raise get_credentials_exception()
|
||
```
|
||
|
||
# 四、注意事项
|
||
1. 如果整个路由模块都需要Token验证,可以在APIRouter中添加依赖:
|
||
|
||
```python
|
||
protected_router = APIRouter(dependencies=[Depends(verify_token)])
|
||
``` |