安全体系——认证、权限与限流完整实现

概述

安全性是REST API设计的核心考量之一。Django REST Framework提供了完整的安全体系架构,包括认证(Authentication)、权限(Permission)和限流(Throttling)三个核心组件。本文将深入剖析这套安全体系的工作原理、实现机制和最佳实践。


一、安全体系架构总览

1.1 三层安全架构

请求处理流程

HTTP请求

认证层

认证成功?

权限层

返回401

有权限?

限流层

返回403

未超限?

执行视图

返回429

返回响应

1.2 安全组件职责

组件职责失败响应核心问题
认证验证用户身份401 Unauthorized你是谁?
权限检查访问权限403 Forbidden你能做什么?
限流控制请求频率429 Too Many Requests你能做多少次?

1.3 安全检查执行顺序

"""
安全检查执行顺序(在APIView.initial()方法中):

1. perform_content_negotiation() - 内容协商
2. determine_version() - 版本检查
3. perform_authentication() - 认证检查
4. check_permissions() - 权限检查
5. check_throttles() - 限流检查

代码位置:rest_framework/views.py
"""

def initial(self, request, *args, **kwargs):
    """
    请求预处理方法
    """
    # 1. 内容协商
    self.perform_content_negotiation(request)
    
    # 2. 版本检查
    self.determine_version(request, *args, **kwargs)
    
    # 3. 认证检查
    self.perform_authentication(request)
    
    # 4. 权限检查
    self.check_permissions(request)
    
    # 5. 限流检查
    self.check_throttles(request)

二、认证系统深度解析

2.1 认证系统架构

BaseAuthentication

+authenticate(request) : Tuple<User, Auth>

+authenticate_header(request) : str

BasicAuthentication

+www_authenticate_realm: str

+authenticate(request)

TokenAuthentication

+keyword: str

+model: Token

+authenticate(request)

+get_model()

SessionAuthentication

+authenticate(request)

+enforce_csrf(request)

RemoteUserAuthentication

+header: str

+authenticate(request)

JWTAuthentication

+keyword: str

+algorithm: str

+secret_key: str

+authenticate(request)

+get_user(payload)

+create_token(user)

2.2 BaseAuthentication源码解析

from abc import ABC, abstractmethod
from rest_framework.exceptions import AuthenticationFailed, NotAuthenticated


class BaseAuthentication(ABC):
    """
    认证基类
    
    所有认证类必须继承此类并实现authenticate方法
    
    authenticate方法返回值:
    - (user, auth) 元组:认证成功
    - None:跳过当前认证器,继续下一个
    - 抛出AuthenticationFailed:认证失败
    
    参考:rest_framework/authentication.py
    """
    
    def authenticate(self, request):
        """
        认证方法
        
        参数:
            request: Request对象
            
        返回:
            (user, auth) 元组 或 None
            
        抛出:
            AuthenticationFailed: 认证失败
            NotAuthenticated: 未认证
        """
        raise NotImplementedError(".authenticate() must be overridden.")
    
    def authenticate_header(self, request):
        """
        返回WWW-Authenticate头值
        
        用于401响应时返回认证方案信息
        如果返回None,则不添加WWW-Authenticate头
        """
        return None


# ===== 认证流程源码 =====
class APIView:
    def perform_authentication(self, request):
        """
        执行认证
        
        遍历所有认证器,直到有一个成功或全部失败
        """
        request.user  # 触发延迟认证
    
    def get_authenticators(self):
        """
        获取认证器实例列表
        """
        return [auth() for auth in self.authentication_classes]


class Request:
    @property
    def user(self):
        """
        延迟加载用户
        
        首次访问时执行认证
        """
        if not hasattr(self, '_user'):
            self._authenticate()
        return self._user
    
    def _authenticate(self):
        """
        执行认证流程
        
        遍历所有认证器:
        1. 如果返回(user, auth),认证成功
        2. 如果返回None,继续下一个认证器
        3. 如果抛出异常,认证失败
        """
        for authenticator in self.authenticators:
            try:
                user_auth_tuple = authenticator.authenticate(self)
            except AuthenticationFailed as exc:
                raise exc
            
            if user_auth_tuple is not None:
                self._user, self._auth = user_auth_tuple
                return
        
        # 所有认证器都返回None,设置为匿名用户
        self._user = self._default_user()
        self._auth = self._default_auth()

2.3 内置认证类详解

2.3.1 BasicAuthentication
from rest_framework.authentication import BasicAuthentication
import base64


class BasicAuthentication(BaseAuthentication):
    """
    HTTP Basic认证
    
    请求头格式:
        Authorization: Basic base64(username:password)
    
    示例:
        用户名: admin
        密码: password123
        Base64编码: YWRtaW46cGFzc3dvcmQxMjM=
        请求头: Authorization: Basic YWRtaW46cGFzc3dvcmQxMjM=
    
    特点:
        1. 简单易用
        2. 每次请求都传输凭据
        3. 必须配合HTTPS使用
        4. 不支持登出(浏览器缓存凭据)
    
    适用场景:
        - 内部API
        - 开发测试环境
        - 简单的脚本调用
    """
    
    www_authenticate_realm = 'api'
    
    def authenticate(self, request):
        """
        认证流程:
        1. 从请求头获取Authorization
        2. 解码Base64获取用户名密码
        3. 验证用户凭据
        """
        auth_header = request.META.get('HTTP_AUTHORIZATION')
        
        if not auth_header:
            return None
        
        # 解析认证头
        parts = auth_header.split()
        if len(parts) != 2 or parts[0].lower() != 'basic':
            return None
        
        # Base64解码
        try:
            decoded = base64.b64decode(parts[1]).decode('utf-8')
            username, password = decoded.split(':', 1)
        except (ValueError, UnicodeDecodeError):
            raise AuthenticationFailed('无效的认证头')
        
        # 验证用户
        return self.authenticate_credentials(username, password)
    
    def authenticate_credentials(self, username, password):
        """
        验证用户凭据
        """
        from django.contrib.auth import authenticate
        
        user = authenticate(username=username, password=password)
        
        if user is None:
            raise AuthenticationFailed('用户名或密码错误')
        
        if not user.is_active:
            raise AuthenticationFailed('用户已被禁用')
        
        return (user, None)
    
    def authenticate_header(self, request):
        """
        返回WWW-Authenticate头
        
        用于浏览器弹出认证对话框
        """
        return f'Basic realm="{self.www_authenticate_realm}"'


# ===== 使用示例 =====
# curl请求示例
# curl -u admin:password123 http://api.example.com/books/
# 或
# curl -H "Authorization: Basic YWRtaW46cGFzc3dvcmQxMjM=" http://api.example.com/books/
2.3.2 TokenAuthentication
from rest_framework.authentication import TokenAuthentication
from rest_framework.authtoken.models import Token


class TokenAuthentication(BaseAuthentication):
    """
    Token认证
    
    请求头格式:
        Authorization: Token <token_value>
    
    示例:
        Authorization: Token 9944b09199c62bcf9418ad846dd0e4bbdfc6ee4b
    
    特点:
        1. 无状态,服务器不存储会话
        2. Token可存储在数据库中
        3. 支持Token过期和撤销
        4. 适合API认证
    
    工作流程:
        1. 用户登录,服务器生成Token
        2. 客户端存储Token
        3. 每次请求携带Token
        4. 服务器验证Token有效性
    """
    
    keyword = 'Token'
    model = Token
    
    def authenticate(self, request):
        """
        认证流程:
        1. 从请求头获取Token
        2. 查询Token对应的用户
        3. 返回用户和Token对象
        """
        auth_header = request.META.get('HTTP_AUTHORIZATION')
        
        if not auth_header:
            return None
        
        # 解析认证头
        parts = auth_header.split()
        if len(parts) != 2 or parts[0].lower() != self.keyword.lower():
            return None
        
        token = parts[1]
        return self.authenticate_credentials(token)
    
    def authenticate_credentials(self, key):
        """
        验证Token
        """
        try:
            token = self.model.objects.select_related('user').get(key=key)
        except self.model.DoesNotExist:
            raise AuthenticationFailed('无效的Token')
        
        if not token.user.is_active:
            raise AuthenticationFailed('用户已被禁用')
        
        return (token.user, token)
    
    def authenticate_header(self, request):
        return self.keyword


