fix: 查重 worker 无限重试 bug + 补 3 个 API/repository 单测 (#1661 follow-up) #1680
@@ -175,19 +175,21 @@ def process_duplication_check(self: Task, record_id: str) -> dict:
|
||||
logger.error("Duplication check failed for record %s: %s", record_id, e, exc_info=True)
|
||||
if session is not None:
|
||||
session.rollback()
|
||||
# 本次是最后一次执行机会(retries 从 0 计数,达到 max_retries 说明重试已耗尽),
|
||||
# 标记 failed;否则保持 pending 由 Celery 60 秒后重试
|
||||
try:
|
||||
if "repo" in locals() and self.request.retries >= self.max_retries:
|
||||
# 超过重试上限:标记 failed 并返回失败结果,不再 retry
|
||||
if "repo" in locals() and self.request.retries >= self.max_retries:
|
||||
try:
|
||||
failed_record = repo.get(record_id)
|
||||
if failed_record is not None and failed_record.status != "failed":
|
||||
failed_record.mark_failed(f"查重失败(已重试{self.max_retries}次): {e}")
|
||||
repo.update(failed_record)
|
||||
session.commit()
|
||||
except Exception as inner:
|
||||
logger.error("Failed to mark duplication record %s as failed: %s", record_id, inner)
|
||||
session.rollback()
|
||||
raise self.retry(exc=e, countdown=60) from e
|
||||
except Exception as inner:
|
||||
logger.error("Failed to mark duplication record %s as failed: %s", record_id, inner)
|
||||
session.rollback()
|
||||
return {"ok": False, "record_id": record_id, "status": "failed", "error": str(e)}
|
||||
# 未达上限:60 秒后重试
|
||||
raise self.retry(exc=e, countdown=60) from e
|
||||
return {"ok": False, "record_id": record_id, "status": "failed", "error": str(e)}
|
||||
|
||||
finally:
|
||||
if session is not None:
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""#1661 查重 API enqueue 及仓储 commit 覆盖测试。
|
||||
|
||||
覆盖:
|
||||
- upload 接口在成功后调用 celery_app.send_task
|
||||
- retry 接口在成功后调用 celery_app.send_task
|
||||
- duplication_repository.update() 正确调用 session.commit()
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
ROOT = os.path.join(os.path.dirname(__file__), "..", "..")
|
||||
sys.path.insert(0, os.path.join(ROOT, "apps", "api"))
|
||||
sys.path.insert(0, os.path.join(ROOT, "packages"))
|
||||
|
||||
from app.api.routes.duplication import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import get_duplication_repository
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from packages.domain.duplication import DuplicationRecord
|
||||
from packages.domain.entities import User
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_test_user():
|
||||
return User(id="user-1", username="testuser", email="test@example.com", display_name="Test User")
|
||||
|
||||
|
||||
def _make_auth_user():
|
||||
return AuthenticatedUser(user=_make_test_user(), session_id="test-session", token_type="bearer")
|
||||
|
||||
|
||||
def _make_record(status="pending"):
|
||||
record = DuplicationRecord.create(
|
||||
user_id="user-1",
|
||||
filename="test.mp4",
|
||||
file_size=1024,
|
||||
storage_key="duplication/abc/test.mp4",
|
||||
)
|
||||
if status != "pending":
|
||||
record.status = status
|
||||
return record
|
||||
|
||||
|
||||
def _build_client(auth_user, repo, storage=None):
|
||||
"""构建带 dependency_overrides 的 TestClient。"""
|
||||
app = FastAPI()
|
||||
app.include_router(router, prefix="/duplication")
|
||||
app.dependency_overrides[get_current_user] = lambda: auth_user
|
||||
app.dependency_overrides[get_duplication_repository] = lambda: repo
|
||||
if storage is not None:
|
||||
app.dependency_overrides[get_storage_service] = lambda: storage
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. Upload endpoint enqueues celery task
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_upload_enqueue_calls_celery_task():
|
||||
"""POST /duplication/upload 成功创建记录后必须调用 send_task。"""
|
||||
record = _make_record()
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.create.return_value = record
|
||||
|
||||
fake_storage = MagicMock()
|
||||
fake_auth = _make_auth_user()
|
||||
|
||||
client = _build_client(fake_auth, fake_repo, fake_storage)
|
||||
|
||||
with patch("app.api.routes.duplication.celery_app") as mock_celery:
|
||||
response = client.post(
|
||||
"/duplication/upload",
|
||||
files={"file": ("test.mp4", b"fake-video-content", "video/mp4")},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
mock_celery.send_task.assert_called_once_with(
|
||||
"worker.process_duplication_check",
|
||||
args=[record.id],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Retry endpoint enqueues celery task
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_retry_enqueue_calls_celery_task():
|
||||
"""POST /duplication/records/{id}/retry 成功后必须调用 send_task。"""
|
||||
record = _make_record(status="failed")
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = record
|
||||
|
||||
# RetryDuplicationUseCase.execute 内部调用 repo.get → record.reset_for_retry → repo.update
|
||||
updated = _make_record()
|
||||
updated.id = record.id
|
||||
updated.status = "pending"
|
||||
fake_repo.update.return_value = updated
|
||||
|
||||
fake_auth = _make_auth_user()
|
||||
|
||||
client = _build_client(fake_auth, fake_repo)
|
||||
|
||||
with patch("app.api.routes.duplication.celery_app") as mock_celery:
|
||||
response = client.post(f"/duplication/records/{record.id}/retry")
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
mock_celery.send_task.assert_called_once_with(
|
||||
"worker.process_duplication_check",
|
||||
args=[record.id],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Repository update calls session.commit()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_repository_update_calls_session_commit():
|
||||
"""duplication_repository 的 update 方法必须调用 session.commit()。"""
|
||||
from packages.adapters.sqlalchemy_impl.duplication_repository import (
|
||||
SQLAlchemyDuplicationRecordRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import DuplicationRecordModel
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_model = MagicMock(spec=DuplicationRecordModel)
|
||||
mock_model.id = "rec-1"
|
||||
|
||||
mock_session.query.return_value.filter.return_value.first.return_value = mock_model
|
||||
|
||||
repo = SQLAlchemyDuplicationRecordRepository(mock_session)
|
||||
|
||||
record = DuplicationRecord.create(
|
||||
user_id="user-1",
|
||||
filename="test.mp4",
|
||||
file_size=1024,
|
||||
storage_key="duplication/abc/test.mp4",
|
||||
)
|
||||
record.status = "completed"
|
||||
record.duplicate_rate = 42.0
|
||||
record.duplicate_count = 1
|
||||
record.visual_similarity = 0.85
|
||||
record.match_count = 2
|
||||
|
||||
result = repo.update(record)
|
||||
|
||||
mock_session.commit.assert_called()
|
||||
assert result.visual_similarity == 0.85
|
||||
assert result.match_count == 2
|
||||
Reference in New Issue
Block a user