diff --git a/packages/shared/ai_service.py b/packages/shared/ai_service.py index c8a066b1d..32014d23d 100755 --- a/packages/shared/ai_service.py +++ b/packages/shared/ai_service.py @@ -339,7 +339,7 @@ def _call_ai_cover_service( logger.info("调用 MediaKit 抽帧: plan_id=%s video=%s", plan_id, primary_video_url[:80]) frames = client.extract_frames( video_url=primary_video_url, - strategy="SceneChange", + strategy="TimeInterval", max_frames=5, ) diff --git a/packages/shared/mediakit_client.py b/packages/shared/mediakit_client.py index 1bb6bf617..fb4effd36 100755 --- a/packages/shared/mediakit_client.py +++ b/packages/shared/mediakit_client.py @@ -23,6 +23,15 @@ from packages.shared.config import get_shared_settings logger = logging.getLogger(__name__) +# MediaKit 服务端偶发 OOM 导致的可重试错误关键词 +_RETRYABLE_ERROR_KEYWORDS = ("signal: killed", "InternalError", "OOM", "out of memory") + + +def _is_retryable_error(error_msg: str) -> bool: + """判断错误是否可重试(MediaKit 服务端偶发 OOM 等).""" + error_lower = error_msg.lower() + return any(kw.lower() in error_lower for kw in _RETRYABLE_ERROR_KEYWORDS) + class MediaKitClient: """MediaKit API 客户端. @@ -45,23 +54,25 @@ class MediaKitClient: def extract_frames( self, video_url: str, - strategy: str = "SceneChange", + strategy: str = "TimeInterval", max_frames: int = 10, poll_interval: float = 2.0, max_poll_attempts: int = 30, + max_retries: int = 1, ) -> Optional[List[Dict[str, Any]]]: """调用 MediaKit 视频抽帧接口. Args: video_url: 视频 URL(需可公开访问) - strategy: 抽帧策略 - - TimeInterval: 按固定时间间隔 + strategy: 抽帧策略(默认 TimeInterval,比 SceneChange 更稳定不易 OOM) + - TimeInterval: 按固定时间间隔(推荐,稳定性好) - SpecifiedTime: 按指定时间点 - SpecifiedFrames: 首尾帧 + 指定帧数 - - SceneChange: 场景变化检测(推荐用于封面选取) + - SceneChange: 场景变化检测(封面选取可用,但高分辨率视频易 OOM) max_frames: 最大返回帧数 poll_interval: 轮询间隔(秒) max_poll_attempts: 最大轮询次数 + max_retries: 失败后自动重试次数(仅对可重试错误如 OOM 生效) Returns: 帧列表 [{"image_url": "...", "timestamp": 1.5}, ...] @@ -71,13 +82,42 @@ class MediaKitClient: logger.warning("MediaKit 未配置,跳过抽帧") return None - # 提交抽帧任务 - task_id = self._submit_extract_task(video_url, strategy, max_frames) - if not task_id: + for attempt in range(1 + max_retries): + # 提交抽帧任务 + task_id = self._submit_extract_task(video_url, strategy, max_frames) + if not task_id: + return None + + # 轮询任务状态 + result, error_msg = self._poll_task_result_with_error(task_id, poll_interval, max_poll_attempts) + + if result is not None: + return result + + # 任务失败,判断是否可重试 + if error_msg and _is_retryable_error(error_msg) and attempt < max_retries: + logger.warning( + "MediaKit 抽帧遇到可重试错误,%ds 后重试: " "task_id=%s attempt=%d/%d error=%s", + 2, + task_id, + attempt + 1, + max_retries, + error_msg, + ) + time.sleep(2) + continue + + # 不可重试或已用尽重试次数 + if error_msg: + logger.error( + "MediaKit 抽帧最终失败: task_id=%s retryable=%s error=%s", + task_id, + _is_retryable_error(error_msg), + error_msg, + ) return None - # 轮询任务状态 - return self._poll_task_result(task_id, poll_interval, max_poll_attempts) + return None def _submit_extract_task( self, @@ -120,8 +160,9 @@ class MediaKitClient: video_urls: List[str], prompt: str, level: str = "Economy", - poll_interval: float = 3.0, - max_poll_attempts: int = 60, + poll_interval: float = 2.0, + max_poll_attempts: int = 15, + max_retries: int = 1, ) -> Optional[List[str]]: """调用 MediaKit 视频理解智能策略 API. @@ -134,6 +175,7 @@ class MediaKitClient: level: 分析档位 Economy / Balanced / Quality poll_interval: 轮询间隔(秒) max_poll_attempts: 最大轮询次数 + max_retries: 失败后自动重试次数(仅对可重试错误如 OOM 生效) Returns: 分析结果列表,每个元素对应 video_urls 中同索引视频的分析文本。 @@ -146,11 +188,40 @@ class MediaKitClient: if not video_urls: return None - task_id = self._submit_video_understand_task(video_urls, prompt, level) - if not task_id: + for attempt in range(1 + max_retries): + task_id = self._submit_video_understand_task(video_urls, prompt, level) + if not task_id: + return None + + result, error_msg = self._poll_video_understand_result_with_error(task_id, poll_interval, max_poll_attempts) + + if result is not None: + return result + + # 任务失败,判断是否可重试 + if error_msg and _is_retryable_error(error_msg) and attempt < max_retries: + logger.warning( + "MediaKit 视频理解遇到可重试错误,%ds 后重试: " "task_id=%s attempt=%d/%d error=%s", + 2, + task_id, + attempt + 1, + max_retries, + error_msg, + ) + time.sleep(2) + continue + + # 不可重试或已用尽重试次数 + if error_msg: + logger.error( + "MediaKit 视频理解最终失败: task_id=%s retryable=%s error=%s", + task_id, + _is_retryable_error(error_msg), + error_msg, + ) return None - return self._poll_video_understand_result(task_id, poll_interval, max_poll_attempts) + return None def _submit_video_understand_task( self, @@ -193,13 +264,16 @@ class MediaKitClient: logger.exception("MediaKit 视频理解任务提交异常: %s", str(e)) return None - def _poll_video_understand_result( + def _poll_video_understand_result_with_error( self, task_id: str, poll_interval: float, max_poll_attempts: int, - ) -> Optional[List[str]]: - """轮询视频理解任务结果,返回 contents 列表.""" + ) -> tuple[Optional[List[str]], Optional[str]]: + """轮询视频理解任务结果,返回 (contents, error_msg). + + 成功时返回 (contents, None),失败时返回 (None, error_message)。 + """ url = f"{self.base_url}/tasks/{task_id}" headers = { "Authorization": f"Bearer {self.api_key}", @@ -220,16 +294,18 @@ class MediaKitClient: task_id, len(contents), ) - return contents if contents else None + return contents if contents else None, None elif status == "failed": - error_msg = data.get("error", "unknown error") + error_detail = data.get("error", "unknown error") + error_node = data.get("error_node", "") + full_error = f"{error_detail}" + (f" [node={error_node}]" if error_node else "") logger.error( "MediaKit 视频理解任务失败: task_id=%s error=%s", task_id, - error_msg, + full_error, ) - return None + return None, full_error logger.debug( "MediaKit 视频理解进行中: task_id=%s status=%s attempt=%d/%d", @@ -248,20 +324,34 @@ class MediaKitClient: ) time.sleep(poll_interval) + error_msg = f"超时: max_attempts={max_poll_attempts}" logger.error( "MediaKit 视频理解超时: task_id=%s max_attempts=%d", task_id, max_poll_attempts, ) - return None + return None, error_msg - def _poll_task_result( + def _poll_video_understand_result( self, task_id: str, poll_interval: float, max_poll_attempts: int, - ) -> Optional[List[Dict[str, Any]]]: - """轮询抽帧任务状态,返回结果.""" + ) -> Optional[List[str]]: + """轮询视频理解任务结果,返回 contents 列表(兼容旧接口).""" + result, _ = self._poll_video_understand_result_with_error(task_id, poll_interval, max_poll_attempts) + return result + + def _poll_task_result_with_error( + self, + task_id: str, + poll_interval: float, + max_poll_attempts: int, + ) -> tuple[Optional[List[Dict[str, Any]]], Optional[str]]: + """轮询抽帧任务状态,返回 (snapshots, error_msg). + + 成功时返回 (snapshots, None),失败时返回 (None, error_message)。 + """ url = f"{self.base_url}/tasks/{task_id}" headers = { "Authorization": f"Bearer {self.api_key}", @@ -282,12 +372,18 @@ class MediaKitClient: task_id, len(snapshots), ) - return snapshots + return snapshots, None elif status == "failed": - error_msg = data.get("error", "unknown error") - logger.error("MediaKit 抽帧任务失败: task_id=%s error=%s", task_id, error_msg) - return None + error_detail = data.get("error", "unknown error") + error_node = data.get("error_node", "") + full_error = f"{error_detail}" + (f" [node={error_node}]" if error_node else "") + logger.error( + "MediaKit 抽帧任务失败: task_id=%s error=%s", + task_id, + full_error, + ) + return None, full_error # status == "processing" or "pending" logger.debug( @@ -300,11 +396,30 @@ class MediaKitClient: time.sleep(poll_interval) except Exception as e: - logger.exception("MediaKit 轮询异常: task_id=%s error=%s", task_id, str(e)) + logger.exception( + "MediaKit 抽帧轮询异常: task_id=%s error=%s", + task_id, + str(e), + ) time.sleep(poll_interval) - logger.error("MediaKit 抽帧超时: task_id=%s max_attempts=%d", task_id, max_poll_attempts) - return None + error_msg = f"超时: max_attempts={max_poll_attempts}" + logger.error( + "MediaKit 抽帧超时: task_id=%s max_attempts=%d", + task_id, + max_poll_attempts, + ) + return None, error_msg + + def _poll_task_result( + self, + task_id: str, + poll_interval: float, + max_poll_attempts: int, + ) -> Optional[List[Dict[str, Any]]]: + """轮询抽帧任务状态,返回结果(兼容旧接口).""" + result, _ = self._poll_task_result_with_error(task_id, poll_interval, max_poll_attempts) + return result # ── 单例管理 ──────────────────────────────────────────────────────────────── diff --git a/tests/unit/test_mediakit_retry.py b/tests/unit/test_mediakit_retry.py new file mode 100644 index 000000000..79e4dcabb --- /dev/null +++ b/tests/unit/test_mediakit_retry.py @@ -0,0 +1,401 @@ +"""MediaKit 重试逻辑 + 错误日志增强测试. + +覆盖范围: +1. _is_retryable_error — 可重试错误判断 +2. extract_frames 重试 — signal: killed / InternalError 自动重试一次 +3. analyze_videos 重试 — 同上 +4. 不可重试错误不触发重试 +5. 默认抽帧策略为 TimeInterval +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, Mock, patch + +import pytest + +from packages.shared.mediakit_client import ( + MediaKitClient, + _is_retryable_error, +) + +# ── _is_retryable_error ────────────────────────────────────────────────────── + + +class TestIsRetryableError: + """可重试错误判断.""" + + def test_signal_killed(self): + """signal: killed 可重试.""" + assert _is_retryable_error("signal: killed") is True + + def test_internal_error(self): + """InternalError 可重试.""" + assert _is_retryable_error("InternalError: something went wrong") is True + + def test_oom(self): + """OOM 可重试.""" + assert _is_retryable_error("ExtractFrames OOM") is True + + def test_out_of_memory(self): + """out of memory 可重试.""" + assert _is_retryable_error("process out of memory") is True + + def test_model_timeout_not_retryable(self): + """model timeout 不可重试.""" + assert _is_retryable_error("model timeout") is False + + def test_network_error_not_retryable(self): + """普通网络错误不可重试.""" + assert _is_retryable_error("Connection refused") is False + + def test_empty_string(self): + """空字符串不可重试.""" + assert _is_retryable_error("") is False + + def test_case_insensitive(self): + """大小写不敏感.""" + assert _is_retryable_error("SIGNAL: KILLED") is True + assert _is_retryable_error("internalerror") is True + + +# ── extract_frames 重试逻辑 ────────────────────────────────────────────────── + + +class TestExtractFramesRetry: + """extract_frames 重试逻辑测试.""" + + def _make_client(self, api_key: str = "test-key") -> MediaKitClient: + with patch("packages.shared.mediakit_client.get_shared_settings") as mock_settings: + mock_settings.return_value = MagicMock( + mediakit_api_key=api_key, + mediakit_base_url="https://mock.mediakit.com/api/v1", + mediakit_timeout=30, + ) + return MediaKitClient() + + @patch("packages.shared.mediakit_client.time.sleep") + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_retry_on_signal_killed_then_success(self, mock_post, mock_get, mock_sleep): + """signal: killed 后重试成功.""" + # 第一次提交成功 + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-retry-1"}, + raise_for_status=MagicMock(), + ) + + # 第一次轮询返回 failed(signal: killed),第二次返回 success + call_count = [0] + + def get_side_effect(*args, **kwargs): + call_count[0] += 1 + if call_count[0] == 1: + # 第一次轮询:失败 + return MagicMock( + status_code=200, + json=lambda: {"status": "failed", "error": "signal: killed"}, + raise_for_status=MagicMock(), + ) + else: + # 第二次轮询(重试后):成功 + return MagicMock( + status_code=200, + json=lambda: { + "status": "success", + "result": {"snapshots": [{"image_url": "https://example.com/frame.jpg", "timestamp": 1.0}]}, + }, + raise_for_status=MagicMock(), + ) + + mock_get.side_effect = get_side_effect + + client = self._make_client() + frames = client.extract_frames( + video_url="https://example.com/video.mp4", + strategy="TimeInterval", + max_frames=5, + poll_interval=0.01, + max_retries=1, + ) + + assert frames is not None + assert len(frames) == 1 + assert frames[0]["image_url"] == "https://example.com/frame.jpg" + + # 验证 sleep 被调用(重试间隔) + mock_sleep.assert_called() + + @patch("packages.shared.mediakit_client.time.sleep") + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_no_retry_on_non_retryable_error(self, mock_post, mock_get, mock_sleep): + """不可重试错误直接返回 None,不重试.""" + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-fail"}, + raise_for_status=MagicMock(), + ) + + mock_get.return_value = MagicMock( + status_code=200, + json=lambda: {"status": "failed", "error": "model timeout"}, + raise_for_status=MagicMock(), + ) + + client = self._make_client() + frames = client.extract_frames( + video_url="https://example.com/video.mp4", + strategy="TimeInterval", + max_frames=5, + poll_interval=0.01, + max_retries=1, + ) + + assert frames is None + # 不可重试错误不应该调用 sleep(不重试) + mock_sleep.assert_not_called() + + @patch("packages.shared.mediakit_client.time.sleep") + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_retry_exhausted_returns_none(self, mock_post, mock_get, mock_sleep): + """重试次数用尽后返回 None.""" + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-fail-retry"}, + raise_for_status=MagicMock(), + ) + + # 每次都返回 signal: killed + mock_get.return_value = MagicMock( + status_code=200, + json=lambda: {"status": "failed", "error": "signal: killed"}, + raise_for_status=MagicMock(), + ) + + client = self._make_client() + frames = client.extract_frames( + video_url="https://example.com/video.mp4", + strategy="TimeInterval", + max_frames=5, + poll_interval=0.01, + max_retries=1, # 最多重试 1 次 + ) + + assert frames is None + # post 被调用 2 次(原始 + 1 次重试) + assert mock_post.call_count == 2 + + def test_default_strategy_is_time_interval(self): + """默认抽帧策略为 TimeInterval(不是 SceneChange).""" + with patch("packages.shared.mediakit_client.httpx.post") as mock_post: + with patch("packages.shared.mediakit_client.httpx.get") as mock_get: + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-default"}, + raise_for_status=MagicMock(), + ) + mock_get.return_value = MagicMock( + status_code=200, + json=lambda: { + "status": "success", + "result": {"snapshots": [{"image_url": "url", "timestamp": 0.0}]}, + }, + raise_for_status=MagicMock(), + ) + + client = self._make_client() + # 不传 strategy,使用默认值 + client.extract_frames( + video_url="https://example.com/video.mp4", + poll_interval=0.01, + ) + + # 验证提交时使用了 TimeInterval + call_args = mock_post.call_args + payload = call_args.kwargs["json"] + assert payload["strategy"] == "TimeInterval" + + +# ── analyze_videos 重试逻辑 ────────────────────────────────────────────────── + + +class TestAnalyzeVideosRetry: + """analyze_videos 重试逻辑测试.""" + + def _make_client(self, api_key: str = "test-key") -> MediaKitClient: + with patch("packages.shared.mediakit_client.get_shared_settings") as mock_settings: + mock_settings.return_value = MagicMock( + mediakit_api_key=api_key, + mediakit_base_url="https://mock.mediakit.com/api/v1", + mediakit_timeout=30, + ) + return MediaKitClient() + + @patch("packages.shared.mediakit_client.time.sleep") + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_retry_on_internal_error_then_success(self, mock_post, mock_get, mock_sleep): + """InternalError 后重试成功.""" + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-retry-vu"}, + raise_for_status=MagicMock(), + ) + + call_count = [0] + + def get_side_effect(*args, **kwargs): + call_count[0] += 1 + if call_count[0] == 1: + return MagicMock( + status_code=200, + json=lambda: {"status": "failed", "error": "InternalError: service busy"}, + raise_for_status=MagicMock(), + ) + else: + return MagicMock( + status_code=200, + json=lambda: { + "status": "completed", + "result": {"contents": ["视频展示了一只猫在沙发上睡觉"]}, + }, + raise_for_status=MagicMock(), + ) + + mock_get.side_effect = get_side_effect + + client = self._make_client() + result = client.analyze_videos( + video_urls=["https://example.com/cat.mp4"], + prompt="描述视频内容", + poll_interval=0.01, + max_retries=1, + ) + + assert result is not None + assert len(result) == 1 + assert "猫" in result[0] + mock_sleep.assert_called() + + @patch("packages.shared.mediakit_client.time.sleep") + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_no_retry_on_unknown_error(self, mock_post, mock_get, mock_sleep): + """未知错误不重试.""" + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-unknown-fail"}, + raise_for_status=MagicMock(), + ) + + mock_get.return_value = MagicMock( + status_code=200, + json=lambda: {"status": "failed", "error": "unknown error"}, + raise_for_status=MagicMock(), + ) + + client = self._make_client() + result = client.analyze_videos( + video_urls=["https://example.com/v.mp4"], + prompt="describe", + poll_interval=0.01, + max_retries=1, + ) + + assert result is None + mock_sleep.assert_not_called() + # post 只调用 1 次(不重试) + assert mock_post.call_count == 1 + + +# ── 错误日志增强 ───────────────────────────────────────────────────────────── + + +class TestErrorLogging: + """错误日志包含完整信息.""" + + def _make_client(self, api_key: str = "test-key") -> MediaKitClient: + with patch("packages.shared.mediakit_client.get_shared_settings") as mock_settings: + mock_settings.return_value = MagicMock( + mediakit_api_key=api_key, + mediakit_base_url="https://mock.mediakit.com/api/v1", + mediakit_timeout=30, + ) + return MediaKitClient() + + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_extract_frames_error_includes_task_id(self, mock_post, mock_get, caplog): + """抽帧失败日志包含 task_id 和完整 error.""" + import logging + + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-log-test"}, + raise_for_status=MagicMock(), + ) + + mock_get.return_value = MagicMock( + status_code=200, + json=lambda: { + "status": "failed", + "error": "signal: killed", + "error_node": "ExtractFrames", + }, + raise_for_status=MagicMock(), + ) + + client = self._make_client() + + with caplog.at_level(logging.ERROR): + client.extract_frames( + video_url="https://example.com/video.mp4", + strategy="TimeInterval", + max_frames=5, + poll_interval=0.01, + max_retries=0, # 不重试,直接失败 + ) + + # 验证日志中包含 task_id 和 error 信息 + error_logs = [r.message for r in caplog.records if r.levelno >= logging.ERROR] + assert any("task-log-test" in msg for msg in error_logs) + assert any("signal: killed" in msg for msg in error_logs) + + @patch("packages.shared.mediakit_client.httpx.get") + @patch("packages.shared.mediakit_client.httpx.post") + def test_analyze_videos_error_includes_full_error(self, mock_post, mock_get, caplog): + """视频理解失败日志包含完整 error 和 error_node.""" + import logging + + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: {"task_id": "task-vu-log"}, + raise_for_status=MagicMock(), + ) + + mock_get.return_value = MagicMock( + status_code=200, + json=lambda: { + "status": "failed", + "error": "signal: killed in ExtractFrames", + "error_node": "ExtractFrames", + }, + raise_for_status=MagicMock(), + ) + + client = self._make_client() + + with caplog.at_level(logging.ERROR): + client.analyze_videos( + video_urls=["https://example.com/v.mp4"], + prompt="describe", + poll_interval=0.01, + max_retries=0, + ) + + error_logs = [r.message for r in caplog.records if r.levelno >= logging.ERROR] + assert any("task-vu-log" in msg for msg in error_logs) + assert any("ExtractFrames" in msg for msg in error_logs)