# ===== Token模型 =====
from django.db import models
from django.conf import settings


class Token(models.Model):
    """
    Token模型
    
    字段:
        key: Token字符串(唯一)
        user: 关联用户
        created: 创建时间
    """
    key = models.CharField(max_length=40, primary_key=True)
    user = models.OneToOneField(
        settings.AUTH_USER_MODEL,
        on_delete=models.CASCADE,
        related_name='auth_token'
    )
    created = models.DateTimeField(auto_now_add=True)
    
    def save(self, *args, **kwargs):
        if not self.key:
            self.key = self.generate_key()
        return super().save(*args, **kwargs)
    
    @staticmethod
    def generate_key():
        """生成随机Token"""
        import secrets
        return secrets.token_hex(20)


# ===== Token生成视图 =====
from rest_framework.decorators import api_view, authentication_classes, permission_classes
from rest_framework.permissions import AllowAny
from rest_framework.response import Response


@api_view(['POST'])
@authentication_classes([])  # 不需要认证
@permission_classes([AllowAny])
def obtain_auth_token(request):
    """
    获取Token视图
    
    请求体:
        {
            "username": "admin",
            "password": "password123"
        }
    
    响应:
        {
            "token": "9944b09199c62bcf9418ad846dd0e4bbdfc6ee4b"
        }
    """
    from django.contrib.auth import authenticate
    
    username = request.data.get('username')
    password = request.data.get('password')
    
    if not username or not password:
        return Response(
            {'error': '用户名和密码不能为空'},
            status=400
        )
    
    user = authenticate(username=username, password=password)
    
    if user is None:
        return Response(
            {'error': '用户名或密码错误'},
            status=401
        )
    
    if not user.is_active:
        return Response(
            {'error': '用户已被禁用'},
            status=401
        )
    
    # 获取或创建Token
    token, created = Token.objects.get_or_create(user=user)
    
    return Response({'token': token.key})


@api_view(['POST'])
def revoke_auth_token(request):
    """
    撤销Token(登出)
    """
    request.user.auth_token.delete()
    return Response({'message': 'Token已撤销'})


# ===== 使用示例 =====
# 获取Token
# curl -X POST -H "Content-Type: application/json" \
#      -d '{"username":"admin","password":"password123"}' \
#      http://api.example.com/api-token-auth/

# 使用Token访问API
# curl -H "Authorization: Token 9944b09199c62bcf9418ad846dd0e4bbdfc6ee4b" \
#      http://api.example.com/books/
2.3.3 SessionAuthentication
from rest_framework.authentication import SessionAuthentication
from django.middleware.csrf import CsrfViewMiddleware


class SessionAuthentication(BaseAuthentication):
    """
    Django Session认证
    
    特点:
        1. 使用Django内置的Session机制
        2. 自动进行CSRF保护
        3. 适合前后端不分离的项目
        4. 依赖浏览器Cookie
    
    工作流程:
        1. 用户登录,服务器创建Session
        2. Session ID存储在Cookie中
        3. 每次请求自动携带Cookie
        4. 服务器验证Session和CSRF Token
    
    注意事项:
        - 需要配置CSRF中间件
        - 非浏览器客户端需要处理CSRF
        - 需要配置CORS(跨域场景)
    """
    
    def authenticate(self, request):
        """
        认证流程:
        1. 从Session获取用户
        2. 验证CSRF Token
        """
        # 获取Session中的用户
        user = getattr(request._request, 'user', None)
        
        if not user or not user.is_authenticated:
            return None
        
        # CSRF检查
        self.enforce_csrf(request._request)
        
        return (user, None)
    
    def enforce_csrf(self, request):
        """
        执行CSRF检查
        
        CSRF保护机制:
        1. 服务器生成CSRF Token
        2. Token存储在Cookie和表单中
        3. 提交时验证两个Token是否一致
        """
        from django.core.exceptions import PermissionDenied
        
        reason = CsrfViewMiddleware().process_view(
            request, None, (), {}
        )
        
        if reason:
            raise PermissionDenied(f'CSRF验证失败: {reason}')


# ===== CSRF配置 =====
# settings.py
MIDDLEWARE = [
    'django.middleware.security.SecurityMiddleware',
    'django.contrib.sessions.middleware.SessionMiddleware',
    'django.middleware.common.CommonMiddleware',
    'django.middleware.csrf.CsrfViewMiddleware',  # CSRF中间件
    'django.contrib.auth.middleware.AuthenticationMiddleware',
    'django.contrib.messages.middleware.MessageMiddleware',
]

# ===== 前端处理CSRF =====
# 方式1:从Cookie获取CSRF Token
function getCookie(name) {
    let cookieValue = null;
    if (document.cookie && document.cookie !== '') {
        const cookies = document.cookie.split(';');
        for (let i = 0; i < cookies.length; i++) {
            const cookie = cookies[i].trim();
            if (cookie.substring(0, name.length + 1) === (name + '=')) {
                cookieValue = decodeURIComponent(cookie.substring(name.length + 1));
                break;
            }
        }
    }
    return cookieValue;
}

// 在请求头中添加CSRF Token
fetch('/api/books/', {
    method: 'POST',
    headers: {
        'Content-Type': 'application/json',
        'X-CSRFToken': getCookie('csrftoken'),
    },
    body: JSON.stringify({title: 'New Book'}),
    credentials: 'include',  // 包含Cookie
});


# ===== 方式2:从响应头获取CSRF Token =====
# 视图中设置CSRF Token
from django.utils.decorators import method_decorator
from django.views.decorators.csrf import ensure_csrf_cookie

class BookViewSet(viewsets.ModelViewSet):
    authentication_classes = [SessionAuthentication]
    
    @method_decorator(ensure_csrf_cookie)
    def list(self, request, *args, **kwargs):
        """确保设置CSRF Cookie"""
        return super().list(request, *args, **kwargs)

2.4 自定义认证类

2.4.1 JWT认证实现
import jwt
from datetime import datetime, timedelta
from rest_framework.authentication import BaseAuthentication
from rest_framework.exceptions import AuthenticationFailed
from django.conf import settings
from django.contrib.auth import get_user_model

User = get_user_model()


