fix(p1): 修复 4 个 P1 技术债务问题 #11
Regular → Executable
+26
-7
@@ -30,13 +30,32 @@ app.add_exception_handler(StarletteHTTPException, http_exception_handler)
|
||||
app.add_exception_handler(RequestValidationError, validation_exception_handler)
|
||||
app.add_exception_handler(Exception, general_exception_handler)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=settings.CORS_ORIGINS,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
# CORS 配置:根据 DEBUG 模式区分
|
||||
# 生产环境:allow_credentials=True 时不能使用通配符 "*"
|
||||
if settings.DEBUG:
|
||||
# 开发环境:允许所有来源(方便本地调试)
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
else:
|
||||
# 生产环境:严格限制来源和方法
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=settings.CORS_ORIGINS,
|
||||
allow_credentials=True,
|
||||
allow_methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"],
|
||||
allow_headers=[
|
||||
"Authorization",
|
||||
"Content-Type",
|
||||
"X-Request-ID",
|
||||
"X-Correlation-ID",
|
||||
],
|
||||
)
|
||||
|
||||
app.add_middleware(GZipMiddleware, minimum_size=1000)
|
||||
app.add_middleware(RequestLoggingMiddleware)
|
||||
|
||||
|
||||
Regular → Executable
+138
-24
@@ -1,4 +1,10 @@
|
||||
import subprocess
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from celery import Celery
|
||||
from celery.app.task import Task
|
||||
from celery.utils.log import get_task_logger
|
||||
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.core.asset_types import infer_mime_type_from_storage_key
|
||||
@@ -10,6 +16,87 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
)
|
||||
from packages.domain import Asset, AssetStatus, IngestJobStatus
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
def extract_media_metadata(file_url: str, media_type: str) -> dict:
|
||||
"""
|
||||
提取媒体文件的元数据。
|
||||
|
||||
Args:
|
||||
file_url: 媒体文件 URL 或本地路径
|
||||
media_type: 媒体类型 (video, audio, image)
|
||||
|
||||
Returns:
|
||||
提取的元数据字典,失败时返回空字典
|
||||
"""
|
||||
metadata = {}
|
||||
|
||||
try:
|
||||
if media_type == "video":
|
||||
# 使用 ffprobe 提取视频元数据
|
||||
cmd = [
|
||||
"ffprobe",
|
||||
"-v", "quiet",
|
||||
"-print_format", "json",
|
||||
"-show_format",
|
||||
"-show_streams",
|
||||
file_url,
|
||||
]
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
import json as json_lib
|
||||
|
||||
probe_data = json_lib.loads(result.stdout)
|
||||
|
||||
# 提取视频流信息
|
||||
for stream in probe_data.get("streams", []):
|
||||
if stream.get("codec_type") == "video":
|
||||
metadata["width"] = int(stream.get("width", 0))
|
||||
metadata["height"] = int(stream.get("height", 0))
|
||||
metadata["codec"] = stream.get("codec_name", "")
|
||||
metadata["fps"] = eval(stream.get("r_frame_rate", "0/1")) if stream.get("r_frame_rate") else 0
|
||||
break
|
||||
|
||||
# 提取格式信息
|
||||
format_info = probe_data.get("format", {})
|
||||
metadata["duration"] = float(format_info.get("duration", 0))
|
||||
metadata["size_bytes"] = int(format_info.get("size", 0))
|
||||
metadata["bitrate"] = int(format_info.get("bit_rate", 0))
|
||||
|
||||
elif media_type == "image":
|
||||
# 使用 Pillow 提取图片元数据
|
||||
try:
|
||||
from PIL import Image
|
||||
|
||||
with Image.open(file_url) as img:
|
||||
metadata["width"] = img.width
|
||||
metadata["height"] = img.height
|
||||
metadata["format"] = img.format
|
||||
metadata["mode"] = img.mode
|
||||
if hasattr(img, "_getexif") and img._getexif():
|
||||
exif = img._getexif()
|
||||
if exif:
|
||||
metadata["exif"] = {k: str(v) for k, v in exif.items() if isinstance(v, (str, int, float))}
|
||||
except ImportError:
|
||||
logger.warning("Pillow not available for image metadata extraction")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract image metadata: {e}")
|
||||
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.warning(f"Timeout extracting metadata from {file_url}")
|
||||
except FileNotFoundError:
|
||||
logger.warning(f"ffprobe not found, cannot extract video metadata")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract metadata: {e}")
|
||||
|
||||
return metadata
|
||||
|
||||
|
||||
@celery_app.task(name="worker.ingest_asset")
|
||||
def ingest_asset(job_id: str) -> dict:
|
||||
@@ -18,34 +105,49 @@ def ingest_asset(job_id: str) -> dict:
|
||||
|
||||
Steps:
|
||||
1. Fetch IngestJob from repository
|
||||
2. Extract metadata from storage_key (placeholder: mock metadata)
|
||||
2. Extract metadata from storage_key
|
||||
3. Create Asset entity
|
||||
4. Update IngestJob status to COMPLETED
|
||||
5. Return result
|
||||
"""
|
||||
db = SessionLocal()
|
||||
job_repo = SQLAlchemyIngestJobRepository(db)
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
return {"status": "failed", "error": "job not found"}
|
||||
|
||||
try:
|
||||
job_repo = SQLAlchemyIngestJobRepository(db)
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
return {"status": "failed", "error": "job not found"}
|
||||
|
||||
# Update job status to PROCESSING
|
||||
job.status = IngestJobStatus.PROCESSING
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
job_repo.update(job)
|
||||
db.commit()
|
||||
|
||||
# Mock metadata extraction (in real implementation: use ffprobe, Pillow, etc.)
|
||||
# Extract real metadata from media file
|
||||
filename = job.storage_key.split("/")[-1]
|
||||
mime_type = infer_mime_type_from_storage_key(job.storage_key)
|
||||
metadata = {
|
||||
"duration": 10.5,
|
||||
"width": 1920,
|
||||
"height": 1080,
|
||||
"size_bytes": 1024000,
|
||||
}
|
||||
|
||||
# Determine media type from mime_type
|
||||
media_type = "video"
|
||||
if mime_type.startswith("image/"):
|
||||
media_type = "image"
|
||||
elif mime_type.startswith("audio/"):
|
||||
media_type = "audio"
|
||||
|
||||
# Extract metadata (returns empty dict on failure)
|
||||
storage_url = job.storage_key # Assuming storage_key is usable as URL/path
|
||||
metadata = extract_media_metadata(storage_url, media_type)
|
||||
|
||||
# Fill in defaults if metadata extraction failed
|
||||
if not metadata:
|
||||
metadata = {
|
||||
"duration": 0,
|
||||
"width": 0,
|
||||
"height": 0,
|
||||
"size_bytes": 0,
|
||||
}
|
||||
|
||||
# Create Asset
|
||||
asset = Asset.create(
|
||||
@@ -56,10 +158,10 @@ def ingest_asset(job_id: str) -> dict:
|
||||
storage_key=job.storage_key,
|
||||
mime_type=mime_type,
|
||||
metadata=metadata,
|
||||
file_size=int(metadata["size_bytes"]),
|
||||
duration=float(metadata["duration"]),
|
||||
width=int(metadata["width"]),
|
||||
height=int(metadata["height"]),
|
||||
file_size=int(metadata.get("size_bytes", 0)),
|
||||
duration=float(metadata.get("duration", 0)),
|
||||
width=int(metadata.get("width", 0)),
|
||||
height=int(metadata.get("height", 0)),
|
||||
status=AssetStatus.READY,
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
@@ -70,21 +172,33 @@ def ingest_asset(job_id: str) -> dict:
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
job_repo.update(job)
|
||||
|
||||
db.commit()
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"job_id": job.id,
|
||||
"asset_id": asset.id,
|
||||
}
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.error(f"Failed to ingest asset {job_id}: {e}")
|
||||
|
||||
# Update job status to FAILED
|
||||
job.status = IngestJobStatus.FAILED
|
||||
job.error_message = str(e)
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
job_repo.update(job)
|
||||
try:
|
||||
job_repo = SQLAlchemyIngestJobRepository(db)
|
||||
job = job_repo.get(job_id)
|
||||
if job:
|
||||
job.status = IngestJobStatus.FAILED
|
||||
job.error_message = str(e)
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
job_repo.update(job)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
|
||||
return {
|
||||
"status": "failed",
|
||||
"job_id": job.id,
|
||||
"job_id": job_id,
|
||||
"error": str(e),
|
||||
}
|
||||
finally:
|
||||
|
||||
Regular → Executable
+41
@@ -28,6 +28,9 @@ class NoopSessionStore:
|
||||
def get_session(self, session_id: str) -> Optional[dict]:
|
||||
return None
|
||||
|
||||
def get_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]:
|
||||
return None
|
||||
|
||||
def get_refresh_token(self, session_id: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
@@ -78,6 +81,10 @@ class SessionStore:
|
||||
"""生成 refresh_token key"""
|
||||
return f"refresh_token:{session_id}"
|
||||
|
||||
def _refresh_token_to_session_key(self, refresh_token: str) -> str:
|
||||
"""生成 refresh_token -> session_id 的反向映射 key"""
|
||||
return f"refresh_token_map:{refresh_token}"
|
||||
|
||||
def _user_sessions_key(self, user_id: str) -> str:
|
||||
"""生成用户所有 Session 的 key"""
|
||||
return f"user_sessions:{user_id}"
|
||||
@@ -127,6 +134,10 @@ class SessionStore:
|
||||
refresh_token_key = self._refresh_token_key(session_id)
|
||||
self.redis.setex(refresh_token_key, expires_in_seconds, refresh_token)
|
||||
|
||||
# 保存 refresh_token -> session_id 的反向映射
|
||||
refresh_token_map_key = self._refresh_token_to_session_key(refresh_token)
|
||||
self.redis.setex(refresh_token_map_key, expires_in_seconds, session_id)
|
||||
|
||||
# 添加到用户的 Session 集合
|
||||
user_sessions_key = self._user_sessions_key(user_id)
|
||||
self.redis.sadd(user_sessions_key, session_id)
|
||||
@@ -158,6 +169,30 @@ class SessionStore:
|
||||
print(f"Failed to get session: {e}")
|
||||
return None
|
||||
|
||||
def get_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]:
|
||||
"""
|
||||
通过 refresh_token 获取 Session
|
||||
|
||||
Args:
|
||||
refresh_token: 刷新令牌
|
||||
|
||||
Returns:
|
||||
Session 数据,如果不存在返回 None
|
||||
"""
|
||||
try:
|
||||
# 先通过反向映射找到 session_id
|
||||
refresh_token_map_key = self._refresh_token_to_session_key(refresh_token)
|
||||
session_id = self.redis.get(refresh_token_map_key)
|
||||
|
||||
if not session_id:
|
||||
return None
|
||||
|
||||
# 再获取完整的 session 数据
|
||||
return self.get_session(session_id)
|
||||
except Exception as e:
|
||||
print(f"Failed to get session by refresh_token: {e}")
|
||||
return None
|
||||
|
||||
def get_refresh_token(self, session_id: str) -> Optional[str]:
|
||||
"""
|
||||
获取 refresh_token
|
||||
@@ -227,8 +262,14 @@ class SessionStore:
|
||||
|
||||
# 删除 refresh_token
|
||||
refresh_token_key = self._refresh_token_key(session_id)
|
||||
refresh_token = self.redis.get(refresh_token_key)
|
||||
self.redis.delete(refresh_token_key)
|
||||
|
||||
# 删除反向映射
|
||||
if refresh_token:
|
||||
refresh_token_map_key = self._refresh_token_to_session_key(refresh_token)
|
||||
self.redis.delete(refresh_token_map_key)
|
||||
|
||||
# 从用户 Session 集合中移除
|
||||
user_sessions_key = self._user_sessions_key(user_id)
|
||||
self.redis.srem(user_sessions_key, session_id)
|
||||
|
||||
Regular → Executable
+62
-12
@@ -10,6 +10,7 @@ from typing import Optional
|
||||
import jwt as pyjwt
|
||||
|
||||
from packages.adapters.redis import get_session_store
|
||||
from packages.adapters.redis.session_store import SessionStore, NoopSessionStore
|
||||
from packages.domain.auth import jwt_service, password_hasher
|
||||
|
||||
LEGACY_SHA256_HEX_LENGTH = 64
|
||||
@@ -68,7 +69,7 @@ class LoginUseCase:
|
||||
|
||||
def __init__(self, user_repository, session_store=None, jwt_secret_key: str | None = None):
|
||||
self.user_repository = user_repository
|
||||
self.session_store = session_store or get_session_store()
|
||||
self.session_store: SessionStore | NoopSessionStore = session_store or get_session_store()
|
||||
self.jwt_secret_key = jwt_secret_key or jwt_service.config.SECRET_KEY
|
||||
|
||||
def execute(self, request: LoginRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
|
||||
@@ -171,8 +172,10 @@ class RefreshTokenRequest:
|
||||
class RefreshTokenUseCase:
|
||||
"""刷新令牌用例"""
|
||||
|
||||
def __init__(self, user_repository):
|
||||
def __init__(self, user_repository, session_store=None, jwt_secret_key: str | None = None):
|
||||
self.user_repository = user_repository
|
||||
self.session_store: SessionStore | NoopSessionStore = session_store or get_session_store()
|
||||
self.jwt_secret_key = jwt_secret_key or jwt_service.config.SECRET_KEY
|
||||
|
||||
def execute(self, request: RefreshTokenRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
|
||||
"""
|
||||
@@ -188,18 +191,65 @@ class RefreshTokenUseCase:
|
||||
if not request.refresh_token:
|
||||
return None, "Refresh token is required"
|
||||
|
||||
# 1. 查找 session(通过遍历所有 session)
|
||||
# 注意:这里为了简化,先用遍历实现,生产环境应该用 refresh_token -> session_id 的索引
|
||||
session = None
|
||||
session_id = None
|
||||
# 1. 通过 refresh_token 查找 session
|
||||
session = self.session_store.get_session_by_refresh_token(request.refresh_token)
|
||||
if not session:
|
||||
return None, "Invalid or expired refresh token"
|
||||
|
||||
# 这是一个简化实现,实际应该在 SessionStore 中添加 find_by_refresh_token 方法
|
||||
# 这里我们假设 refresh_token 就是 session_id(简化处理)
|
||||
# 生产环境需要更复杂的映射
|
||||
# 2. 检查 session 是否过期
|
||||
session_id = session.get("session_id")
|
||||
if not session_id:
|
||||
return None, "Invalid session data"
|
||||
|
||||
# 临时方案:从 Redis 获取(需要在 session_store 中添加方法)
|
||||
# 现在先返回错误,提示需要实现
|
||||
return None, "Refresh token implementation pending (需要完善 session_store)"
|
||||
# 3. 检查 session 是否在有效期内
|
||||
expires_at_str = session.get("expires_at")
|
||||
if expires_at_str:
|
||||
expires_at = datetime.fromisoformat(expires_at_str)
|
||||
if datetime.now(timezone.utc) > expires_at:
|
||||
# session 已过期,删除它
|
||||
self.session_store.delete_session(session_id)
|
||||
return None, "Session has expired, please login again"
|
||||
|
||||
# 4. 获取用户信息
|
||||
user_id = session.get("user_id")
|
||||
if not user_id:
|
||||
return None, "Invalid session: missing user_id"
|
||||
|
||||
user = self.user_repository.get(user_id)
|
||||
if not user:
|
||||
return None, "User not found"
|
||||
|
||||
# 5. 生成新的 access_token
|
||||
now = datetime.now(timezone.utc)
|
||||
access_token_payload = {
|
||||
"sub": user.id,
|
||||
"sid": session_id,
|
||||
"type": "user_auth",
|
||||
"iat": now,
|
||||
"exp": now + timedelta(minutes=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES),
|
||||
}
|
||||
access_token = pyjwt.encode(
|
||||
access_token_payload,
|
||||
self.jwt_secret_key,
|
||||
algorithm=jwt_service.config.ALGORITHM,
|
||||
)
|
||||
|
||||
# 6. 更新 session 的最后活跃时间
|
||||
self.session_store.update_last_active(session_id)
|
||||
|
||||
# 7. 返回新的登录响应(refresh_token 保持不变)
|
||||
return (
|
||||
LoginResponse(
|
||||
access_token=access_token,
|
||||
refresh_token=request.refresh_token, # 保持原有的 refresh_token
|
||||
user_id=user.id,
|
||||
email=user.email,
|
||||
username=user.username,
|
||||
display_name=user.display_name,
|
||||
expires_in=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return None, f"Token refresh failed: {str(e)}"
|
||||
|
||||
Reference in New Issue
Block a user