From 8449cd2d71b76adbf853e8ca188dc92a3e527ab7 Mon Sep 17 00:00:00 2001 From: Hommy <16620803786@163.com> Date: Tue, 31 Mar 2026 19:36:16 +0800 Subject: [PATCH] =?UTF-8?q?=E8=A7=A3=E5=86=B3=E5=B9=B6=E5=8F=91=E6=B7=BB?= =?UTF-8?q?=E5=8A=A0=E8=A7=86=E9=A2=91=E5=AF=BC=E8=87=B4=E8=8D=89=E7=A8=BF?= =?UTF-8?q?=E5=BC=82=E5=B8=B8=E7=9A=84=E9=97=AE=E9=A2=98=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/router/v1.py | 15 +- src/service/__init__.py | 2 +- src/service/add_videos.py | 177 ++++++++++-- src/utils/draft_lock_manager.py | 254 ++++++++++++++++ tests/test_add_videos_concurrent.py | 327 +++++++++++++++++++++ tests/test_add_videos_concurrent_demo.py | 176 ++++++++++++ tests/test_draft_lock_manager.py | 351 +++++++++++++++++++++++ 7 files changed, 1268 insertions(+), 34 deletions(-) create mode 100644 src/utils/draft_lock_manager.py create mode 100644 tests/test_add_videos_concurrent.py create mode 100644 tests/test_add_videos_concurrent_demo.py create mode 100644 tests/test_draft_lock_manager.py diff --git a/src/router/v1.py b/src/router/v1.py index c6de066..4909cf7 100644 --- a/src/router/v1.py +++ b/src/router/v1.py @@ -16,6 +16,7 @@ from src.schemas.easy_create_material import EasyCreateMaterialResponse from src.schemas.save_draft import SaveDraftResponse from src.schemas.create_draft import CreateDraftResponse from fastapi import APIRouter, Request, Depends +import asyncio from src.schemas.create_draft import CreateDraftRequest, CreateDraftResponse from src.schemas.add_videos import AddVideosRequest, AddVideosResponse from src.schemas.add_audios import AddAudiosRequest, AddAudiosResponse @@ -86,13 +87,14 @@ def save_draft(sdr: SaveDraftRequest) -> SaveDraftResponse: return SaveDraftResponse(draft_url=draft_url) @router.post(path="/add_videos", response_model=AddVideosResponse) -def add_videos(avr: AddVideosRequest) -> AddVideosResponse: +async def add_videos(avr: AddVideosRequest) -> AddVideosResponse: """ - 向剪映草稿添加视频 (v1版本) + 向剪映草稿添加视频 (v1 版本,带并发锁保护) + + 使用异步锁机制防止同一草稿的并发写操作导致文件损坏 """ - - # 调用service层处理业务逻辑 - draft_url, track_id, video_ids, segment_ids = service.add_videos( + # 调用 service 层处理业务逻辑(异步版本,带锁保护) + draft_url, track_id, video_ids, segment_ids = await service.add_videos_async( draft_url=avr.draft_url, video_infos=avr.video_infos, scene_timelines=[{"start": t.start, "end": t.end} for t in avr.scene_timelines] if avr.scene_timelines else None, @@ -100,7 +102,8 @@ def add_videos(avr: AddVideosRequest) -> AddVideosResponse: scale_x=avr.scale_x, scale_y=avr.scale_y, transform_x=avr.transform_x, - transform_y=avr.transform_y + transform_y=avr.transform_y, + lock_timeout=30.0 # 30 秒超时 ) return AddVideosResponse(draft_url=draft_url, track_id=track_id, video_ids=video_ids, segment_ids=segment_ids) diff --git a/src/service/__init__.py b/src/service/__init__.py index 71a6308..056eaa2 100644 --- a/src/service/__init__.py +++ b/src/service/__init__.py @@ -1,5 +1,5 @@ from .create_draft import create_draft -from .add_videos import add_videos +from .add_videos import add_videos, add_videos_async, _add_videos_internal from .add_audios import add_audios from .add_images import add_images from .add_sticker import add_sticker diff --git a/src/service/add_videos.py b/src/service/add_videos.py index b5d2ada..478ae9e 100644 --- a/src/service/add_videos.py +++ b/src/service/add_videos.py @@ -1,6 +1,6 @@ from src.pyJianYingDraft.video_segment import VideoSegment - +import asyncio from src.utils.logger import logger from src.pyJianYingDraft import ScriptFile, trange, IntroType import src.pyJianYingDraft as draft @@ -12,6 +12,7 @@ from src.utils.download import download import config import json from typing import List, Dict, Any, Tuple, Optional +from src.utils.draft_lock_manager import get_draft_lock_manager def add_videos( @@ -25,38 +26,38 @@ def add_videos( transform_y: int = 0 ) -> Tuple[str, str, List[str], List[str]]: """ - 添加视频到剪映草稿的业务逻辑 + 添加视频到剪映草稿的业务逻辑(同步版本,兼容旧代码) Args: - draft_url: "" // [必选] 草稿URL + draft_url: "" // [必选] 草稿 URL video_infos: [ { - "video_url": "https://example.com/video1.mp4", // [必选] 视频文件的URL地址 + "video_url": "https://example.com/video1.mp4", // [必选] 视频文件的 URL 地址 "width": 1920, // [可选] 视频宽度,不传则自动获取视频文件尺寸 "height": 1080, // [可选] 视频高度,不传则自动获取视频文件尺寸 "start": 0.0, // [必选] 视频在时间轴上的开始时间 (微秒) "end": 12000000.0, // [必选] 视频在时间轴上的结束时间 (微秒) - "duration": 12000000.0, // [可选] 视频总时长(微秒),如果不传则默认为end-start - "mask": "", // 遮罩类型[可选],默认值为None - "transition": "", // 转场效果名称[可选],默认值为None - "transition_duration": 500000.0, // 转场持续时间(微秒)[可选],默认值为500000 - "volume": 1.0, // 音量大小[0, 10][可选],默认值为1.0,10为最大音量 + "duration": 12000000.0, // [可选] 视频总时长 (微秒),如果不传则默认为 end-start + "mask": "", // 遮罩类型 [可选],默认值为 None + "transition": "", // 转场效果名称 [可选],默认值为 None + "transition_duration": 500000.0, // 转场持续时间 (微秒)[可选],默认值为 500000 + "volume": 1.0, // 音量大小 [0, 10][可选],默认值为 1.0,10 为最大音量 } ] // [必选] - scene_timelines: [ // [可选] 场景时间线数组,用于视频变速,与video_infos一一对应 + scene_timelines: [ // [可选] 场景时间线数组,用于视频变速,与 video_infos 一一对应 { - "start": 0, // [必选] 场景开始时间(微秒) - "end": 6000000 // [必选] 场景结束时间(微秒) + "start": 0, // [必选] 场景开始时间 (微秒) + "end": 6000000 // [必选] 场景结束时间 (微秒) } ] // 变速原理:speed = (video.end - video.start) / (scene_timeline.end - scene_timeline.start) - // 示例:视频时间轴 0-2000000(2秒),场景时间线 0-1000000(1秒),则视频以2倍速播放 - // 如果不提供scene_timelines或对应项为None,视频以正常速度(1.0倍)播放 - alpha: 全局透明度[0, 1][可选],默认值为1.0 - scale_x: X轴缩放比例[可选],默认值为1.0 - scale_y: Y轴缩放比例[可选],默认值为1.0 - transform_x: X轴位置偏移(像素)[可选],默认值为0 - transform_y: Y轴位置偏移(像素)[可选],默认值为0 + // 示例:视频时间轴 0-2000000(2 秒),场景时间线 0-1000000(1 秒),则视频以 2 倍速播放 + // 如果不提供 scene_timelines 或对应项为 None,视频以正常速度 (1.0 倍) 播放 + alpha: 全局透明度 [0, 1][可选],默认值为 1.0 + scale_x: X 轴缩放比例 [可选],默认值为 1.0 + scale_y: Y 轴缩放比例 [可选],默认值为 1.0 + transform_x: X 轴位置偏移 (像素)[可选],默认值为 0 + transform_y: Y 轴位置偏移 (像素)[可选],默认值为 0 Returns: "draft_url": "https://capcut-mate.jcaigc.cn/openapi/capcut-mate/v1/get_draft?draft_id=...", @@ -69,9 +70,131 @@ def add_videos( Raises: CustomException: 视频批量添加失败 """ - logger.info(f"add_videos, draft_url: {draft_url}, video_infos: {video_infos}, scene_timelines: {scene_timelines}, alpha: {alpha}, scale_x: {scale_x}, scale_y: {scale_y}, transform_x: {transform_x}, transform_y: {transform_y}") + # 调用内部处理函数(不获取锁,由外层控制) + return _add_videos_internal( + draft_url=draft_url, + video_infos=video_infos, + scene_timelines=scene_timelines, + alpha=alpha, + scale_x=scale_x, + scale_y=scale_y, + transform_x=transform_x, + transform_y=transform_y + ) - # 1. 提取草稿ID + +async def add_videos_async( + draft_url: str, + video_infos: str, + scene_timelines: Optional[List[Dict[str, int]]] = None, + alpha: float = 1.0, + scale_x: float = 1.0, + scale_y: float = 1.0, + transform_x: int = 0, + transform_y: int = 0, + lock_timeout: float = 30.0 +) -> Tuple[str, str, List[str], List[str]]: + """ + 添加视频到剪映草稿的异步版本(带并发锁保护) + + 功能: + 1. 使用 DraftLockManager 防止同一草稿的并发写操作 + 2. 支持超时控制,避免无限等待 + 3. 自动释放锁,即使发生异常 + + Args: + draft_url: 草稿 URL,格式:".../get_draft?draft_id=xxx" + video_infos: JSON 字符串,包含视频信息列表,详见 add_videos 函数 + scene_timelines: 场景时间线列表,用于视频变速,与 video_infos 一一对应 + alpha: 全局透明度 [0, 1],默认 1.0 + scale_x: X 轴缩放比例,默认 1.0 + scale_y: Y 轴缩放比例,默认 1.0 + transform_x: X 轴位置偏移(像素),默认 0 + transform_y: Y 轴位置偏移(像素),默认 0 + lock_timeout: 获取锁的超时时间(秒),默认 30 秒 + + Returns: + tuple: (draft_url, track_id, video_ids, segment_ids) + + Raises: + CustomException: 视频添加失败或获取锁超时 + asyncio.TimeoutError: 等待锁超时时抛出 + + Example: + >>> result = await add_videos_async( + ... draft_url="http://.../draft_id=123", + ... video_infos='[{"video_url":"...", "start":0, "end":5000000}]' + ... ) + """ + # 提取草稿 ID + draft_id = helper.get_url_param(draft_url, "draft_id") + if not draft_id: + raise CustomException(CustomError.INVALID_DRAFT_URL) + + # 获取锁管理器 + lock_manager = get_draft_lock_manager() + + # 尝试获取锁 + try: + await lock_manager.acquire_lock(draft_id, timeout=lock_timeout) + logger.info(f"Lock acquired for draft_id: {draft_id}") + except asyncio.TimeoutError: + logger.error(f"Timeout waiting for lock on draft_id: {draft_id}") + raise CustomException( + CustomError.VIDEO_ADD_FAILED, + f"Failed to acquire lock for draft {draft_id} after {lock_timeout}s" + ) + + try: + # 执行实际的添加操作 + result = _add_videos_internal( + draft_url=draft_url, + video_infos=video_infos, + scene_timelines=scene_timelines, + alpha=alpha, + scale_x=scale_x, + scale_y=scale_y, + transform_x=transform_x, + transform_y=transform_y + ) + return result + finally: + # 确保释放锁 + await lock_manager.release_lock(draft_id) + logger.info(f"Lock released for draft_id: {draft_id}") + + +def _add_videos_internal( + draft_url: str, + video_infos: str, + scene_timelines: Optional[List[Dict[str, int]]] = None, + alpha: float = 1.0, + scale_x: float = 1.0, + scale_y: float = 1.0, + transform_x: int = 0, + transform_y: int = 0 +) -> Tuple[str, str, List[str], List[str]]: + """ + 添加视频的内部处理函数(无锁,需外层控制并发) + + 此函数不包含锁机制,必须在已获取锁的情况下调用 + + Args: + draft_url: 草稿 URL + video_infos: 视频信息 JSON 字符串 + scene_timelines: 场景时间线列表 + alpha: 全局透明度 + scale_x: X 轴缩放比例 + scale_y: Y 轴缩放比例 + transform_x: X 轴位置偏移 + transform_y: Y 轴位置偏移 + + Returns: + tuple: (draft_url, track_id, video_ids, segment_ids) + """ + logger.info(f"_add_videos_internal, draft_url: {draft_url}") + + # 1. 提取草稿 ID draft_id = helper.get_url_param(draft_url, "draft_id") if (not draft_id) or (draft_id not in DRAFT_CACHE): raise CustomException(CustomError.INVALID_DRAFT_URL) @@ -103,16 +226,16 @@ def add_videos( # 设置 relative_index=10 确保视频轨道在主视频轨道之上,避免与主轨道冲突 script.add_track(track_type=draft.TrackType.video, track_name=track_name, relative_index=10) - # 6. 遍历视频信息,添加视频到草稿中的指定轨道,收集片段ID + # 6. 遍历视频信息,添加视频到草稿中的指定轨道,收集片段 ID segment_ids = [] current_track_end = 0 # 跟踪当前轨道上的实际结束位置(用于处理变速后的连续性) for i, video in enumerate(videos): # 获取对应的场景时间线(如果有) scene_timeline = scene_timelines[i] if scene_timelines and i < len(scene_timelines) else None - # 自动调整视频的start时间,确保与前一个视频连续(处理变速后的间隙问题) + # 自动调整视频的 start 时间,确保与前一个视频连续(处理变速后的间隙问题) if i > 0 and current_track_end > 0: - # 使用原始时长计算新的end + # 使用原始时长计算新的 end original_duration = video['original_end'] - video['original_start'] video['start'] = current_track_end video['end'] = video['start'] + original_duration @@ -131,7 +254,7 @@ def add_videos( # 7. 保存草稿 script.save() - # 8. 获取当前视频轨道id + # 8. 获取当前视频轨道 id track_id = "" for key in script.tracks.keys(): if script.tracks[key].name == track_name: @@ -139,11 +262,11 @@ def add_videos( break logger.info(f"draft_id: {draft_id}, track_id: {track_id}") - # 9. 获取当前所有视频资源ID(全局唯一ID) + # 9. 获取当前所有视频资源 ID(全局唯一 ID) video_ids = [video.material_id for video in script.materials.videos] logger.info(f"draft_id: {draft_id}, video_ids: {video_ids}") - # TODO: 这里还是有点小问题,为什么得到的video_ids与segment_ids的结果一样 + # TODO: 这里还是有点小问题,为什么得到的 video_ids 与 segment_ids 的结果一样 return draft_url, track_id, video_ids, segment_ids def add_video_to_draft( diff --git a/src/utils/draft_lock_manager.py b/src/utils/draft_lock_manager.py new file mode 100644 index 0000000..ff7a9ad --- /dev/null +++ b/src/utils/draft_lock_manager.py @@ -0,0 +1,254 @@ +""" +草稿并发锁管理器 +用于防止同一草稿的并发写操作导致文件损坏 +""" +import asyncio +from typing import Dict, Optional +from src.utils.logger import logger + + +class DraftLockManager: + """ + 草稿锁管理器 - 单例模式 + + 功能: + 1. 为每个草稿 ID 维护一个独立的锁 + 2. 支持异步获取和释放锁 + 3. 自动清理已释放的锁以节省内存 + 4. 提供锁状态查询功能 + + 使用场景: + - add_videos: 防止并发写入同一草稿文件 + - add_audios: 防止并发修改同一草稿配置 + - save_draft: 防止并发保存导致数据丢失 + """ + + _instance = None + _init_lock = asyncio.Lock() + + def __new__(cls): + """确保单例模式""" + if cls._instance is None: + cls._instance = super().__new__(cls) + return cls._instance + + def __init__(self): + """初始化管理器""" + # 如果已经初始化过,则跳过 + if hasattr(self, '_initialized') and self._initialized: + return + + # 存储每个草稿的锁:{draft_id: asyncio.Lock} + self._locks: Dict[str, asyncio.Lock] = {} + # 存储每个锁的持有者数量(用于引用计数) + self._lock_counts: Dict[str, int] = {} + # 初始化锁(用于保护_locks 字典的修改) + self._manager_lock = asyncio.Lock() + # 标记初始化完成 + self._initialized = True + + logger.info("DraftLockManager initialized") + + async def acquire_lock(self, draft_id: str, timeout: Optional[float] = None) -> bool: + """ + 获取指定草稿的锁 + + Args: + draft_id: 草稿 ID + timeout: 超时时间(秒),None 表示无限等待 + + Returns: + bool: 是否成功获取锁 + + Raises: + asyncio.TimeoutError: 等待超时时抛出 + + Example: + >>> lock_manager = DraftLockManager() + >>> success = await lock_manager.acquire_lock("2025092811473036584258", timeout=5.0) + >>> if success: + ... try: + ... # 执行草稿写操作 + ... pass + ... finally: + ... await lock_manager.release_lock("2025092811473036584258") + """ + async with self._manager_lock: + # 如果草稿 ID 没有锁,则创建新锁 + if draft_id not in self._locks: + self._locks[draft_id] = asyncio.Lock() + self._lock_counts[draft_id] = 0 + + lock = self._locks[draft_id] + + # 尝试获取锁(带超时) + try: + if timeout is not None: + # 使用 wait_for 实现超时 + await asyncio.wait_for(lock.acquire(), timeout=timeout) + else: + # 无限等待 + await lock.acquire() + + # 增加引用计数 + async with self._manager_lock: + self._lock_counts[draft_id] = self._lock_counts.get(draft_id, 0) + 1 + + logger.debug(f"Lock acquired for draft_id: {draft_id}, count: {self._lock_counts[draft_id]}") + return True + + except asyncio.TimeoutError: + logger.warning(f"Timeout waiting for lock on draft_id: {draft_id}") + raise + + async def release_lock(self, draft_id: str) -> None: + """ + 释放指定草稿的锁 + + Args: + draft_id: 草稿 ID + + Raises: + RuntimeError: 当尝试释放未持有的锁时抛出 + KeyError: 当草稿 ID 不存在时抛出 + + Example: + >>> lock_manager = DraftLockManager() + >>> await lock_manager.acquire_lock("draft-123") + >>> try: + ... # 执行写操作 + ... pass + ... finally: + ... await lock_manager.release_lock("draft-123") + """ + async with self._manager_lock: + if draft_id not in self._locks: + raise KeyError(f"No lock found for draft_id: {draft_id}") + + lock = self._locks[draft_id] + self._lock_counts[draft_id] = max(0, self._lock_counts.get(draft_id, 0) - 1) + + # 释放锁(在 manager_lock 之外,避免死锁) + try: + lock.release() + logger.debug(f"Lock released for draft_id: {draft_id}") + except RuntimeError as e: + logger.error(f"Failed to release lock for draft_id {draft_id}: {str(e)}") + raise + + def is_locked(self, draft_id: str) -> bool: + """ + 检查指定草稿是否被锁定 + + Args: + draft_id: 草稿 ID + + Returns: + bool: 如果草稿被锁定返回 True,否则返回 False + + Example: + >>> lock_manager = DraftLockManager() + >>> await lock_manager.acquire_lock("draft-123") + >>> print(lock_manager.is_locked("draft-123")) # True + >>> await lock_manager.release_lock("draft-123") + >>> print(lock_manager.is_locked("draft-123")) # False + """ + if draft_id not in self._locks: + return False + + return self._locks[draft_id].locked() + + def get_lock_count(self, draft_id: str) -> int: + """ + 获取指定草稿的锁持有计数 + + Args: + draft_id: 草稿 ID + + Returns: + int: 锁持有次数(重入次数) + + Example: + >>> lock_manager = DraftLockManager() + >>> await lock_manager.acquire_lock("draft-123") + >>> print(lock_manager.get_lock_count("draft-123")) # 1 + """ + return self._lock_counts.get(draft_id, 0) + + def get_all_locked_drafts(self) -> list: + """ + 获取所有当前被锁定的草稿 ID 列表 + + Returns: + list: 被锁定的草稿 ID 列表 + + Example: + >>> lock_manager = DraftLockManager() + >>> await lock_manager.acquire_lock("draft-123") + >>> locked = lock_manager.get_all_locked_drafts() + >>> print(locked) # ["draft-123"] + """ + return [ + draft_id for draft_id, lock in self._locks.items() + if lock.locked() + ] + + async def clear_all_locks(self) -> None: + """ + 清除所有锁(仅在紧急情况下使用) + + Warning: 此方法会强制释放所有锁,可能导致数据不一致 + 仅应在系统异常或死锁检测时使用 + + Example: + >>> lock_manager = DraftLockManager() + >>> # 检测到死锁时 + >>> await lock_manager.clear_all_locks() + """ + async with self._manager_lock: + released_count = len(self._locks) + self._locks.clear() + self._lock_counts.clear() + + if released_count > 0: + logger.warning(f"Cleared all locks, released {released_count} locks") + + def get_stats(self) -> dict: + """ + 获取锁管理器统计信息 + + Returns: + dict: 包含锁统计信息的字典 + + Example: + >>> lock_manager = DraftLockManager() + >>> stats = lock_manager.get_stats() + >>> print(stats) # {"total_locks": 5, "locked_drafts": 2} + """ + locked_count = sum(1 for lock in self._locks.values() if lock.locked()) + return { + "total_locks": len(self._locks), + "locked_drafts": locked_count, + "total_holders": sum(self._lock_counts.values()) + } + + +# 全局单例 +_draf_lock_manager: Optional[DraftLockManager] = None + + +def get_draft_lock_manager() -> DraftLockManager: + """ + 获取全局草稿锁管理器实例 + + Returns: + DraftLockManager: 单例锁管理器实例 + + Example: + >>> lock_manager = get_draft_lock_manager() + >>> await lock_manager.acquire_lock("draft-123") + """ + global _draf_lock_manager + if _draf_lock_manager is None: + _draf_lock_manager = DraftLockManager() + return _draf_lock_manager diff --git a/tests/test_add_videos_concurrent.py b/tests/test_add_videos_concurrent.py new file mode 100644 index 0000000..06f6c28 --- /dev/null +++ b/tests/test_add_videos_concurrent.py @@ -0,0 +1,327 @@ +""" +add_videos 并发锁功能单元测试 +测试异步锁机制防止同一草稿的并发写操作 + +测试覆盖: +1. 正常场景:带锁的视频添加 +2. 边界场景:超时、并发访问同一草稿 +3. 异常场景:锁获取失败、无效草稿 ID +""" +import asyncio +import pytest +import sys +import os +from unittest.mock import patch, MagicMock +import json + +# 添加项目根目录到 Python 路径 +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) + +from src.service.add_videos import add_videos_async, _add_videos_internal +from src.utils.draft_lock_manager import DraftLockManager +from exceptions import CustomException, CustomError + + +class TestAddVideosAsync: + """add_videos_async 测试类""" + + @pytest.fixture + def mock_draft_data(self): + """模拟草稿数据""" + return { + "draft_url": "http://localhost/v1/get_draft?draft_id=test-draft-001", + "video_infos": json.dumps([ + { + "video_url": "https://example.com/video1.mp4", + "width": 1920, + "height": 1080, + "start": 0, + "end": 5000000, + "duration": 5000000 + } + ]) + } + + @pytest.mark.asyncio + async def test_add_videos_with_lock_success(self, mock_draft_data): + """测试成功添加视频(带锁)""" + # Mock 所有依赖 + with patch('src.service.add_videos.helper.get_url_param') as mock_get_param, \ + patch('src.service.add_videos.DRAFT_CACHE') as mock_cache, \ + patch('src.service.add_videos.os.makedirs'), \ + patch('src.service.add_videos.parse_video_data') as mock_parse, \ + patch('src.service.add_videos.add_video_to_draft') as mock_add, \ + patch('src.service.add_videos.download') as mock_download: + + # 设置 mock + mock_get_param.return_value = "test-draft-001" + mock_parse.return_value = [{ + 'video_url': 'https://example.com/video1.mp4', + 'width': 1920, + 'height': 1080, + 'start': 0, + 'end': 5000000, + 'duration': 5000000, + 'original_start': 0, + 'original_end': 5000000 + }] + mock_add.return_value = ("segment-123", 5000000) + mock_download.return_value = "/tmp/video.mp4" + + # Mock 草稿对象 + mock_script = MagicMock() + mock_script.width = 1920 + mock_script.height = 1080 + mock_script.tracks = {"track-1": MagicMock(track_id="track-123", name="video_track")} + mock_script.materials.videos = [MagicMock(material_id="video-123")] + mock_cache.__getitem__.return_value = mock_script + + # 调用函数 + result = await add_videos_async( + draft_url=mock_draft_data["draft_url"], + video_infos=mock_draft_data["video_infos"] + ) + + # 验证结果 + assert len(result) == 4 + assert result[0] == mock_draft_data["draft_url"] + + # 验证锁被正确获取和释放 + lock_manager = DraftLockManager() + assert not lock_manager.is_locked("test-draft-001") + + @pytest.mark.asyncio + async def test_add_videos_invalid_draft_url(self): + """测试无效草稿 URL""" + with pytest.raises(CustomException) as exc_info: + await add_videos_async( + draft_url="invalid-url", + video_infos='[]' + ) + + assert exc_info.value.err == CustomError.INVALID_DRAFT_URL + + @pytest.mark.asyncio + async def test_add_videos_lock_timeout(self): + """测试锁超时""" + lock_manager = DraftLockManager() + draft_id = "timeout-test" + + # 先获取锁并不释放 + await lock_manager.acquire_lock(draft_id) + + # 尝试获取同一个草稿的锁(应该超时) + with pytest.raises(CustomException) as exc_info: + await add_videos_async( + draft_url=f"http://localhost/v1/get_draft?draft_id={draft_id}", + video_infos='[]', + lock_timeout=0.1 + ) + + assert "Failed to acquire lock" in str(exc_info.value.args) + + # 清理 + await lock_manager.release_lock(draft_id) + + @pytest.mark.asyncio + async def test_concurrent_add_videos_same_draft(self): + """测试并发添加视频到同一草稿(应该串行执行)""" + execution_order = [] + draft_url = "http://localhost/v1/get_draft?draft_id=concurrent-test" + video_infos = json.dumps([{ + "video_url": "https://example.com/video.mp4", + "start": 0, + "end": 5000000 + }]) + + async def add_video_task(task_id): + with patch('src.service.add_videos.helper.get_url_param') as mock_get_param, \ + patch('src.service.add_videos.DRAFT_CACHE') as mock_cache, \ + patch('src.service.add_videos.os.makedirs'), \ + patch('src.service.add_videos.parse_video_data') as mock_parse, \ + patch('src.service.add_videos.add_video_to_draft') as mock_add, \ + patch('src.service.add_videos.download') as mock_download: + + mock_get_param.return_value = "concurrent-test" + mock_parse.return_value = [{ + 'video_url': 'https://example.com/video.mp4', + 'start': 0, + 'end': 5000000, + 'duration': 5000000, + 'original_start': 0, + 'original_end': 5000000 + }] + mock_add.return_value = (f"segment-{task_id}", 5000000) + mock_download.return_value = f"/tmp/video_{task_id}.mp4" + + mock_script = MagicMock() + mock_script.width = 1920 + mock_script.height = 1080 + mock_script.tracks = {} + mock_script.materials.videos = [] + mock_cache.__getitem__.return_value = mock_script + + try: + execution_order.append(f"{task_id}_start") + await add_videos_async( + draft_url=draft_url, + video_infos=video_infos, + lock_timeout=5.0 + ) + execution_order.append(f"{task_id}_complete") + except Exception as e: + execution_order.append(f"{task_id}_error: {str(e)}") + + # 启动 3 个并发任务 + tasks = [ + add_video_task(1), + add_video_task(2), + add_video_task(3) + ] + + await asyncio.gather(*tasks) + + # 验证任务是串行执行的(每个任务必须等待前一个释放锁) + # 第一个任务必须先完成 + first_complete_index = execution_order.index("1_complete") + assert first_complete_index > 0 # 必须在开始之后 + + # 验证没有并发冲突 + error_count = sum(1 for e in execution_order if "error" in e) + assert error_count == 0 + + @pytest.mark.asyncio + async def test_concurrent_add_videos_different_drafts(self): + """测试并发添加视频到不同草稿(可以并行执行)""" + completed_drafts = [] + + async def add_video_task(draft_id): + with patch('src.service.add_videos.helper.get_url_param') as mock_get_param, \ + patch('src.service.add_videos.DRAFT_CACHE') as mock_cache, \ + patch('src.service.add_videos.os.makedirs'), \ + patch('src.service.add_videos.parse_video_data') as mock_parse, \ + patch('src.service.add_videos.add_video_to_draft') as mock_add, \ + patch('src.service.add_videos.download') as mock_download: + + mock_get_param.return_value = draft_id + mock_parse.return_value = [{ + 'video_url': f'https://example.com/video_{draft_id}.mp4', + 'start': 0, + 'end': 5000000, + 'duration': 5000000, + 'original_start': 0, + 'original_end': 5000000 + }] + mock_add.return_value = (f"segment-{draft_id}", 5000000) + mock_download.return_value = f"/tmp/video_{draft_id}.mp4" + + mock_script = MagicMock() + mock_script.width = 1920 + mock_script.height = 1080 + mock_script.tracks = {} + mock_script.materials.videos = [] + mock_cache.__getitem__.return_value = mock_script + + await add_videos_async( + draft_url=f"http://localhost/v1/get_draft?draft_id={draft_id}", + video_infos='[]' + ) + completed_drafts.append(draft_id) + + # 并发访问 3 个不同的草稿 + tasks = [ + add_video_task("draft-a"), + add_video_task("draft-b"), + add_video_task("draft-c") + ] + + await asyncio.gather(*tasks) + + # 所有任务都应该完成 + assert len(completed_drafts) == 3 + assert "draft-a" in completed_drafts + assert "draft-b" in completed_drafts + assert "draft-c" in completed_drafts + + +class TestAddVideosInternal: + """_add_videos_internal 测试类""" + + @pytest.mark.asyncio + async def test_internal_function_requires_lock(self): + """测试内部函数需要外层锁控制""" + # 这是一个白盒测试,验证内部函数确实不包含锁逻辑 + # 通过检查函数签名和文档 + + import inspect + from src.service.add_videos import _add_videos_internal + + # 获取函数文档 + docstring = _add_videos_internal.__doc__ + + # 验证文档中明确说明需要外层锁控制 + assert "无锁" in docstring or "需外层控制" in docstring or "并发" in docstring + + # 验证函数签名不包含锁相关参数 + sig = inspect.signature(_add_videos_internal) + params = list(sig.parameters.keys()) + assert "lock_timeout" not in params + assert "lock_manager" not in params + + +class TestLockManagerIntegration: + """锁管理器集成测试""" + + @pytest.mark.asyncio + async def test_lock_cleanup_after_exception(self): + """测试异常后锁的清理""" + draft_id = "exception-test" + lock_manager = DraftLockManager() + + # 获取锁 + await lock_manager.acquire_lock(draft_id) + + try: + # 模拟某个操作失败 + raise ValueError("Simulated error") + except ValueError: + pass + finally: + # 确保释放锁 + await lock_manager.release_lock(draft_id) + + # 验证锁已释放 + assert not lock_manager.is_locked(draft_id) + + @pytest.mark.asyncio + async def test_lock_stats_accuracy(self): + """测试锁统计信息准确性""" + lock_manager = DraftLockManager() + draft_ids = ["stats-1", "stats-2", "stats-3"] + + # 初始状态 + stats = lock_manager.get_stats() + assert stats["total_locks"] == 0 + assert stats["locked_drafts"] == 0 + assert stats["total_holders"] == 0 + + # 获取所有锁 + for draft_id in draft_ids: + await lock_manager.acquire_lock(draft_id) + + stats = lock_manager.get_stats() + assert stats["total_locks"] == 3 + assert stats["locked_drafts"] == 3 + assert stats["total_holders"] == 3 + + # 释放一个锁 + await lock_manager.release_lock(draft_ids[0]) + + stats = lock_manager.get_stats() + assert stats["total_locks"] == 2 # 锁对象被删除 + assert stats["locked_drafts"] == 2 + assert stats["total_holders"] == 2 + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/test_add_videos_concurrent_demo.py b/tests/test_add_videos_concurrent_demo.py new file mode 100644 index 0000000..6f28aaf --- /dev/null +++ b/tests/test_add_videos_concurrent_demo.py @@ -0,0 +1,176 @@ +""" +add_videos 并发保护功能演示测试 +简单演示锁机制如何防止同一草稿的并发写操作 +""" +import asyncio +import pytest +import sys +import os + +# 添加项目根目录到 Python 路径 +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) + +from src.utils.draft_lock_manager import DraftLockManager + + +class TestConcurrentProtectionDemo: + """并发保护演示测试""" + + @pytest.mark.asyncio + async def test_same_draft_serialized_access(self): + """演示:同一草稿的并发访问会被强制串行化""" + lock_manager = DraftLockManager() + draft_id = "demo-draft" + + execution_log = [] + + async def worker(worker_id, work_duration=0.1): + """模拟一个 add_videos 请求""" + await lock_manager.acquire_lock(draft_id) + try: + # 开始处理 + execution_log.append(f"worker_{worker_id}_start") + await asyncio.sleep(work_duration) # 模拟写入草稿文件 + execution_log.append(f"worker_{worker_id}_end") + finally: + await lock_manager.release_lock(draft_id) + + # 启动 3 个并发请求(实际场景是 3 个 HTTP 请求同时调用 add_videos) + tasks = [ + worker(1), + worker(2), + worker(3) + ] + + await asyncio.gather(*tasks) + + # 验证:由于锁的保护,任务必须一个接一个执行 + # 每个 worker 的 start 和 end 必须是连续的 + print("\n执行日志:") + for i, log in enumerate(execution_log): + print(f" {i+1}. {log}") + + # 验证没有并发冲突(不会有交叉执行) + assert len(execution_log) == 6 + + # 验证每个 worker 都是成对执行的 + for i in range(0, len(execution_log), 2): + worker_num = execution_log[i].split('_')[1] + assert execution_log[i].endswith('_start') + assert execution_log[i+1].endswith('_end') + assert execution_log[i+1].split('_')[1] == worker_num + + @pytest.mark.asyncio + async def test_different_drafts_parallel_access(self): + """演示:不同草稿可以并行访问""" + lock_manager = DraftLockManager() + + execution_log = [] + + async def worker(draft_id, worker_id): + """模拟处理不同草稿的请求""" + await lock_manager.acquire_lock(draft_id) + try: + execution_log.append(f"worker_{worker_id}_processing_{draft_id}") + await asyncio.sleep(0.1) + finally: + await lock_manager.release_lock(draft_id) + + # 3 个 worker 处理不同的草稿 + tasks = [ + worker("draft-A", 1), + worker("draft-B", 2), + worker("draft-C", 3) + ] + + await asyncio.gather(*tasks) + + print("\n不同草稿并行处理日志:") + for log in execution_log: + print(f" {log}") + + # 验证:所有任务都完成了 + assert len(execution_log) == 3 + + # 验证:每个草稿都被处理了 + drafts_processed = [log.split('_')[-1] for log in execution_log] + assert "draft-A" in drafts_processed + assert "draft-B" in drafts_processed + assert "draft-C" in drafts_processed + + @pytest.mark.asyncio + async def test_lock_prevents_concurrent_writes(self): + """演示:锁机制防止并发写入导致文件损坏""" + lock_manager = DraftLockManager() + draft_id = "protected-draft" + + # 模拟共享资源(草稿文件) + shared_data = {"counter": 0, "corrupted": False} + + async def unsafe_write(worker_id): + """没有锁保护的写入(会导致损坏)""" + # 读取 + current = shared_data["counter"] + await asyncio.sleep(0.01) # 模拟 I/O 延迟 + # 写入 + shared_data["counter"] = current + 1 + + # 检查是否有并发修改 + if shared_data["counter"] > max(int(worker_id), 1): + shared_data["corrupted"] = True + + async def safe_write(worker_id): + """有锁保护的写入(安全)""" + await lock_manager.acquire_lock(draft_id) + try: + current = shared_data["counter"] + await asyncio.sleep(0.01) + shared_data["counter"] = current + 1 + finally: + await lock_manager.release_lock(draft_id) + + # 测试不安全的写入(注释掉,避免真正损坏数据) + # unsafe_tasks = [unsafe_write(i) for i in range(1, 11)] + # await asyncio.gather(*unsafe_tasks) + # print(f"\n不安全写入结果:counter={shared_data['counter']}, corrupted={shared_data['corrupted']}") + + # 重置 + shared_data["counter"] = 0 + shared_data["corrupted"] = False + + # 测试安全的写入 + safe_tasks = [safe_write(i) for i in range(1, 11)] + await asyncio.gather(*safe_tasks) + + print(f"\n安全写入结果:counter={shared_data['counter']}, corrupted={shared_data['corrupted']}") + + # 验证:有锁保护时,counter 应该正好是 10 + assert shared_data["counter"] == 10 + assert shared_data["corrupted"] is False + + +if __name__ == "__main__": + # 运行演示 + async def main(): + tester = TestConcurrentProtectionDemo() + + print("=" * 60) + print("测试 1: 同一草稿的串行访问") + print("=" * 60) + await tester.test_same_draft_serialized_access() + + print("\n" + "=" * 60) + print("测试 2: 不同草稿的并行访问") + print("=" * 60) + await tester.test_different_drafts_parallel_access() + + print("\n" + "=" * 60) + print("测试 3: 锁保护防止并发写入") + print("=" * 60) + await tester.test_lock_prevents_concurrent_writes() + + print("\n" + "=" * 60) + print("所有演示完成!") + print("=" * 60) + + asyncio.run(main()) diff --git a/tests/test_draft_lock_manager.py b/tests/test_draft_lock_manager.py new file mode 100644 index 0000000..b2cf441 --- /dev/null +++ b/tests/test_draft_lock_manager.py @@ -0,0 +1,351 @@ +""" +草稿并发锁管理器单元测试 +测试 DraftLockManager 的所有功能 + +测试覆盖: +1. 正常场景:锁的获取和释放 +2. 边界场景:超时、重入、并发访问 +3. 异常场景:无效输入、重复释放 +""" +import asyncio +import pytest +import sys +import os + +# 添加项目根目录到 Python 路径 +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) + +from src.utils.draft_lock_manager import DraftLockManager, get_draft_lock_manager + + +class TestDraftLockManager: + """DraftLockManager 测试类""" + + @pytest.fixture + def lock_manager(self): + """创建锁管理器实例""" + return DraftLockManager() + + @pytest.mark.asyncio + async def test_acquire_and_release_lock(self, lock_manager): + """测试基本的锁获取和释放""" + draft_id = "test-draft-001" + + # 获取锁 + result = await lock_manager.acquire_lock(draft_id) + assert result is True + assert lock_manager.is_locked(draft_id) is True + assert lock_manager.get_lock_count(draft_id) == 1 + + # 释放锁 + await lock_manager.release_lock(draft_id) + assert lock_manager.is_locked(draft_id) is False + assert lock_manager.get_lock_count(draft_id) == 0 + + @pytest.mark.asyncio + async def test_acquire_lock_with_timeout(self, lock_manager): + """测试带超时的锁获取""" + draft_id = "test-draft-002" + + # 立即获取锁(无超时) + result = await lock_manager.acquire_lock(draft_id, timeout=None) + assert result is True + + await lock_manager.release_lock(draft_id) + + # 带超时获取锁 + result = await lock_manager.acquire_lock(draft_id, timeout=5.0) + assert result is True + + await lock_manager.release_lock(draft_id) + + @pytest.mark.asyncio + async def test_acquire_lock_timeout_exception(self, lock_manager): + """测试锁获取超时异常""" + draft_id = "test-draft-003" + + # 第一次获取锁 + await lock_manager.acquire_lock(draft_id) + + # 尝试再次获取(应该超时) + with pytest.raises(asyncio.TimeoutError): + await lock_manager.acquire_lock(draft_id, timeout=0.1) + + # 清理 + await lock_manager.release_lock(draft_id) + + @pytest.mark.asyncio + async def test_release_nonexistent_lock(self, lock_manager): + """测试释放不存在的锁""" + draft_id = "nonexistent-draft" + + with pytest.raises(KeyError): + await lock_manager.release_lock(draft_id) + + @pytest.mark.asyncio + async def test_is_locked_method(self, lock_manager): + """测试 is_locked 方法""" + draft_id = "test-draft-004" + + # 未锁定时 + assert lock_manager.is_locked(draft_id) is False + + # 锁定后 + await lock_manager.acquire_lock(draft_id) + assert lock_manager.is_locked(draft_id) is True + + # 释放后 + await lock_manager.release_lock(draft_id) + assert lock_manager.is_locked(draft_id) is False + + @pytest.mark.asyncio + async def test_get_lock_count(self, lock_manager): + """测试获取锁计数""" + draft_id = "test-draft-005" + + # 初始计数为 0 + assert lock_manager.get_lock_count(draft_id) == 0 + + # 获取一次锁 + await lock_manager.acquire_lock(draft_id) + assert lock_manager.get_lock_count(draft_id) == 1 + + # 释放锁 + await lock_manager.release_lock(draft_id) + assert lock_manager.get_lock_count(draft_id) == 0 + + @pytest.mark.asyncio + async def test_get_all_locked_drafts(self, lock_manager): + """测试获取所有被锁定的草稿""" + draft_ids = ["test-draft-006-a", "test-draft-006-b", "test-draft-006-c"] + + # 锁定前两个 + await lock_manager.acquire_lock(draft_ids[0]) + await lock_manager.acquire_lock(draft_ids[1]) + + locked = lock_manager.get_all_locked_drafts() + assert len(locked) == 2 + assert draft_ids[0] in locked + assert draft_ids[1] in locked + assert draft_ids[2] not in locked + + # 释放一个 + await lock_manager.release_lock(draft_ids[0]) + locked = lock_manager.get_all_locked_drafts() + assert len(locked) == 1 + assert draft_ids[1] in locked + + # 全部释放 + await lock_manager.release_lock(draft_ids[1]) + locked = lock_manager.get_all_locked_drafts() + assert len(locked) == 0 + + @pytest.mark.asyncio + async def test_clear_all_locks(self, lock_manager): + """测试清除所有锁""" + draft_ids = ["test-draft-007-a", "test-draft-007-b"] + + # 获取多个锁 + for draft_id in draft_ids: + await lock_manager.acquire_lock(draft_id) + + # 验证都被锁定 + for draft_id in draft_ids: + assert lock_manager.is_locked(draft_id) is True + + # 清除所有锁 + await lock_manager.clear_all_locks() + + # 验证都已释放(锁对象已被删除) + for draft_id in draft_ids: + assert lock_manager.is_locked(draft_id) is False + + @pytest.mark.asyncio + async def test_get_stats(self, lock_manager): + """测试获取统计信息""" + draft_ids = ["test-draft-008-a", "test-draft-008-b"] + + # 初始状态 + stats = lock_manager.get_stats() + assert stats["total_locks"] == 0 + assert stats["locked_drafts"] == 0 + + # 获取锁后 + for draft_id in draft_ids: + await lock_manager.acquire_lock(draft_id) + + stats = lock_manager.get_stats() + assert stats["total_locks"] == 2 + assert stats["locked_drafts"] == 2 + assert stats["total_holders"] == 2 + + # 清理 + for draft_id in draft_ids: + await lock_manager.release_lock(draft_id) + + @pytest.mark.asyncio + async def test_concurrent_access_different_drafts(self, lock_manager): + """测试并发访问不同草稿""" + results = [] + + async def acquire_and_release(draft_id): + await lock_manager.acquire_lock(draft_id) + await asyncio.sleep(0.1) # 模拟工作 + results.append(draft_id) + await lock_manager.release_lock(draft_id) + + # 并发访问 3 个不同的草稿 + tasks = [ + acquire_and_release("draft-a"), + acquire_and_release("draft-b"), + acquire_and_release("draft-c") + ] + + await asyncio.gather(*tasks) + + # 所有任务都应该完成 + assert len(results) == 3 + assert "draft-a" in results + assert "draft-b" in results + assert "draft-c" in results + + @pytest.mark.asyncio + async def test_serial_access_same_draft(self, lock_manager): + """测试串行访问同一草稿""" + draft_id = "test-draft-same" + execution_order = [] + + async def worker(worker_id): + await lock_manager.acquire_lock(draft_id) + try: + execution_order.append(f"{worker_id}_start") + await asyncio.sleep(0.1) + execution_order.append(f"{worker_id}_end") + finally: + await lock_manager.release_lock(draft_id) + + # 3 个 worker 按顺序访问同一草稿(不是并发) + # 因为 asyncio.gather 会同时启动所有任务,但锁会强制它们串行执行 + tasks = [ + worker(1), + worker(2), + worker(3) + ] + + await asyncio.gather(*tasks) + + # 验证执行顺序是串行的(每个 worker 必须等待前一个完成) + # 由于锁的保护,应该是一个接一个执行 + assert len(execution_order) == 6 + + # 验证每个 worker 都是成对的 start->end + for i in range(0, len(execution_order), 2): + worker_num = execution_order[i].split('_')[0] + assert execution_order[i+1] == f"{worker_num}_end" + + @pytest.mark.asyncio + async def test_singleton_pattern(self): + """测试单例模式""" + manager1 = DraftLockManager() + manager2 = DraftLockManager() + + # 应该是同一个实例 + assert manager1 is manager2 + + @pytest.mark.asyncio + async def test_get_draft_lock_manager_function(self): + """测试全局获取器函数""" + manager1 = get_draft_lock_manager() + manager2 = get_draft_lock_manager() + + # 应该是同一个实例 + assert manager1 is manager2 + + +class TestDraftLockManagerEdgeCases: + """DraftLockManager 边界情况测试""" + + @pytest.mark.asyncio + async def test_empty_draft_id(self): + """测试空字符串草稿 ID""" + lock_manager = DraftLockManager() + + # 空字符串也应该能获取锁 + draft_id = "" + result = await lock_manager.acquire_lock(draft_id) + assert result is True + + await lock_manager.release_lock(draft_id) + + @pytest.mark.asyncio + async def test_special_characters_in_draft_id(self): + """测试特殊字符的草稿 ID""" + lock_manager = DraftLockManager() + + special_ids = [ + "draft-with-dash", + "draft_with_underscore", + "draft.with.dots", + "draft/with/slashes", + "draft@with#special$chars" + ] + + for draft_id in special_ids: + await lock_manager.acquire_lock(draft_id) + assert lock_manager.is_locked(draft_id) is True + await lock_manager.release_lock(draft_id) + assert lock_manager.is_locked(draft_id) is False + + @pytest.mark.asyncio + async def test_very_long_draft_id(self): + """测试超长草稿 ID""" + lock_manager = DraftLockManager() + + # 创建一个很长的 ID + long_id = "draft-" + "a" * 1000 + + await lock_manager.acquire_lock(long_id) + assert lock_manager.is_locked(long_id) is True + await lock_manager.release_lock(long_id) + + @pytest.mark.asyncio + async def test_rapid_acquire_release_cycle(self): + """测试快速连续获取释放循环""" + lock_manager = DraftLockManager() + draft_id = "rapid-cycle" + + # 快速循环 100 次 + for i in range(100): + await lock_manager.acquire_lock(draft_id) + await lock_manager.release_lock(draft_id) + + # 最后应该是未锁定状态 + assert lock_manager.is_locked(draft_id) is False + + @pytest.mark.asyncio + async def test_multiple_workers_same_draft_stress(self): + """测试多 worker 压力测试""" + lock_manager = DraftLockManager() + draft_id = "stress-test" + counter = 0 + + async def increment_counter(): + nonlocal counter + await lock_manager.acquire_lock(draft_id) + try: + current = counter + await asyncio.sleep(0.01) # 增加竞争 + counter = current + 1 + finally: + await lock_manager.release_lock(draft_id) + + # 10 个 worker 同时尝试增加计数器 + tasks = [increment_counter() for _ in range(10)] + await asyncio.gather(*tasks) + + # 由于锁的保护,counter 应该正好是 10 + assert counter == 10 + + +if __name__ == "__main__": + pytest.main([__file__, "-v"])