"""成片中心新功能单元测试。 覆盖:分页列表、复核状态更新、缩略图更新、批量获取、use case。 """ import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) import pytest from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository from packages.adapters.sqlalchemy_impl.models import Base from packages.application.generated_videos import ( GetVideosByIdsUseCase, ListGeneratedVideosPaginatedUseCase, UpdateVideoReviewStatusUseCase, ) from packages.domain import GeneratedVideo def _repository(): engine = create_engine("sqlite:///:memory:") Base.metadata.create_all(engine) session = sessionmaker(bind=engine)() return SQLAlchemyGeneratedVideoRepository(session) def _create_video(repo, project_id="proj-1", status="completed", review_status="pending_review", idx=1): video = GeneratedVideo.create( project_id=project_id, generation_task_id=f"task-{idx}", name=f"video-{idx}.mp4", file_url=f"generated/video-{idx}.mp4", file_size=1024 * idx, duration=10.0 * idx, width=1280, height=720, fps=25.0, ) video.status = status video.review_status = review_status repo.create(video) return video class TestGeneratedVideoRepository: """GeneratedVideoRepository 新方法测试。""" def test_list_paginated_default(self): repo = _repository() for i in range(5): _create_video(repo, idx=i) items, total = repo.list_paginated(page=1, page_size=3) assert total == 5 assert len(items) == 3 # 按 generated_at 倒序,最新的在前 assert items[0].name == "video-4.mp4" def test_list_paginated_by_project(self): repo = _repository() _create_video(repo, project_id="proj-a", idx=1) _create_video(repo, project_id="proj-a", idx=2) _create_video(repo, project_id="proj-b", idx=3) items, total = repo.list_paginated(project_id="proj-a") assert total == 2 assert all(i.project_id == "proj-a" for i in items) def test_list_paginated_by_status(self): repo = _repository() _create_video(repo, status="completed", idx=1) _create_video(repo, status="completed", idx=2) _create_video(repo, status="failed", idx=3) items, total = repo.list_paginated(status="completed") assert total == 2 assert all(i.status == "completed" for i in items) def test_list_paginated_by_review_status(self): repo = _repository() _create_video(repo, review_status="pending_review", idx=1) _create_video(repo, review_status="approved", idx=2) _create_video(repo, review_status="rejected", idx=3) items, total = repo.list_paginated(review_status="approved") assert total == 1 assert items[0].review_status == "approved" def test_list_paginated_multi_filter(self): repo = _repository() _create_video(repo, project_id="p1", status="completed", review_status="approved", idx=1) _create_video(repo, project_id="p1", status="completed", review_status="pending_review", idx=2) _create_video(repo, project_id="p2", status="completed", review_status="approved", idx=3) items, total = repo.list_paginated(project_id="p1", review_status="approved") assert total == 1 assert items[0].project_id == "p1" assert items[0].review_status == "approved" def test_update_review_status(self): repo = _repository() video = _create_video(repo, idx=1) result = repo.update_review_status(video.id, "approved") assert result is not None assert result.review_status == "approved" # 验证持久化 saved = repo.get(video.id) assert saved.review_status == "approved" def test_update_review_status_not_found(self): repo = _repository() result = repo.update_review_status("nonexistent", "approved") assert result is None def test_update_thumbnail(self): repo = _repository() video = _create_video(repo, idx=1) assert video.thumbnail_url is None ok = repo.update_thumbnail(video.id, "https://oss/thumb.jpg") assert ok is True saved = repo.get(video.id) assert saved.thumbnail_url == "https://oss/thumb.jpg" def test_update_thumbnail_not_found(self): repo = _repository() ok = repo.update_thumbnail("nonexistent", "https://oss/thumb.jpg") assert ok is False def test_get_by_ids(self): repo = _repository() v1 = _create_video(repo, idx=1) v2 = _create_video(repo, idx=2) v3 = _create_video(repo, idx=3) result = repo.get_by_ids([v1.id, v3.id]) assert len(result) == 2 ids = {v.id for v in result} assert v1.id in ids assert v3.id in ids def test_get_by_ids_empty(self): repo = _repository() result = repo.get_by_ids([]) assert result == [] class TestGeneratedVideoUseCases: """Use case 层测试。""" def test_list_paginated_use_case(self): repo = _repository() for i in range(10): _create_video(repo, idx=i) use_case = ListGeneratedVideosPaginatedUseCase(repo) items, total = use_case.execute(page=2, page_size=3) assert total == 10 assert len(items) == 3 def test_list_paginated_use_case_page_clamp(self): repo = _repository() use_case = ListGeneratedVideosPaginatedUseCase(repo) # page < 1 应该被修正为 1 items, total = use_case.execute(page=0, page_size=20) assert total == 0 def test_list_paginated_use_case_page_size_clamp(self): repo = _repository() use_case = ListGeneratedVideosPaginatedUseCase(repo) # page_size > 100 应该被修正为 20 for i in range(30): _create_video(repo, idx=i) items, total = use_case.execute(page=1, page_size=200) assert total == 30 assert len(items) == 20 # clamp 到默认 20 def test_update_review_status_use_case(self): repo = _repository() video = _create_video(repo, idx=1) use_case = UpdateVideoReviewStatusUseCase(repo) result = use_case.execute(video.id, "approved") assert result is not None assert result.review_status == "approved" def test_update_review_status_use_case_invalid_status(self): repo = _repository() video = _create_video(repo, idx=1) use_case = UpdateVideoReviewStatusUseCase(repo) with pytest.raises(ValueError, match="无效的 review_status"): use_case.execute(video.id, "invalid_status") def test_update_review_status_use_case_empty_id(self): repo = _repository() use_case = UpdateVideoReviewStatusUseCase(repo) with pytest.raises(ValueError, match="video_id 不能为空"): use_case.execute("", "approved") def test_get_by_ids_use_case(self): repo = _repository() v1 = _create_video(repo, idx=1) v2 = _create_video(repo, idx=2) use_case = GetVideosByIdsUseCase(repo) result = use_case.execute([v1.id, v2.id]) assert len(result) == 2