class JWTAuthentication(BaseAuthentication):
    """
    JWT(JSON Web Token)认证
    
    请求头格式:
        Authorization: Bearer <jwt_token>
    
    JWT结构:
        Header.Payload.Signature
    
    特点:
        1. 无状态,不需要服务器存储
        2. 包含用户信息和过期时间
        3. 支持跨域
        4. 适合分布式系统
    
    参考:
        RFC 7519: https://tools.ietf.org/html/rfc7519
    """
    
    keyword = 'Bearer'
    algorithm = 'HS256'
    
    # 安全配置:使用独立的JWT密钥,而非Django的SECRET_KEY
    # 在settings.py中配置: JWT_SECRET_KEY = 'your-jwt-secret-key'
    # 如果未配置,则回退使用SECRET_KEY(不推荐生产环境)
    secret_key = getattr(settings, 'JWT_SECRET_KEY', settings.SECRET_KEY)
    
    access_token_lifetime = timedelta(hours=1)
    refresh_token_lifetime = timedelta(days=7)
    
    # 安全建议:
    # 1. 生产环境必须配置独立的JWT_SECRET_KEY
    # 2. JWT_SECRET_KEY应与SECRET_KEY不同
    # 3. 定期轮换密钥
    # 4. 使用强随机字符串作为密钥(至少32字符)
    
    def authenticate(self, request):
        """
        认证流程:
        1. 从请求头获取Token
        2. 解码验证Token
        3. 返回用户对象
        """
        auth_header = request.META.get('HTTP_AUTHORIZATION')
        
        if not auth_header:
            return None
        
        # 解析认证头
        parts = auth_header.split()
        if len(parts) != 2 or parts[0].lower() != self.keyword.lower():
            return None
        
        token = parts[1]
        
        try:
            # 解码Token
            payload = jwt.decode(
                token,
                self.secret_key,
                algorithms=[self.algorithm]
            )
        except jwt.ExpiredSignatureError:
            raise AuthenticationFailed('Token已过期')
        except jwt.InvalidTokenError as e:
            raise AuthenticationFailed(f'无效的Token: {str(e)}')
        
        # 获取用户
        user = self.get_user(payload)
        
        return (user, token)
    
    def get_user(self, payload):
        """
        从payload获取用户
        """
        user_id = payload.get('user_id')
        
        if not user_id:
            raise AuthenticationFailed('Token中缺少用户标识')
        
        try:
            user = User.objects.get(pk=user_id)
        except User.DoesNotExist:
            raise AuthenticationFailed('用户不存在')
        
        if not user.is_active:
            raise AuthenticationFailed('用户已被禁用')
        
        return user
    
    @classmethod
    def create_access_token(cls, user):
        """
        创建访问Token
        
        有效期短,用于API访问
        """
        payload = {
            'user_id': user.pk,
            'username': user.username,
            'exp': datetime.utcnow() + cls.access_token_lifetime,
            'iat': datetime.utcnow(),
            'type': 'access',
        }
        return jwt.encode(payload, cls.secret_key, algorithm=cls.algorithm)
    
    @classmethod
    def create_refresh_token(cls, user):
        """
        创建刷新Token
        
        有效期长,用于刷新访问Token
        """
        payload = {
            'user_id': user.pk,
            'exp': datetime.utcnow() + cls.refresh_token_lifetime,
            'iat': datetime.utcnow(),
            'type': 'refresh',
        }
        return jwt.encode(payload, cls.secret_key, algorithm=cls.algorithm)
    
    @classmethod
    def verify_refresh_token(cls, token):
        """
        验证刷新Token
        """
        try:
            payload = jwt.decode(
                token,
                cls.secret_key,
                algorithms=[cls.algorithm]
            )
            
            if payload.get('type') != 'refresh':
                raise AuthenticationFailed('无效的刷新Token')
            
            return payload
        except jwt.ExpiredSignatureError:
            raise AuthenticationFailed('刷新Token已过期')
        except jwt.InvalidTokenError:
            raise AuthenticationFailed('无效的刷新Token')


# ===== JWT登录视图 =====
from rest_framework.decorators import api_view, authentication_classes, permission_classes
from rest_framework.permissions import AllowAny
from rest_framework.response import Response
from django.contrib.auth import authenticate


@api_view(['POST'])
@authentication_classes([])
@permission_classes([AllowAny])
def jwt_login(request):
    """
    JWT登录视图
    
    请求体:
        {
            "username": "admin",
            "password": "password123"
        }
    
    响应:
        {
            "access": "eyJ0eXAiOiJKV1QiLCJhbGc...",
            "refresh": "eyJ0eXAiOiJKV1QiLCJhbGc...",
            "user": {
                "id": 1,
                "username": "admin"
            }
        }
    """
    username = request.data.get('username')
    password = request.data.get('password')
    
    if not username or not password:
        return Response(
            {'error': '用户名和密码不能为空'},
            status=400
        )
    
    user = authenticate(username=username, password=password)
    
    if user is None:
        return Response(
            {'error': '用户名或密码错误'},
            status=401
        )
    
    if not user.is_active:
        return Response(
            {'error': '用户已被禁用'},
            status=401
        )
    
    # 生成Token
    access_token = JWTAuthentication.create_access_token(user)
    refresh_token = JWTAuthentication.create_refresh_token(user)
    
    return Response({
        'access': access_token,
        'refresh': refresh_token,
        'user': {
            'id': user.pk,
            'username': user.username,
            'email': user.email,
        }
    })


@api_view(['POST'])
@authentication_classes([])
@permission_classes([AllowAny])
def jwt_refresh(request):
    """
    刷新Token视图
    
    请求体:
        {
            "refresh": "eyJ0eXAiOiJKV1QiLCJhbGc..."
        }
    
    响应:
        {
            "access": "eyJ0eXAiOiJKV1QiLCJhbGc..."
        }
    """
    refresh_token = request.data.get('refresh')
    
    if not refresh_token:
        return Response(
            {'error': '缺少刷新Token'},
            status=400
        )
    
    try:
        payload = JWTAuthentication.verify_refresh_token(refresh_token)
        user = User.objects.get(pk=payload['user_id'])
        
        # 生成新的访问Token
        access_token = JWTAuthentication.create_access_token(user)
        
        return Response({
            'access': access_token,
        })
    except AuthenticationFailed as e:
        return Response(
            {'error': str(e)},
            status=401
        )
    except User.DoesNotExist:
        return Response(
            {'error': '用户不存在'},
            status=401
        )


# ===== 完整的JWT Token刷新机制 =====
class JWTTokenManager:
    """
    JWT Token管理器
    
    提供完整的Token生命周期管理:
    1. Token生成
    2. Token刷新
    3. Token撤销(黑名单)
    4. Token验证
    """
    
    # Token黑名单(生产环境应使用Redis)
    _blacklist = set()
    
    @classmethod
    def generate_tokens(cls, user):
        """
        生成完整的Token对
        
        返回:
            {
                'access': '...',
                'refresh': '...',
                'expires_in': 3600  # access token有效期(秒)
            }
        """
        access_token = JWTAuthentication.create_access_token(user)
        refresh_token = JWTAuthentication.create_refresh_token(user)
        
        return {
            'access': access_token,
            'refresh': refresh_token,
            'expires_in': int(JWTAuthentication.access_token_lifetime.total_seconds()),
            'token_type': 'Bearer',
        }
    
    @classmethod
    def refresh_access_token(cls, refresh_token):
        """
        刷新访问Token
        
        返回新的access token
        可选:同时刷新refresh token(滑动刷新)
        """
        # 检查黑名单
        if refresh_token in cls._blacklist:
            raise AuthenticationFailed('Token已被撤销')
        
        payload = JWTAuthentication.verify_refresh_token(refresh_token)
        user = User.objects.get(pk=payload['user_id'])
        
        return {
            'access': JWTAuthentication.create_access_token(user),
            'expires_in': int(JWTAuthentication.access_token_lifetime.total_seconds()),
        }
    
    @classmethod
    def revoke_token(cls, token):
        """
        撤销Token(加入黑名单)
        
        用于登出功能
        """
        cls._blacklist.add(token)
    
    @classmethod
    def is_token_revoked(cls, token):
        """检查Token是否已被撤销"""
        return token in cls._blacklist


