fix(auth): stabilize audit and token timing

This commit is contained in:
2026-06-11 11:17:11 +08:00
parent f61be3685f
commit 9be0cd34cf
2 changed files with 38 additions and 9 deletions
+10 -6
View File
@@ -1,4 +1,4 @@
from datetime import datetime, timedelta from datetime import datetime, timedelta, timezone
from typing import Optional, Union, Any from typing import Optional, Union, Any
from jose import jwt from jose import jwt
@@ -8,6 +8,10 @@ from app.core.config import settings
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
def _utc_now() -> datetime:
return datetime.now(timezone.utc)
def create_access_token( def create_access_token(
subject: Union[str, Any], expires_delta: Optional[timedelta] = None subject: Union[str, Any], expires_delta: Optional[timedelta] = None
) -> str: ) -> str:
@@ -22,9 +26,9 @@ def create_access_token(
JWT token 字符串 JWT token 字符串
""" """
if expires_delta: if expires_delta:
expire = datetime.now() + expires_delta expire = _utc_now() + expires_delta
else: else:
expire = datetime.now() + timedelta( expire = _utc_now() + timedelta(
minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES
) )
@@ -32,7 +36,7 @@ def create_access_token(
"exp": expire, "exp": expire,
"sub": str(subject), "sub": str(subject),
"type": "access", "type": "access",
"iat": datetime.now(), "iat": _utc_now(),
} }
encoded_jwt = jwt.encode( encoded_jwt = jwt.encode(
to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM
@@ -50,13 +54,13 @@ def create_refresh_token(subject: Union[str, Any]) -> str:
Returns: Returns:
JWT refresh token 字符串 JWT refresh token 字符串
""" """
expire = datetime.now() + timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS) expire = _utc_now() + timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS)
to_encode = { to_encode = {
"exp": expire, "exp": expire,
"sub": str(subject), "sub": str(subject),
"type": "refresh", "type": "refresh",
"iat": datetime.now(), "iat": _utc_now(),
} }
encoded_jwt = jwt.encode( encoded_jwt = jwt.encode(
to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM
+28 -3
View File
@@ -61,11 +61,20 @@ class AuditMiddleware(BaseHTTPMiddleware):
"/api/v1/openproject/", "/api/v1/openproject/",
"/openproject/", "/openproject/",
} }
EXCLUDED_PATH_PREFIXES = (
)
async def dispatch(self, request: Request, call_next: Callable) -> Response: async def dispatch(self, request: Request, call_next: Callable) -> Response:
# 提取开始时间 # 提取开始时间
start_time = time.time() start_time = time.time()
# 流式 Copilot 请求前置排除,避免读取/改写 body 影响 SSE 生命周期
if self._is_excluded_path(request.url.path):
response = await call_next(request)
process_time = time.time() - start_time
response.headers["X-Process-Time"] = str(process_time)
return response
# 1. 预判是否需要读取Body (针对写操作) # 1. 预判是否需要读取Body (针对写操作)
# 注意:我们暂时移除早期的 return,因为需要等待路由匹配后才能检查 Tag # 注意:我们暂时移除早期的 return,因为需要等待路由匹配后才能检查 Tag
should_capture_body = request.method in ["POST", "PUT", "PATCH"] should_capture_body = request.method in ["POST", "PUT", "PATCH"]
@@ -74,13 +83,24 @@ class AuditMiddleware(BaseHTTPMiddleware):
if should_capture_body: if should_capture_body:
try: try:
# 注意:读取 body 后需要重新设置,避免影响后续处理 # 注意:读取 body 后需要重新设置,避免影响后续处理
original_receive = request._receive
body = await request.body() body = await request.body()
if body: if body:
request_data = json.loads(body.decode()) request_data = json.loads(body.decode())
# 重新构造请求以供后续使用 # 重新构造请求以供后续使用:仅回放一次,后续回落原始 receive
body_sent = False
async def receive(): async def receive():
return {"type": "http.request", "body": body} nonlocal body_sent
if not body_sent:
body_sent = True
return {
"type": "http.request",
"body": body,
"more_body": False,
}
return await original_receive()
request._receive = receive request._receive = receive
except Exception as e: except Exception as e:
@@ -90,7 +110,7 @@ class AuditMiddleware(BaseHTTPMiddleware):
response = await call_next(request) response = await call_next(request)
# 3. 决定是否审计 # 3. 决定是否审计
if request.url.path in self.EXCLUDED_PATHS: if self._is_excluded_path(request.url.path):
process_time = time.time() - start_time process_time = time.time() - start_time
response.headers["X-Process-Time"] = str(process_time) response.headers["X-Process-Time"] = str(process_time)
return response return response
@@ -150,6 +170,11 @@ class AuditMiddleware(BaseHTTPMiddleware):
return response return response
def _is_excluded_path(self, path: str) -> bool:
if path in self.EXCLUDED_PATHS:
return True
return any(path.startswith(prefix) for prefix in self.EXCLUDED_PATH_PREFIXES)
def _resolve_project_id(self, request: Request) -> UUID | None: def _resolve_project_id(self, request: Request) -> UUID | None:
project_header = request.headers.get("X-Project-Id") project_header = request.headers.get("X-Project-Id")
if not project_header: if not project_header: