diff --git a/apps/worker/worker_app/tasks/duplication_check.py b/apps/worker/worker_app/tasks/duplication_check.py index dcfa49c17..3746f7108 100644 --- a/apps/worker/worker_app/tasks/duplication_check.py +++ b/apps/worker/worker_app/tasks/duplication_check.py @@ -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: diff --git a/tests/unit/test_duplication_api_enqueue.py b/tests/unit/test_duplication_api_enqueue.py new file mode 100644 index 000000000..9d58798a6 --- /dev/null +++ b/tests/unit/test_duplication_api_enqueue.py @@ -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