Files
blog-press/docs/FastAPI/OAuth2.md
2025-11-29 21:49:12 +08:00

165 lines
5.2 KiB
Markdown
Raw Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# 一、基础概念
## 1.1 OAuth2 Password Bearer 模式​
  用于**用户名+密码**登录,获取**access_token**。
## 1.2 FastAPI 的 OAuth2PasswordBearer
  从请求头**Authorization**中提取TokenToken格式为 **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)])
```