fix: resolve CI check failures - format code and fix line length issues
This commit is contained in:
@@ -1,6 +1,4 @@
|
||||
"""
|
||||
请求日志中间件
|
||||
"""
|
||||
"""请求日志和限流中间件"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
@@ -14,39 +12,39 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# 敏感参数名称模式(不区分大小写)
|
||||
SENSITIVE_PARAM_PATTERNS = re.compile(
|
||||
r'^(password|passwd|pwd|token|secret|key|authorization|auth|api_key|apikey|'
|
||||
r'access_token|refresh_token|accesstoken|refreshtoken|'
|
||||
r'session_id|sessionid|sid|cookie|csrf|xsrf|bearer)$',
|
||||
re.IGNORECASE
|
||||
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('?'):
|
||||
|
||||
if query_string.startswith("?"):
|
||||
query_string = query_string[1:]
|
||||
|
||||
|
||||
if not query_string:
|
||||
return query_string
|
||||
|
||||
parts = query_string.split('&')
|
||||
|
||||
parts = query_string.split("&")
|
||||
filtered_parts = []
|
||||
|
||||
|
||||
for part in parts:
|
||||
if '=' in part:
|
||||
key, value = part.split('=', 1)
|
||||
if "=" in part:
|
||||
key, value = part.split("=", 1)
|
||||
if SENSITIVE_PARAM_PATTERNS.match(key):
|
||||
filtered_parts.append(f"{key}=***")
|
||||
else:
|
||||
@@ -54,8 +52,8 @@ def filter_sensitive_params(query_string: Optional[str]) -> Optional[str]:
|
||||
else:
|
||||
# 没有 = 的参数,保留原样
|
||||
filtered_parts.append(part)
|
||||
|
||||
return '&'.join(filtered_parts)
|
||||
|
||||
return "&".join(filtered_parts)
|
||||
|
||||
|
||||
class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
||||
@@ -68,51 +66,57 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
||||
# 过滤 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}")
|
||||
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:
|
||||
@@ -123,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"
|
||||
),
|
||||
}
|
||||
},
|
||||
)
|
||||
@@ -136,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
|
||||
|
||||
Reference in New Issue
Block a user