diff --git a/apps/api/main.py b/apps/api/main.py old mode 100644 new mode 100755 index 865c89184..b8dc0d2fd --- a/apps/api/main.py +++ b/apps/api/main.py @@ -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) diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py old mode 100644 new mode 100755 index f6ffe9965..36e441db3 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -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: diff --git a/packages/adapters/redis/session_store.py b/packages/adapters/redis/session_store.py old mode 100644 new mode 100755 index 1378ca8c8..e252a38a3 --- a/packages/adapters/redis/session_store.py +++ b/packages/adapters/redis/session_store.py @@ -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) diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py old mode 100644 new mode 100755 index 409f4fed6..3aed170fc --- a/packages/application/auth/login_use_case.py +++ b/packages/application/auth/login_use_case.py @@ -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)}"