# ===== 使用Redis存储黑名单(生产环境推荐)=====
class RedisTokenBlacklist:
    """
    基于Redis的Token黑名单
    
    优势:
    1. 持久化存储
    2. 自动过期
    3. 高性能
    """
    
    def __init__(self, redis_client=None):
        self.redis = redis_client
        self.prefix = 'jwt_blacklist:'
    
    def add(self, token, expires_in=None):
        """添加到黑名单"""
        key = f'{self.prefix}{token}'
        if expires_in:
            self.redis.setex(key, expires_in, '1')
        else:
            self.redis.set(key, '1')
    
    def contains(self, token):
        """检查是否在黑名单中"""
        key = f'{self.prefix}{token}'
        return bool(self.redis.exists(key))
    
    def remove(self, token):
        """从黑名单移除"""
        key = f'{self.prefix}{token}'
        self.redis.delete(key)


# ===== 登出视图 =====
@api_view(['POST'])
def jwt_logout(request):
    """
    登出视图
    
    撤销当前Token
    """
    refresh_token = request.data.get('refresh')
    
    if refresh_token:
        JWTTokenManager.revoke_token(refresh_token)
    
    # 也可以撤销access token(需要从请求头获取)
    auth_header = request.META.get('HTTP_AUTHORIZATION', '')
    if auth_header.startswith('Bearer '):
        access_token = auth_header.split()[1]
        JWTTokenManager.revoke_token(access_token)
    
    return Response({'message': '登出成功'})


# ===== URL配置 =====
"""
# urls.py
from django.urls import path

urlpatterns = [
    path('auth/login/', jwt_login, name='jwt-login'),
    path('auth/refresh/', jwt_refresh, name='jwt-refresh'),
    path('auth/logout/', jwt_logout, name='jwt-logout'),
]
"""
2.4.2 API Key认证
from rest_framework.authentication import BaseAuthentication
from rest_framework.exceptions import AuthenticationFailed
from django.contrib.auth import get_user_model
from django.db import models
from django.conf import settings
from django.utils import timezone
import hashlib
import time

User = get_user_model()


class APIKey(models.Model):
    """
    API Key模型
    """
    key = models.CharField(max_length=64, unique=True, primary_key=True)
    user = models.ForeignKey(
        settings.AUTH_USER_MODEL,
        on_delete=models.CASCADE,
        related_name='api_keys'
    )
    name = models.CharField(max_length=100)
    is_active = models.BooleanField(default=True)
    expires_at = models.DateTimeField(null=True, blank=True)
    last_used_at = models.DateTimeField(null=True, blank=True)
    created_at = models.DateTimeField(auto_now_add=True)
    
    # 权限范围
    scopes = models.JSONField(default=list)
    
    # 限流配置
    rate_limit = models.IntegerField(default=1000)  # 每小时请求数
    
    class Meta:
        verbose_name = 'API Key'
        verbose_name_plural = 'API Keys'


class APIKeyAuthentication(BaseAuthentication):
    """
    API Key认证
    
    请求头格式:
        X-API-Key: <api_key>
    
    或查询参数:
        ?api_key=<api_key>
    
    特点:
        1. 适合服务间调用
        2. 可设置权限范围
        3. 可设置过期时间
        4. 可追踪使用情况
    """
    
    keyword = 'X-API-Key'
    query_param = 'api_key'
    
    def authenticate(self, request):
        """
        认证流程:
        1. 从请求头或查询参数获取API Key
        2. 验证Key有效性
        3. 返回用户和Key对象
        """
        # 优先从请求头获取
        api_key = request.META.get(f'HTTP_{self.keyword.upper().replace("-", "_")}')
        
        # 其次从查询参数获取
        if not api_key:
            api_key = request.query_params.get(self.query_param)
        
        if not api_key:
            return None
        
        return self.authenticate_credentials(api_key)
    
    def authenticate_credentials(self, key):
        """
        验证API Key
        """
        try:
            api_key = APIKey.objects.select_related('user').get(key=key)
        except APIKey.DoesNotExist:
            raise AuthenticationFailed('无效的API Key')
        
        # 检查是否激活
        if not api_key.is_active:
            raise AuthenticationFailed('API Key已被禁用')
        
        # 检查是否过期
        if api_key.expires_at and api_key.expires_at < timezone.now():
            raise AuthenticationFailed('API Key已过期')
        
        # 检查用户状态
        if not api_key.user.is_active:
            raise AuthenticationFailed('用户已被禁用')
        
        # 更新最后使用时间
        api_key.last_used_at = timezone.now()
        api_key.save(update_fields=['last_used_at'])
        
        return (api_key.user, api_key)
    
    def authenticate_header(self, request):
        return self.keyword


# ===== API Key生成 =====
def generate_api_key():
    """生成API Key"""
    import secrets
    import string
    
    # 生成32字符的随机字符串
    alphabet = string.ascii_letters + string.digits
    random_str = ''.join(secrets.choice(alphabet) for _ in range(32))
    
    # 添加前缀和时间戳
    timestamp = str(int(time.time()))
    raw_key = f"sk_{timestamp}_{random_str}"
    
    # 可选:进行哈希处理
    # hashed = hashlib.sha256(raw_key.encode()).hexdigest()
    
    return raw_key


# ===== 创建API Key视图 =====
@api_view(['POST'])
def create_api_key(request):
    """
    创建API Key
    
    请求体:
        {
            "name": "My App",
            "scopes": ["read", "write"],
            "expires_in_days": 365
        }
    """
    name = request.data.get('name')
    scopes = request.data.get('scopes', [])
    expires_in_days = request.data.get('expires_in_days')
    
    if not name:
        return Response({'error': '名称不能为空'}, status=400)
    
    # 生成Key
    key = generate_api_key()
    
    # 计算过期时间
    expires_at = None
    if expires_in_days:
        expires_at = timezone.now() + timedelta(days=expires_in_days)
    
    # 创建API Key
    api_key = APIKey.objects.create(
        key=key,
        user=request.user,
        name=name,
        scopes=scopes,
        expires_at=expires_at,
    )
    
    return Response({
        'key': api_key.key,
        'name': api_key.name,
        'scopes': api_key.scopes,
        'expires_at': api_key.expires_at,
    })

2.5 多认证方式组合

from rest_framework.authentication import (
    SessionAuthentication,
    TokenAuthentication,
    BasicAuthentication,
)
from rest_framework.permissions import IsAuthenticated
from rest_framework import viewsets


class BookViewSet(viewsets.ModelViewSet):
    """
    多认证方式组合示例
    
    认证顺序:
    1. SessionAuthentication
    2. TokenAuthentication
    3. BasicAuthentication
    
    只要有一个认证成功即可
    """
    
    authentication_classes = [
        SessionAuthentication,
        TokenAuthentication,
        BasicAuthentication,
        JWTAuthentication,
    ]
    
    permission_classes = [IsAuthenticated]
    queryset = Book.objects.all()
    serializer_class = BookSerializer


# ===== 根据请求类型选择认证方式 =====
class SmartAuthentication(BaseAuthentication):
    """
    智能认证
    
    根据请求特征自动选择认证方式
    """
    
    def authenticate(self, request):
        auth_header = request.META.get('HTTP_AUTHORIZATION', '')
        
        # JWT认证
        if auth_header.lower().startswith('bearer '):
            return JWTAuthentication().authenticate(request)
        
        # Token认证
        if auth_header.lower().startswith('token '):
            return TokenAuthentication().authenticate(request)
        
        # Basic认证
        if auth_header.lower().startswith('basic '):
            return BasicAuthentication().authenticate(request)
        
        # Session认证(浏览器请求)
        if request.META.get('HTTP_COOKIE'):
            return SessionAuthentication().authenticate(request)
        
        # API Key认证
        if request.META.get('HTTP_X_API_KEY'):
            return APIKeyAuthentication().authenticate(request)
        
        return None


