fix: 查重 worker 无限重试 bug + 补 3 个 API/repository 单测 (#1661 follow-up) #1680

Merged
auto-approve-bot merged 2 commits from fix/duplication-ci-coverage-retry into develop 2026-09-04 11:38:12 +08:00
2 changed files with 178 additions and 8 deletions
@@ -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:
+168
View File
@@ -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