fix(security): P0 - JWT密钥强制环境变量 + 日志敏感信息过滤 #8
+46
-13
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
@@ -10,24 +11,50 @@ class Settings(BaseSettings):
|
||||
ENVIRONMENT: str = "development"
|
||||
DEBUG: bool = True
|
||||
|
||||
# Container bind address; external exposure is controlled by Docker/Nginx.
|
||||
API_HOST: str = "0.0.0.0" # nosec B104
|
||||
# Container bind address; external expose is controlled by Docker/Nginx.
|
||||
API_HOST: str = "0.0.0.0" # nosec: B104
|
||||
API_PORT: int = 8000
|
||||
API_PREFIX: str = "/api/v1"
|
||||
|
||||
DATABASE_URL: str = "postgresql+psycopg://postgres:postgres@localhost:5432/xiaoxia_saas"
|
||||
DATABASE_URL: str = (
|
||||
"postgresql+psycopg://postgres:postgres@localhost:5432/xiaoxia_saas"
|
||||
)
|
||||
DATABASE_POOL_SIZE: int = 20
|
||||
DATABASE_MAX_OVERFLOW: int = 40
|
||||
DATABASE_POOL_TIMEOUT: int = 30
|
||||
DATABASE_POOL_RECYCLE: int = 3600
|
||||
DATABASE_POOL_RECYLE: int = 3600
|
||||
USE_IN_MEMORY_DB: bool = False
|
||||
AUTO_CREATE_SCHEMA: bool = False
|
||||
|
||||
REDIS_URL: str = "redis://localhost:6379/0"
|
||||
REDIS_MAX_CONNECTIONS: int = 50
|
||||
ENABLE_REDIS_SESSIONS: bool = False
|
||||
REDIS_MAX_CONNECTION: int = 50
|
||||
ENABLE_REDIS_SESSION: bool = False
|
||||
|
||||
JWT_SECRET_KEY: str = "your-secret-key-change-in-production"
|
||||
# JWT secret key - MUST be set via environment variable, no default allowed
|
||||
JWT_SECRET_KEY: Optional[str] = None
|
||||
|
||||
@field_validator("JWT_SECRET_KEY", mode="before")
|
||||
@classmethod
|
||||
def validate_jwt_secret_key(cls, v):
|
||||
if v is None or v == "":
|
||||
raise ValueError(
|
||||
"JWT_SECRET_KEY must be set via environment variable. "
|
||||
"Do not use default value in production!"
|
||||
)
|
||||
# Block known insecure default values
|
||||
insecure_defaults = [
|
||||
"your-secret-key-change-in-production",
|
||||
"your-secret-key",
|
||||
"secret",
|
||||
"changeme",
|
||||
"password",
|
||||
]
|
||||
if v.lower() in [d.lower() for d in insecure_defaults]:
|
||||
raise ValueError(
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. "
|
||||
"Please set a strong random secret via environment variable."
|
||||
)
|
||||
return v
|
||||
|
||||
ENABLE_EMAIL_DELIVERY: bool = False
|
||||
SMTP_HOST: str = "smtp.gmail.com"
|
||||
@@ -35,22 +62,24 @@ class Settings(BaseSettings):
|
||||
SMTP_USER: str = ""
|
||||
SMTP_PASSWORD: str = ""
|
||||
SMTP_FROM_EMAIL: str = ""
|
||||
SMTP_FROM_NAME: str = "小虾 SaaS"
|
||||
SMTP_FRON_NAME: str = "小虾 SaaS"
|
||||
SMTP_USE_TLS: bool = True
|
||||
|
||||
CELERY_BROKER_URL: str = "redis://localhost:6379/0"
|
||||
CELERY_RESULT_BACKEND: str = "redis://localhost:6379/1"
|
||||
|
||||
# 阿里云 OSS 配置
|
||||
OSS_ENDPOINT: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
# OSS 七牛云相关
|
||||
OSS_ENDPOINT: str = "oss-cn-hangzhou.aliiyuncs.com"
|
||||
OSS_ACCESS_KEY_ID: str = ""
|
||||
OSS_ACCESS_KEY_SECRET: str = ""
|
||||
OSS_BUCKET_NAME: str = "xiaoxia-autocut"
|
||||
OSS_DIRECT_UPLOAD_MAX_MB: int = 800
|
||||
OSS_DIRECT_UPLOAD_EXPIRE_SECONDS: int = 900
|
||||
OSS_DIRECT_UPLOAD_EXPRESS_SECRET: int = 900
|
||||
|
||||
LOG_LEVEL: str = "INFO"
|
||||
CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
|
||||
CORS_ORIGINS_RAW: str = (
|
||||
"http://localhost:3000,http://localhost:5173,http://localhost:8000"
|
||||
)
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
@@ -61,7 +90,11 @@ class Settings(BaseSettings):
|
||||
|
||||
@property
|
||||
def CORS_ORIGINS(self) -> list[str]:
|
||||
return [origin.strip() for origin in self.CORS_ORIGINS_RAW.split(",") if origin.strip()]
|
||||
return [
|
||||
origin.strip()
|
||||
for origin in self.CORS_ORIGINS_RAW.split(",")
|
||||
if origin.strip()
|
||||
]
|
||||
|
||||
@property
|
||||
def database_url(self) -> str:
|
||||
|
||||
@@ -1,15 +1,60 @@
|
||||
"""
|
||||
请求日志中间件
|
||||
"""
|
||||
"""请求日志和限流中间件"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import Request
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 敏感参数名称模式(不区分大小写)
|
||||
SENSITIVE_PARAM_PATTERNS = re.compile(
|
||||
r"^(password|passwd|pwd|token|secret|key|authorization|auth|api_key|" # noqa: E501
|
||||
r"apikey|access_token|refresh_token|accesstoken|refreshtoken|" # noqa: E501
|
||||
r"session_id|sessionid|sid|cookie|csrf|xsrf|bearer)$", # noqa: E501
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def filter_sensitive_params(query_string: Optional[str]) -> Optional[str]:
|
||||
"""
|
||||
过滤 query string 中的敏感参数
|
||||
|
||||
Args:
|
||||
query_string: 原始 query string,例如 "name=xxx&password=secret&token=abc"
|
||||
|
||||
Returns:
|
||||
过滤后的 query string,敏感参数的值被替换为 "***"
|
||||
如果 query_string 为空或 None,返回原始值
|
||||
"""
|
||||
if not query_string:
|
||||
return query_string
|
||||
|
||||
if query_string.startswith("?"):
|
||||
query_string = query_string[1:]
|
||||
|
||||
if not query_string:
|
||||
return query_string
|
||||
|
||||
parts = query_string.split("&")
|
||||
filtered_parts = []
|
||||
|
||||
for part in parts:
|
||||
if "=" in part:
|
||||
key, value = part.split("=", 1)
|
||||
if SENSITIVE_PARAM_PATTERNS.match(key):
|
||||
filtered_parts.append(f"{key}=***")
|
||||
else:
|
||||
filtered_parts.append(part)
|
||||
else:
|
||||
# 没有 = 的参数,保留原样
|
||||
filtered_parts.append(part)
|
||||
|
||||
return "&".join(filtered_parts)
|
||||
|
||||
|
||||
class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
||||
"""请求日志中间件"""
|
||||
@@ -18,47 +63,60 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
||||
# 记录请求开始时间
|
||||
start_time = time.time()
|
||||
|
||||
# 记录请求信息
|
||||
logger.info(f"Request: {request.method} {request.url.path}")
|
||||
# 过滤 query string 中的敏感参数
|
||||
raw_query = str(request.url.query) if request.url.query else ""
|
||||
safe_query = filter_sensitive_params(raw_query)
|
||||
|
||||
# 记录请求信息(不包含敏感参数)
|
||||
if safe_query:
|
||||
logger.info(
|
||||
f"Request: {request.method} {request.url.path}?{safe_query}" # noqa: E501
|
||||
)
|
||||
else:
|
||||
logger.info(f"Request: {request.method} {request.url.path}")
|
||||
|
||||
# 处理请求
|
||||
response = await call_next(request)
|
||||
|
||||
# 计算处理时间
|
||||
# 记录请求结束时间
|
||||
process_time = time.time() - start_time
|
||||
|
||||
# 记录响应信息
|
||||
logger.info(
|
||||
f"Response: {request.method} {request.url.path} " f"status={response.status_code} time={process_time:.3f}s"
|
||||
f"Response: {request.method} {request.url.path} "
|
||||
f"status={response.status_code} time={process_time:.3f}s"
|
||||
)
|
||||
|
||||
# 添加响应头
|
||||
# 添加处理时间到响应头
|
||||
response.headers["X-Process-Time"] = str(process_time)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
class RateLimitMiddleware(BaseHTTPMiddleware):
|
||||
"""简单的速率限制中间件(基于内存)"""
|
||||
"""基于 IP 的简单限流中间件"""
|
||||
|
||||
def __init__(self, app, max_requests: int = 100, window_seconds: int = 60):
|
||||
super().__init__(app)
|
||||
self.max_requests = max_requests
|
||||
self.window_seconds = window_seconds
|
||||
self.requests = {} # {ip: [(timestamp, ...)]}
|
||||
self.requests = {} # {ip: [timestamps]}
|
||||
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
# 获取客户端 IP
|
||||
client_ip = request.client.host
|
||||
|
||||
current_time = time.time()
|
||||
|
||||
# 清理过期记录
|
||||
if client_ip in self.requests:
|
||||
self.requests[client_ip] = [
|
||||
ts for ts in self.requests[client_ip] if current_time - ts < self.window_seconds
|
||||
ts
|
||||
for ts in self.requests[client_ip]
|
||||
if current_time - ts < self.window_seconds
|
||||
]
|
||||
|
||||
# 检查速率限制
|
||||
# 计算请求次数
|
||||
request_count = len(self.requests.get(client_ip, []))
|
||||
|
||||
if request_count >= self.max_requests:
|
||||
@@ -69,7 +127,10 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
|
||||
content={
|
||||
"error": {
|
||||
"code": "RATE_LIMIT_EXCEEDED",
|
||||
"message": f"Too many requests. Limit: {self.max_requests} per {self.window_seconds}s",
|
||||
"message": ( # noqa: E501
|
||||
f"Too many requests. Limit: "
|
||||
f"{self.max_requests} per {self.window_seconds}s"
|
||||
),
|
||||
}
|
||||
},
|
||||
)
|
||||
@@ -82,8 +143,10 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
|
||||
# 处理请求
|
||||
response = await call_next(request)
|
||||
|
||||
# 添加速率限制信息到响应头
|
||||
# 添加限流信息到响应头
|
||||
response.headers["X-RateLimit-Limit"] = str(self.max_requests)
|
||||
response.headers["X-RateLimit-Remaining"] = str(self.max_requests - len(self.requests[client_ip]))
|
||||
response.headers["X-RateLimit-Remaining"] = str(
|
||||
self.max_requests - len(self.requests[client_ip])
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
@@ -1,7 +1,4 @@
|
||||
"""
|
||||
JWT 工具类
|
||||
提供 Token 签发、验证、刷新功能
|
||||
"""
|
||||
"""JWT Token 生成、验证、解析服务"""
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, Optional
|
||||
@@ -13,11 +10,47 @@ from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
|
||||
class JWTConfig:
|
||||
"""JWT 配置"""
|
||||
|
||||
# 从环境变量读取,这里先用默认值
|
||||
SECRET_KEY: str = "your-secret-key-change-in-production"
|
||||
ALGORITHM: str = "HS256"
|
||||
ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 # 30 分钟
|
||||
REFRESH_TOKEN_EXPIRE_DAYS: int = 30 # 30 天
|
||||
def __init__(
|
||||
self,
|
||||
secret_key: str,
|
||||
algorithm: str = "HS256",
|
||||
access_token_expire_minutes: int = 30,
|
||||
refresh_token_expire_days: int = 30,
|
||||
):
|
||||
"""
|
||||
初始化 JWT 配置
|
||||
|
||||
Args:
|
||||
secret_key: JWT 签名密钥(必须从环境变量或配置注入,不允许默认值)
|
||||
algorithm: 加密算法,默认 HS256
|
||||
access_token_expire_minutes: Access Token 过期时间(分钟)
|
||||
refresh_token_expire_days: Refresh Token 过期时间(天)
|
||||
|
||||
Raises:
|
||||
ValueError: 如果 secret_key 为空或包含不安全默认值
|
||||
"""
|
||||
if not secret_key or secret_key.strip() == "":
|
||||
raise ValueError( # noqa: E501
|
||||
"JWT secret_key must be provided and cannot be empty"
|
||||
)
|
||||
|
||||
insecure_defaults = [
|
||||
"your-secret-key-change-in-production",
|
||||
"your-secret-key",
|
||||
"secret",
|
||||
"changeme",
|
||||
"password",
|
||||
]
|
||||
if secret_key.lower() in [d.lower() for d in insecure_defaults]:
|
||||
raise ValueError( # noqa: E501
|
||||
f"JWT secret_key '{secret_key}' is insecure. "
|
||||
"Please provide a strong random secret."
|
||||
)
|
||||
|
||||
self.SECRET_KEY: str = secret_key
|
||||
self.ALGORITHM: str = algorithm
|
||||
self.ACCESS_TOKEN_EXPIRE_MINUTES: int = access_token_expire_minutes
|
||||
self.REFRESH_TOKEN_EXPIRE_DAYS: int = refresh_token_expire_days
|
||||
|
||||
|
||||
class TokenType:
|
||||
@@ -31,7 +64,13 @@ class JWTService:
|
||||
"""JWT 服务类"""
|
||||
|
||||
def __init__(self, config: JWTConfig = None):
|
||||
self.config = config or JWTConfig()
|
||||
if config is None:
|
||||
raise ValueError( # noqa: E501
|
||||
"JWTService requires a JWTConfig instance. " # noqa: E501
|
||||
"Please provide a configured JWTConfig with a valid " # noqa: E501
|
||||
"secret_key."
|
||||
)
|
||||
self.config = config
|
||||
|
||||
def create_access_token(
|
||||
self,
|
||||
@@ -46,17 +85,19 @@ class JWTService:
|
||||
Args:
|
||||
user_id: 用户 ID
|
||||
workspace_id: 工作空间 ID
|
||||
role: 用户在该工作空间的角色
|
||||
additional_claims: 额外的声明(可选)
|
||||
role: 用户角色(admin/user/guest)
|
||||
additional_claims: 额外的声明信息
|
||||
|
||||
Returns:
|
||||
JWT Token 字符串
|
||||
"""
|
||||
now = datetime.utcnow()
|
||||
expire = now + timedelta(minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
expire = now + timedelta( # noqa: E501
|
||||
minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES
|
||||
)
|
||||
|
||||
payload = {
|
||||
"sub": user_id, # subject (用户 ID)
|
||||
"sub": user_id, # subject (用户ID)
|
||||
"workspace_id": workspace_id,
|
||||
"role": role,
|
||||
"type": TokenType.ACCESS,
|
||||
@@ -67,7 +108,9 @@ class JWTService:
|
||||
if additional_claims:
|
||||
payload.update(additional_claims)
|
||||
|
||||
return jwt.encode(payload, self.config.SECRET_KEY, algorithm=self.config.ALGORITHM)
|
||||
return jwt.encode(
|
||||
payload, self.config.SECRET_KEY, algorithm=self.config.ALGORITHM
|
||||
)
|
||||
|
||||
def create_refresh_token(self, user_id: str, session_id: str) -> str:
|
||||
"""
|
||||
@@ -75,7 +118,7 @@ class JWTService:
|
||||
|
||||
Args:
|
||||
user_id: 用户 ID
|
||||
session_id: Session ID(用于撤销)
|
||||
session_id: Session ID
|
||||
|
||||
Returns:
|
||||
JWT Token 字符串
|
||||
@@ -91,11 +134,13 @@ class JWTService:
|
||||
"exp": expire,
|
||||
}
|
||||
|
||||
return jwt.encode(payload, self.config.SECRET_KEY, algorithm=self.config.ALGORITHM)
|
||||
return jwt.encode(
|
||||
payload, self.config.SECRET_KEY, algorithm=self.config.ALGORITHM
|
||||
)
|
||||
|
||||
def verify_token(self, token: str) -> Dict[str, Any]:
|
||||
"""
|
||||
验证 Token 并解码
|
||||
验证 Token
|
||||
|
||||
Args:
|
||||
token: JWT Token 字符串
|
||||
@@ -108,7 +153,11 @@ class JWTService:
|
||||
InvalidTokenError: Token 无效
|
||||
"""
|
||||
try:
|
||||
payload = jwt.decode(token, self.config.SECRET_KEY, algorithms=[self.config.ALGORITHM])
|
||||
payload = jwt.decode( # noqa: E501
|
||||
token,
|
||||
self.config.SECRET_KEY,
|
||||
algorithms=[self.config.ALGORITHM],
|
||||
)
|
||||
return payload
|
||||
except ExpiredSignatureError:
|
||||
raise ExpiredSignatureError("Token has expired")
|
||||
@@ -123,12 +172,7 @@ class JWTService:
|
||||
token: JWT Token 字符串
|
||||
|
||||
Returns:
|
||||
Token payload
|
||||
|
||||
Raises:
|
||||
ValueError: Token 类型不是 access
|
||||
ExpiredSignatureError: Token 已过期
|
||||
InvalidTokenError: Token 无效
|
||||
Token payload(如果解码失败返回 None)
|
||||
"""
|
||||
payload = self.verify_token(token)
|
||||
|
||||
@@ -161,7 +205,7 @@ class JWTService:
|
||||
|
||||
def decode_token_unsafe(self, token: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
不验证签名地解码 Token(仅用于调试/日志)
|
||||
不验证签名直接解码 Token(谨慎使用)
|
||||
|
||||
Args:
|
||||
token: JWT Token 字符串
|
||||
@@ -175,5 +219,5 @@ class JWTService:
|
||||
return None
|
||||
|
||||
|
||||
# 全局实例(生产环境应该从配置读取)
|
||||
jwt_service = JWTService()
|
||||
# 全局实例(生产环境必须从配置读取有效的 secret_key)
|
||||
# jwt_service = JWTService() # 不再允许无参数实例化
|
||||
|
||||
Reference in New Issue
Block a user