diff --git a/apps/api/app/api/routes/assets.py b/apps/api/app/api/routes/assets.py index f51a25f0b..2cbfaf184 100755 --- a/apps/api/app/api/routes/assets.py +++ b/apps/api/app/api/routes/assets.py @@ -671,20 +671,28 @@ def create_asset( asset_library_repository: Any = Depends(get_asset_library_repository), project_repository: Any = Depends(get_project_repository), ) -> AssetResponse: - project = project_repository.find_by_id(request.project_id) + # 先获取素材库,用于推导 project_id(前端可能不传) + library = asset_library_repository.get(request.library_id) + if library is None: + raise HTTPException(status_code=404, detail=f"AssetLibrary {request.library_id} not found") + + # project_id 自动推导:优先用请求值,否则从 library 关联的项目获取 + project_id = request.project_id or library.project_id + + project = project_repository.find_by_id(project_id) if project is None: - raise HTTPException(status_code=404, detail=f"Project {request.project_id} not found") + raise HTTPException(status_code=404, detail=f"Project {project_id} not found") if not project.can_access(authenticated_user.user.id): raise HTTPException(status_code=403, detail="Access denied to project") - library = asset_library_repository.get(request.library_id) - if library is None or library.project_id != request.project_id: - raise HTTPException(status_code=404, detail=f"AssetLibrary {request.library_id} not found") + # 确保 library 和 project 归属一致 + if library.project_id != project_id: + raise HTTPException(status_code=400, detail="AssetLibrary does not belong to the specified project") use_case = CreateAssetUseCase(asset_repository) item = use_case.execute( CreateAssetCommand( - project_id=request.project_id, + project_id=project_id, library_id=request.library_id, name=request.name, storage_key=request.storage_key, diff --git a/apps/api/app/schemas/asset.py b/apps/api/app/schemas/asset.py index b839a9e4c..528a77f3d 100755 --- a/apps/api/app/schemas/asset.py +++ b/apps/api/app/schemas/asset.py @@ -2,7 +2,7 @@ from pydantic import BaseModel, Field class CreateAssetRequest(BaseModel): - project_id: str = Field(..., min_length=1) + project_id: str | None = Field(default=None, description="可选,不传时从 library.project_id 自动推导") library_id: str = Field(..., min_length=1) name: str = Field(..., min_length=1, max_length=100) storage_key: str = Field(..., min_length=1, max_length=255) diff --git a/tests/unit/test_create_asset_optional_project_id.py b/tests/unit/test_create_asset_optional_project_id.py new file mode 100644 index 000000000..22ec7614b --- /dev/null +++ b/tests/unit/test_create_asset_optional_project_id.py @@ -0,0 +1,190 @@ +"""测试 create_asset 端点:project_id 可选,从 library 自动推导。""" + +from unittest.mock import MagicMock, patch + +import pytest +from app.api.routes.assets import create_asset +from app.auth import AuthenticatedUser +from app.schemas.asset import CreateAssetRequest +from fastapi import HTTPException + +from packages.domain import AssetStatus, ClassificationStatus + + +@pytest.fixture +def mock_user(): + user = MagicMock(spec=AuthenticatedUser) + user.user.id = "user-123" + return user + + +@pytest.fixture +def mock_library(): + lib = MagicMock() + lib.id = "lib-abc" + lib.project_id = "proj-from-library" + return lib + + +@pytest.fixture +def mock_project(): + proj = MagicMock() + proj.id = "proj-from-library" + proj.can_access.return_value = True + return proj + + +def _make_request(**overrides): + defaults = dict( + library_id="lib-abc", + name="test-audio.mp3", + storage_key="uploads/test.mp3", + mime_type="audio/mpeg", + file_size=1024, + status="uploading", + ) + defaults.update(overrides) + return CreateAssetRequest(**defaults) + + +def test_project_id_derived_from_library_when_not_provided(mock_user, mock_library, mock_project): + """前端不传 project_id 时,从 library.project_id 自动推导。""" + request = _make_request() # project_id 默认 None + + asset_repo = MagicMock() + lib_repo = MagicMock() + lib_repo.get.return_value = mock_library + proj_repo = MagicMock() + proj_repo.find_by_id.return_value = mock_project + + expected_asset = MagicMock() + expected_asset.id = "asset-1" + expected_asset.project_id = "proj-from-library" + expected_asset.library_id = "lib-abc" + expected_asset.name = "test-audio.mp3" + expected_asset.storage_key = "" + expected_asset.mime_type = "audio/mpeg" + expected_asset.metadata = {} + expected_asset.file_size = 1024 + expected_asset.thumbnail_url = None + expected_asset.duration = None + expected_asset.width = None + expected_asset.height = None + expected_asset.fps = None + expected_asset.codec = None + expected_asset.status = AssetStatus.UPLOADING + expected_asset.classification_status = ClassificationStatus.PENDING + expected_asset.quality_score = None + expected_asset.created_at = None + expected_asset.uploaded_by_user_id = "user-123" + expected_asset.tag_ids = [] + with patch("app.api.routes.assets.CreateAssetUseCase") as mock_uc: + mock_uc.return_value.execute.return_value = expected_asset + result = create_asset( + request=request, + authenticated_user=mock_user, + asset_repository=asset_repo, + asset_library_repository=lib_repo, + project_repository=proj_repo, + ) + + # 验证 project_id 被正确推导 + proj_repo.find_by_id.assert_called_once_with("proj-from-library") + # 验证 use case 使用的是推导出的 project_id + cmd = mock_uc.return_value.execute.call_args[0][0] + assert cmd.project_id == "proj-from-library" + + +def test_explicit_project_id_used_when_provided(mock_user, mock_library, mock_project): + """前端显式传 project_id 时,优先使用请求值。""" + mock_project.id = "proj-explicit" + mock_project.can_access.return_value = True + mock_library.project_id = "proj-explicit" # 匹配 + + request = _make_request(project_id="proj-explicit") + + asset_repo = MagicMock() + lib_repo = MagicMock() + lib_repo.get.return_value = mock_library + proj_repo = MagicMock() + proj_repo.find_by_id.return_value = mock_project + + mock_asset = MagicMock() + mock_asset.id = "asset-1" + mock_asset.storage_key = "" + mock_asset.mime_type = "audio/mpeg" + mock_asset.project_id = "proj-explicit" + mock_asset.library_id = "lib-abc" + mock_asset.name = "test" + mock_asset.metadata = {} + mock_asset.file_size = 0 + mock_asset.thumbnail_url = None + mock_asset.duration = None + mock_asset.width = None + mock_asset.height = None + mock_asset.fps = None + mock_asset.codec = None + mock_asset.status = AssetStatus.UPLOADING + mock_asset.classification_status = ClassificationStatus.PENDING + mock_asset.quality_score = None + mock_asset.created_at = None + mock_asset.uploaded_by_user_id = "user-123" + mock_asset.tag_ids = [] + + with patch("app.api.routes.assets.CreateAssetUseCase") as mock_uc: + mock_uc.return_value.execute.return_value = mock_asset + create_asset( + request=request, + authenticated_user=mock_user, + asset_repository=asset_repo, + asset_library_repository=lib_repo, + project_repository=proj_repo, + ) + + proj_repo.find_by_id.assert_called_once_with("proj-explicit") + cmd = mock_uc.return_value.execute.call_args[0][0] + assert cmd.project_id == "proj-explicit" + + +def test_library_not_found_returns_404(mock_user): + """素材库不存在时返回 404。""" + request = _make_request() + + lib_repo = MagicMock() + lib_repo.get.return_value = None + proj_repo = MagicMock() + asset_repo = MagicMock() + + with pytest.raises(HTTPException) as exc_info: + create_asset( + request=request, + authenticated_user=mock_user, + asset_repository=asset_repo, + asset_library_repository=lib_repo, + project_repository=proj_repo, + ) + assert exc_info.value.status_code == 404 + + +def test_library_project_mismatch_returns_400(mock_user, mock_library, mock_project): + """当 library.project_id 与请求的 project_id 不一致时返回 400。""" + mock_library.project_id = "proj-A" + mock_project.id = "proj-B" + + request = _make_request(project_id="proj-B") + + lib_repo = MagicMock() + lib_repo.get.return_value = mock_library + proj_repo = MagicMock() + proj_repo.find_by_id.return_value = mock_project + asset_repo = MagicMock() + + with pytest.raises(HTTPException) as exc_info: + create_asset( + request=request, + authenticated_user=mock_user, + asset_repository=asset_repo, + asset_library_repository=lib_repo, + project_repository=proj_repo, + ) + assert exc_info.value.status_code == 400