# ===== 条件认证 =====
class ConditionalAuthentication(BaseAuthentication):
    """
    条件认证
    
    根据视图或请求条件选择认证方式
    """
    
    def authenticate(self, request):
        view = request.parser_context.get('view')
        
        # 检查视图是否指定了特定认证方式
        if hasattr(view, 'required_auth_type'):
            auth_type = view.required_auth_type
            
            if auth_type == 'jwt':
                return JWTAuthentication().authenticate(request)
            elif auth_type == 'token':
                return TokenAuthentication().authenticate(request)
        
        # 默认使用所有认证方式
        for auth_class in [JWTAuthentication, TokenAuthentication]:
            try:
                result = auth_class().authenticate(request)
                if result:
                    return result
            except AuthenticationFailed:
                continue
        
        return None

三、权限系统深度解析

3.1 权限系统架构

权限检查流程

请求

check_permissions

has_permission?

执行视图

返回403

get_object

check_object_permissions

has_object_permission?

继续执行

返回403

3.2 BasePermission源码解析

from rest_framework.permissions import BasePermission


class BasePermission:
    """
    权限基类
    
    所有权限类必须继承此类
    
    方法:
        has_permission: 视图级权限检查
        has_object_permission: 对象级权限检查
    
    返回值:
        True: 允许访问
        False: 拒绝访问
    """
    
    def has_permission(self, request, view):
        """
        视图级权限检查
        
        参数:
            request: Request对象
            view: 视图对象
            
        返回:
            bool: 是否有权限
            
        说明:
            - 每个请求都会执行此方法
            - 在get_object()之前执行
            - 适用于列表、创建等操作
        """
        return True
    
    def has_object_permission(self, request, view, obj):
        """
        对象级权限检查
        
        参数:
            request: Request对象
            view: 视图对象
            obj: 要检查的对象
            
        返回:
            bool: 是否有权限
            
        说明:
            - 仅在get_object()时执行
            - 适用于详情、更新、删除等操作
            - 必须先通过has_permission
        """
        return True


# ===== 权限检查源码 =====
class APIView:
    def check_permissions(self, request):
        """
        检查视图级权限
        
        遍历所有权限类,任一返回False则拒绝
        """
        for permission in self.get_permissions():
            if not permission.has_permission(request, self):
                self.permission_denied(
                    request,
                    message=getattr(permission, 'message', None),
                    code=getattr(permission, 'code', None)
                )
    
    def check_object_permissions(self, request, obj):
        """
        检查对象级权限
        """
        for permission in self.get_permissions():
            if not permission.has_object_permission(request, self, obj):
                self.permission_denied(
                    request,
                    message=getattr(permission, 'message', None),
                    code=getattr(permission, 'code', None)
                )
    
    def get_permissions(self):
        """获取权限类实例列表"""
        return [permission() for permission in self.permission_classes]

3.3 内置权限类详解

from rest_framework.permissions import (
    AllowAny,
    IsAuthenticated,
    IsAdminUser,
    IsAuthenticatedOrReadOnly,
    DjangoModelPermissions,
    DjangoObjectPermissions,
)


class AllowAny(BasePermission):
    """
    允许所有用户
    
    无任何限制
    """
    def has_permission(self, request, view):
        return True


class IsAuthenticated(BasePermission):
    """
    仅允许已认证用户
    
    最常用的权限类
    """
    def has_permission(self, request, view):
        return bool(request.user and request.user.is_authenticated)


class IsAdminUser(BasePermission):
    """
    仅允许管理员用户
    
    检查user.is_staff
    """
    def has_permission(self, request, view):
        return bool(request.user and request.user.is_staff)


class IsAuthenticatedOrReadOnly(BasePermission):
    """
    已认证用户可写,匿名用户只读
    
    适用于公开API
    """
    def has_permission(self, request, view):
        # SAFE_METHODS = ('GET', 'HEAD', 'OPTIONS')
        if request.method in SAFE_METHODS:
            return True
        return bool(request.user and request.user.is_authenticated)


class DjangoModelPermissions(BasePermission):
    """
    Django模型权限
    
    基于Django的权限系统:
    - add: 创建权限
    - change: 更新权限
    - delete: 删除权限
    - view: 查看权限
    
    权限格式:<app_label>.<action>_<model_name>
    示例:books.add_book, books.change_book
    """
    
    perms_map = {
        'GET': ['%(app_label)s.view_%(model_name)s'],
        'OPTIONS': [],
        'HEAD': [],
        'POST': ['%(app_label)s.add_%(model_name)s'],
        'PUT': ['%(app_label)s.change_%(model_name)s'],
        'PATCH': ['%(app_label)s.change_%(model_name)s'],
        'DELETE': ['%(app_label)s.delete_%(model_name)s'],
    }
    
    def get_required_permissions(self, method, model_cls):
        """
        获取所需权限列表
        """
        kwargs = {
            'app_label': model_cls._meta.app_label,
            'model_name': model_cls._meta.model_name
        }
        
        return [perm % kwargs for perm in self.perms_map.get(method, [])]
    
    def has_permission(self, request, view):
        # 获取模型类
        model_cls = getattr(view, 'queryset', None)
        if model_cls is not None:
            model_cls = model_cls.model
        else:
            return True
        
        # 获取所需权限
        perms = self.get_required_permissions(request.method, model_cls)
        
        if not perms:
            return True
        
        # 检查用户权限
        return request.user.has_perms(perms)


class DjangoObjectPermissions(DjangoModelPermissions):
    """
    Django对象权限
    
    在模型权限基础上增加对象级权限检查
    需要配合django-guardian等第三方库使用
    """
    
    def has_object_permission(self, request, view, obj):
        # 获取所需权限
        perms = self.get_required_permissions(request.method, obj.__class__)
        
        if not perms:
            return True
        
        # 检查对象权限
        return request.user.has_perms(perms, obj)

3.4 自定义权限类

from rest_framework.permissions import BasePermission, SAFE_METHODS


class IsOwner(BasePermission):
    """
    对象所有者权限
    
    只有对象的所有者才能访问
    """
    
    message = '只有所有者才能执行此操作'
    
    def has_permission(self, request, view):
        return request.user.is_authenticated
    
    def has_object_permission(self, request, view, obj):
        # 检查对象是否有owner属性
        if hasattr(obj, 'owner'):
            return obj.owner == request.user
        if hasattr(obj, 'user'):
            return obj.user == request.user
        if hasattr(obj, 'author'):
            return obj.author == request.user
        return False


class IsOwnerOrAdmin(BasePermission):
    """
    所有者或管理员权限
    
    管理员可以访问所有对象
    普通用户只能访问自己的对象
    """
    
    message = '权限不足'
    
    def has_permission(self, request, view):
        return request.user.is_authenticated
    
    def has_object_permission(self, request, view, obj):
        # 管理员有所有权限
        if request.user.is_staff:
            return True
        
        # 检查所有权
        owner = getattr(obj, 'owner', None) or \
                getattr(obj, 'user', None) or \
                getattr(obj, 'author', None)
        
        return owner == request.user


class IsOwnerOrReadOnly(BasePermission):
    """
    所有者可写,其他人只读
    
    适用于博客、评论等内容
    """
    
    def has_permission(self, request, view):
        if request.method in SAFE_METHODS:
            return True
        return request.user.is_authenticated
    
    def has_object_permission(self, request, view, obj):
        # 只读方法允许所有
        if request.method in SAFE_METHODS:
            return True
        
        # 写操作检查所有权
        owner = getattr(obj, 'owner', None) or \
                getattr(obj, 'user', None) or \
                getattr(obj, 'author', None)
        
        return owner == request.user


class HasRolePermission(BasePermission):
    """
    角色权限
    
    基于用户角色进行权限控制
    """
    
    def __init__(self, required_roles=None):
        self.required_roles = required_roles or []
    
    def has_permission(self, request, view):
        if not request.user.is_authenticated:
            return False
        
        # 获取用户角色
        user_roles = getattr(request.user, 'roles', [])
        if hasattr(user_roles, 'all'):
            user_roles = list(user_roles.values_list('name', flat=True))
        
        # 检查是否有所需角色
        return any(role in user_roles for role in self.required_roles)


class HasPermission(BasePermission):
    """
    权限字符串检查
    
    检查用户是否有特定权限
    """
    
    def __init__(self, permission=None):
        self.permission = permission
    
    def has_permission(self, request, view):
        if not self.permission:
            return True
        return request.user.has_perm(self.permission)


class IsVerifiedUser(BasePermission):
    """
    已验证用户权限
    
    检查用户是否已完成邮箱/手机验证
    """
    
    message = '请先完成账户验证'
    
    def has_permission(self, request, view):
        if not request.user.is_authenticated:
            return False
        return getattr(request.user, 'is_verified', True)


class IsPremiumUser(BasePermission):
    """
    高级用户权限
    
    检查用户是否有高级会员资格
    """
    
    message = '此功能仅限高级会员使用'
    
    def has_permission(self, request, view):
        if not request.user.is_authenticated:
            return False
        
        # 检查会员状态
        if hasattr(request.user, 'membership'):
            membership = request.user.membership
            return membership.is_active and membership.is_premium
        
        return False


class TimeRestrictedPermission(BasePermission):
    """
    时间限制权限
    
    只在特定时间段允许访问
    """
    
    message = '当前时间不允许访问'
    
    def __init__(self, start_hour=9, end_hour=18):
        self.start_hour = start_hour
        self.end_hour = end_hour
    
    def has_permission(self, request, view):
        from datetime import datetime
        current_hour = datetime.now().hour
        return self.start_hour <= current_hour < self.end_hour


class IPWhitelistPermission(BasePermission):
    """
    IP白名单权限
    
    只允许特定IP访问
    """
    
    message = 'IP不在白名单中'
    
    def __init__(self, allowed_ips=None):
        self.allowed_ips = allowed_ips or []
    
    def has_permission(self, request, view):
        client_ip = self.get_client_ip(request)
        return client_ip in self.allowed_ips
    
    def get_client_ip(self, request):
        x_forwarded_for = request.META.get('HTTP_X_FORWARDED_FOR')
        if x_forwarded_for:
            return x_forwarded_for.split(',')[0].strip()
        return request.META.get('REMOTE_ADDR')


# ===== 组合权限 =====
class AndPermission(BasePermission):
    """
    AND组合权限
    
    所有权限都必须通过
    """
    
    def __init__(self, *permissions):
        self.permissions = permissions
    
    def has_permission(self, request, view):
        return all(
            perm().has_permission(request, view)
            for perm in self.permissions
        )
    
    def has_object_permission(self, request, view, obj):
        return all(
            perm().has_object_permission(request, view, obj)
            for perm in self.permissions
        )


class OrPermission(BasePermission):
    """
    OR组合权限
    
    任一权限通过即可
    """
    
    def __init__(self, *permissions):
        self.permissions = permissions
    
    def has_permission(self, request, view):
        return any(
            perm().has_permission(request, view)
            for perm in self.permissions
        )
    
    def has_object_permission(self, request, view, obj):
        return any(
            perm().has_object_permission(request, view, obj)
            for perm in self.permissions
        )


# ===== 使用示例 =====
class BookViewSet(viewsets.ModelViewSet):
    queryset = Book.objects.all()
    serializer_class = BookSerializer
    
    # 方式1:单一权限
    permission_classes = [IsAuthenticated]
    
    # 方式2:组合权限
    # permission_classes = [IsAuthenticated, IsOwnerOrAdmin]
    
    # 方式3:动态权限
    def get_permissions(self):
        if self.action == 'list':
            return [AllowAny()]
        elif self.action == 'create':
            return [IsAuthenticated(), IsVerifiedUser()]
        elif self.action in ['update', 'partial_update']:
            return [IsAuthenticated(), IsOwnerOrAdmin()]
        elif self.action == 'destroy':
            return [IsAuthenticated(), IsAdminUser()]
        return [IsAuthenticated()]

四、限流系统深度解析

4.1 限流系统架构

限流检查流程

请求

check_throttles

遍历throttle_classes

allow_request?

继续执行

返回429

返回wait时间

4.2 BaseThrottle源码解析

from rest_framework.throttling import BaseThrottle


class BaseThrottle:
    """
    限流基类
    
    所有限流类必须继承此类
    
    方法:
        allow_request: 是否允许请求
        get_ident: 获取客户端标识
        wait: 返回需要等待的秒数
    """
    
    def allow_request(self, request, view):
        """
        是否允许请求
        
        返回:
            True: 允许
            False: 拒绝
        """
        raise NotImplementedError('.allow_request() must be overridden')
    
    def get_ident(self, request):
        """
        获取客户端标识
        
        用于区分不同客户端
        默认使用IP地址
        """
        # 检查代理
        xff = request.META.get('HTTP_X_FORWARDED_FOR')
        
        if xff:
            # 使用最后一个代理的IP
            return xff.split(',')[-1].strip()
        
        return request.META.get('REMOTE_ADDR')
    
    def wait(self):
        """
        返回需要等待的秒数
        
        在请求被拒绝时调用
        """
        return None


# ===== 限流检查源码 =====
class APIView:
    def check_throttles(self, request):
        """
        检查限流
        """
        for throttle in self.get_throttles():
            if not throttle.allow_request(request, self):
                self.throttled(request, throttle.wait())
    
    def get_throttles(self):
        """获取限流类实例列表"""
        return [throttle() for throttle in self.throttle_classes]
    
    def throttled(self, request, wait):
        """抛出限流异常"""
        from rest_framework.exceptions import Throttled
        raise Throttled(wait)

4.3 SimpleRateThrottle滑动窗口算法

import time
from rest_framework.throttling import SimpleRateThrottle


class SimpleRateThrottle(BaseThrottle):
    """
    简单速率限流
    
    使用滑动窗口算法实现
    
    原理:
        1. 记录每个时间窗口内的请求时间戳
        2. 请求时移除过期的记录
        3. 检查当前窗口内的请求数是否超限
    
    速率格式:
        's': 秒
        'm': 分钟
        'h': 小时
        'd': 天
    
    示例:
        '100/hour': 每小时100次
        '10/minute': 每分钟10次
        '1000/day': 每天1000次
    """
    
    cache = None
    timer = time.time
    cache_format = 'throttle_%(ident)s_%(scope)s'
    scope = None
    THROTTLE_RATES = {}
    
    def __init__(self):
        if not getattr(self, 'rate', None):
            self.rate = self.get_rate()
        
        self.num_requests, self.duration = self.parse_rate(self.rate)
    
    def get_rate(self):
        """
        获取速率配置
        
        从settings中获取对应scope的速率
        """
        if not self.scope:
            raise ValueError('必须设置scope属性')
        
        if self.scope not in self.THROTTLE_RATES:
            raise ValueError(f'未找到scope={self.scope}的速率配置')
        
        return self.THROTTLE_RATES[self.scope]
    
    def parse_rate(self, rate):
        """
        解析速率字符串
        
        '100/hour' → (100, 3600)
        '10/minute' → (10, 60)
        '1000/day' → (1000, 86400)
        """
        if rate is None:
            return None, None
        
        num, period = rate.split('/')
        num_requests = int(num)
        
        duration = {
            's': 1,
            'sec': 1,
            'm': 60,
            'min': 60,
            'h': 3600,
            'hour': 3600,
            'd': 86400,
            'day': 86400,
        }.get(period.lower(), 0)
        
        return num_requests, duration
    
    def allow_request(self, request, view):
        """
        是否允许请求
        
        滑动窗口算法实现
        """
        if self.rate is None:
            return True
        
        # 获取缓存键
        self.key = self.get_cache_key(request, view)
        if self.key is None:
            return True
        
        # 获取历史记录
        self.history = self.cache.get(self.key, [])
        self.now = self.timer()
        
        # 滑动窗口:移除过期记录
        while self.history and self.history[-1] <= self.now - self.duration:
            self.history.pop()
        
        # 检查是否超限
        if len(self.history) >= self.num_requests:
            return False
        
        # 记录本次请求
        self.history.insert(0, self.now)
        self.cache.set(self.key, self.history, self.duration)
        
        return True
    
    def get_cache_key(self, request, view):
        """
        获取缓存键
        
        子类必须实现此方法
        """
        raise NotImplementedError('.get_cache_key() must be overridden')
    
    def wait(self):
        """
        返回需要等待的秒数
        """
        if self.history:
            remaining_duration = self.duration - (self.now - self.history[-1])
        else:
            remaining_duration = self.duration
        
        available_requests = self.num_requests - len(self.history)
        if available_requests <= 0:
            return remaining_duration
        
        return None


# ===== 滑动窗口算法图解 =====
"""
滑动窗口算法示例(限制:5次/分钟)

时间线:
|----60秒窗口----|
0    15    30    45    60    75    90   (秒)

请求记录(假设当前时间90秒):
history = [85, 80, 70, 50, 40, 30, 20, 10]

步骤1:移除过期记录(超过60秒的)
当前时间:90秒
窗口起点:90 - 60 = 30秒
移除:20, 10
剩余:[85, 80, 70, 50, 40, 30]

步骤2:检查数量
len(history) = 6 >= 5(限制)
结果:拒绝请求

步骤3:计算等待时间
最旧记录:30秒
过期时间:30 + 60 = 90秒
等待时间:90 - 90 = 0秒(立即可以重试)

如果最旧记录是35秒:
过期时间:35 + 60 = 95秒
等待时间:95 - 90 = 5秒
"""

4.4 内置限流类详解

from rest_framework.throttling import (
    AnonRateThrottle,
    UserRateThrottle,
    ScopedRateThrottle,
)


class AnonRateThrottle(SimpleRateThrottle):
    """
    匿名用户限流
    
    基于IP地址限流
    适用于未认证用户
    
    配置:
        REST_FRAMEWORK = {
            'DEFAULT_THROTTLE_RATES': {
                'anon': '100/hour',
            }
        }
    """
    
    scope = 'anon'
    
    def get_cache_key(self, request, view):
        # 只对未认证用户生效
        if request.user.is_authenticated:
            return None
        
        return self.cache_format % {
            'ident': self.get_ident(request),
            'scope': self.scope,
        }


class UserRateThrottle(SimpleRateThrottle):
    """
    认证用户限流
    
    基于用户ID限流
    适用于已认证用户
    
    配置:
        REST_FRAMEWORK = {
            'DEFAULT_THROTTLE_RATES': {
                'user': '1000/hour',
            }
        }
    """
    
    scope = 'user'
    
    def get_cache_key(self, request, view):
        # 只对认证用户生效
        if not request.user.is_authenticated:
            return None
        
        return self.cache_format % {
            'ident': request.user.pk,
            'scope': self.scope,
        }


class ScopedRateThrottle(SimpleRateThrottle):
    """
    作用域限流
    
    基于视图的throttle_scope属性限流
    适用于特定功能的限流
    
    配置:
        REST_FRAMEWORK = {
            'DEFAULT_THROTTLE_RATES': {
                'contacts': '10/hour',
                'uploads': '20/day',
            }
        }
        
        class ContactViewSet(viewsets.ModelViewSet):
            throttle_scope = 'contacts'
    """
    
    scope_attr = 'throttle_scope'
    
    def __init__(self):
        pass  # 延迟初始化
    
    def allow_request(self, request, view):
        # 获取视图的scope
        scope = getattr(view, self.scope_attr, None)
        
        if not scope:
            return True
        
        # 动态设置scope和rate
        self.scope = scope
        self.rate = self.get_rate()
        self.num_requests, self.duration = self.parse_rate(self.rate)
        
        return super().allow_request(request, view)
    
    def get_cache_key(self, request, view):
        ident = request.user.pk if request.user.is_authenticated else self.get_ident(request)
        
        return self.cache_format % {
            'ident': ident,
            'scope': self.scope,
        }

4.5 自定义限流类

from rest_framework.throttling import SimpleRateThrottle, BaseThrottle
from django.core.cache import cache
import time


class IPRateThrottle(SimpleRateThrottle):
    """
    IP限流
    
    基于IP地址限流,不区分用户认证状态
    """
    
    scope = 'ip'
    
    def get_cache_key(self, request, view):
        return self.cache_format % {
            'ident': self.get_ident(request),
            'scope': self.scope,
        }


class UserIPRateThrottle(SimpleRateThrottle):
    """
    用户+IP组合限流
    
    已认证用户:基于用户ID
    未认证用户:基于IP
    """
    
    scope = 'user_ip'
    
    def get_cache_key(self, request, view):
        if request.user.is_authenticated:
            ident = f'user_{request.user.pk}'
        else:
            ident = f'ip_{self.get_ident(request)}'
        
        return self.cache_format % {
            'ident': ident,
            'scope': self.scope,
        }


class BurstRateThrottle(UserRateThrottle):
    """
    突发限流
    
    短时间内允许突发请求
    """
    scope = 'burst'


class SustainedRateThrottle(UserRateThrottle):
    """
    持续限流
    
    长时间内的总体限流
    """
    scope = 'sustained'


class MethodRateThrottle(SimpleRateThrottle):
    """
    HTTP方法限流
    
    不同的HTTP方法使用不同的限流策略
    """
    
    def get_cache_key(self, request, view):
        # 根据HTTP方法设置不同的scope
        method = request.method.lower()
        self.scope = f'{method}_rate'
        
        ident = request.user.pk if request.user.is_authenticated else self.get_ident(request)
        
        return self.cache_format % {
            'ident': ident,
            'scope': self.scope,
        }


class PathRateThrottle(SimpleRateThrottle):
    """
    路径限流
    
    不同的API路径使用不同的限流策略
    """
    
    def get_cache_key(self, request, view):
        # 使用请求路径作为scope
        path = request.path.replace('/', '_').strip('_')
        self.scope = f'path_{path}'
        
        ident = request.user.pk if request.user.is_authenticated else self.get_ident(request)
        
        return self.cache_format % {
            'ident': ident,
            'scope': self.scope,
        }


class RoleBasedRateThrottle(SimpleRateThrottle):
    """
    角色限流
    
    不同角色使用不同的限流策略
    """
    
    def get_cache_key(self, request, view):
        if not request.user.is_authenticated:
            self.scope = 'anon'
            ident = self.get_ident(request)
        elif request.user.is_staff:
            self.scope = 'admin'
            ident = request.user.pk
        elif hasattr(request.user, 'is_premium') and request.user.is_premium:
            self.scope = 'premium'
            ident = request.user.pk
        else:
            self.scope = 'user'
            ident = request.user.pk
        
        self.rate = self.get_rate()
        self.num_requests, self.duration = self.parse_rate(self.rate)
        
        return self.cache_format % {
            'ident': ident,
            'scope': self.scope,
        }


class TokenBucketThrottle(BaseThrottle):
    """
    令牌桶算法限流
    
    原理:
        1. 以固定速率向桶中添加令牌
        2. 每个请求消耗一个令牌
        3. 桶满时令牌溢出
        4. 无令牌时拒绝请求
    
    优点:
        - 允许一定程度的突发流量
        - 平滑处理请求
    """
    
    def __init__(self, rate=10, capacity=20):
        """
        参数:
            rate: 令牌生成速率(个/秒)
            capacity: 桶容量
        """
        self.rate = rate
        self.capacity = capacity
        self.cache = cache
    
    def get_cache_key(self, request, view):
        ident = request.user.pk if request.user.is_authenticated else self.get_ident(request)
        return f'token_bucket_{ident}'
    
    def allow_request(self, request, view):
        key = self.get_cache_key(request, view)
        now = time.time()
        
        # 获取当前令牌数和上次更新时间
        bucket = self.cache.get(key, {'tokens': self.capacity, 'last_update': now})
        
        tokens = bucket['tokens']
        last_update = bucket['last_update']
        
        # 计算新增令牌
        elapsed = now - last_update
        new_tokens = elapsed * self.rate
        tokens = min(tokens + new_tokens, self.capacity)
        
        # 检查是否有令牌
        if tokens >= 1:
            tokens -= 1
            self.cache.set(key, {'tokens': tokens, 'last_update': now}, timeout=3600)
            return True
        
        # 无令牌,拒绝请求
        self.cache.set(key, {'tokens': tokens, 'last_update': now}, timeout=3600)
        return False
    
    def wait(self):
        """返回需要等待的秒数"""
        return 1.0 / self.rate


class LeakyBucketThrottle(BaseThrottle):
    """
    漏桶算法限流
    
    原理:
        1. 请求进入桶中排队
        2. 以固定速率从桶中流出处理
        3. 桶满时拒绝新请求
    
    优点:
        - 严格限制流出速率
        - 适合流量整形
    """
    
    def __init__(self, rate=10, capacity=20):
        """
        参数:
            rate: 处理速率(个/秒)
            capacity: 桶容量
        """
        self.rate = rate
        self.capacity = capacity
        self.cache = cache
    
    def get_cache_key(self, request, view):
        ident = request.user.pk if request.user.is_authenticated else self.get_ident(request)
        return f'leaky_bucket_{ident}'
    
    def allow_request(self, request, view):
        key = self.get_cache_key(request, view)
        now = time.time()
        
        # 获取桶状态
        bucket = self.cache.get(key, {'water': 0, 'last_leak': now})
        
        water = bucket['water']
        last_leak = bucket['last_leak']
        
        # 计算漏出的水量
        elapsed = now - last_leak
        leaked = elapsed * self.rate
        water = max(water - leaked, 0)
        
        # 检查桶是否已满
        if water < self.capacity:
            water += 1
            self.cache.set(key, {'water': water, 'last_leak': now}, timeout=3600)
            return True
        
        # 桶满,拒绝请求
        self.cache.set(key, {'water': water, 'last_leak': now}, timeout=3600)
        return False
    
    def wait(self):
        """返回需要等待的秒数"""
        return 1.0 / self.rate


# ===== 使用示例 =====
class BookViewSet(viewsets.ModelViewSet):
    queryset = Book.objects.all()
    serializer_class = BookSerializer
    
    # 方式1:全局限流
    throttle_classes = [AnonRateThrottle, UserRateThrottle]
    
    # 方式2:作用域限流
    throttle_scope = 'books'
    
    # 方式3:动态限流
    def get_throttles(self):
        if self.action == 'create':
            return [BurstRateThrottle(), SustainedRateThrottle()]
        return [UserRateThrottle()]

4.6 限流配置详解

# settings.py

REST_FRAMEWORK = {
    # ===== 限流类配置 =====
    'DEFAULT_THROTTLE_CLASSES': [
        'rest_framework.throttling.AnonRateThrottle',
        'rest_framework.throttling.UserRateThrottle',
    ],
    
    # ===== 限流速率配置 =====
    'DEFAULT_THROTTLE_RATES': {
        # 匿名用户限流
        'anon': '100/hour',
        
        # 认证用户限流
        'user': '1000/hour',
        
        # 突发限流
        'burst': '60/minute',
        
        # 持续限流
        'sustained': '10000/day',
        
        # 作用域限流
        'contacts': '10/hour',
        'uploads': '20/day',
        
        # 角色限流
        'admin': '10000/hour',
        'premium': '5000/hour',
        
        # HTTP方法限流
        'get_rate': '300/minute',
        'post_rate': '30/minute',
        'put_rate': '30/minute',
        'delete_rate': '10/minute',
    },
    
    # ===== 代理配置 =====
    'NUM_PROXIES': None,
    """
    NUM_PROXIES配置说明:
    
    当应用部署在反向代理(如Nginx)后面时,
    需要配置代理数量以正确获取客户端IP。
    
    示例:
    - 无代理:NUM_PROXIES = None 或 0
    - 1层代理:NUM_PROXIES = 1
    - 2层代理:NUM_PROXIES = 2
    
    X-Forwarded-For格式:
    X-Forwarded-For: client, proxy1, proxy2
    
    获取真实IP:
    proxies = NUM_PROXIES or 0
    client_ip = xff.split(',')[-proxies - 1]
    """
}


# ===== 视图级限流配置 =====
class ContactViewSet(viewsets.ModelViewSet):
    queryset = Contact.objects.all()
    serializer_class = ContactSerializer
    
    # 覆盖限流类
    throttle_classes = [ScopedRateThrottle]
    
    # 设置作用域
    throttle_scope = 'contacts'


class UploadViewSet(viewsets.ModelViewSet):
    queryset = Upload.objects.all()
    serializer_class = UploadSerializer
    
    # 多重限流
    throttle_classes = [BurstRateThrottle, SustainedRateThrottle]

五、REST_FRAMEWORK安全配置详解

# settings.py

REST_FRAMEWORK = {
    # ===== 认证配置 =====
    'DEFAULT_AUTHENTICATION_CLASSES': [
        'rest_framework.authentication.SessionAuthentication',
        'rest_framework.authentication.TokenAuthentication',
        # 'rest_framework.authentication.BasicAuthentication',
        # 'core.authentication.JWTAuthentication',
    ],
    
    # 未认证用户的用户类
    'UNAUTHENTICATED_USER': 'django.contrib.auth.models.AnonymousUser',
    
    # 未认证用户的认证凭证
    'UNAUTHENTICATED_TOKEN': None,
    
    # ===== 权限配置 =====
    'DEFAULT_PERMISSION_CLASSES': [
        'rest_framework.permissions.IsAuthenticated',
        # 'rest_framework.permissions.AllowAny',
    ],
    
    # ===== 限流配置 =====
    'DEFAULT_THROTTLE_CLASSES': [
        'rest_framework.throttling.AnonRateThrottle',
        'rest_framework.throttling.UserRateThrottle',
    ],
    
    'DEFAULT_THROTTLE_RATES': {
        'anon': '100/hour',
        'user': '1000/hour',
    },
    
    # 代理数量
    'NUM_PROXIES': None,
    
    # ===== 异常处理 =====
    'EXCEPTION_HANDLER': 'rest_framework.views.exception_handler',
    
    # 认证失败时的响应
    # 'UNAUTHENTICATED_USER': None,  # 返回401而非403
}

总结

Django REST Framework的安全体系提供了完整的认证、权限、限流解决方案:

  1. 认证:验证用户身份,支持多种认证方式(Basic、Token、Session、JWT等)
  2. 权限:控制访问权限,支持视图级和对象级权限检查
  3. 限流:控制请求频率,支持滑动窗口、令牌桶、漏桶等算法

关键要点:

  • 认证是确定"你是谁"
  • 权限是确定"你能做什么"
  • 限流是确定"你能做多少次"
  • 三者按顺序执行,任一失败都会阻止请求
Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