diff --git a/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index 687c0db8f..58269c35c 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -23,6 +23,7 @@ from app.schemas.upload import ( from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase +from packages.domain import Asset, AssetStatus logger = logging.getLogger(__name__) @@ -80,6 +81,40 @@ def _validate_mime_type(content_type: str | None) -> str: return base_type +def _infer_mime_type_from_storage_key(storage_key: str) -> str: + """从 storage_key 推断 MIME 类型(与 worker 端保持一致)。""" + lower_filename = storage_key.rsplit("/", 1)[-1].lower() + _MIME_MAP = { + ".mov": "video/quicktime", ".mp4": "video/mp4", ".avi": "video/x-msvideo", + ".mkv": "video/x-matroska", ".webm": "video/webm", + ".png": "image/png", ".gif": "image/gif", ".bmp": "image/bmp", + ".svg": "image/svg+xml", ".jpg": "image/jpeg", ".jpeg": "image/jpeg", + ".mp3": "audio/mpeg", ".wav": "audio/wav", ".ogg": "audio/ogg", + ".flac": "audio/flac", ".m4a": "audio/x-m4a", + } + for ext, mime in _MIME_MAP.items(): + if lower_filename.endswith(ext): + return mime + return "video/mp4" # default + + +def _create_pending_asset( + asset_repository, project_id, library_id, storage_key, filename, mime_type, user_id, file_hash="" +): + """立即创建一条 PROCESSING 状态的 Asset 记录,使前端能马上看到新素材。""" + asset = Asset.create( + project_id=project_id, + library_id=library_id, + name=filename, + storage_key=storage_key, + mime_type=mime_type, + status=AssetStatus.PROCESSING, + uploaded_by_user_id=user_id, + file_hash=file_hash, + ) + return asset_repository.create(asset) + + def _submit_ingest_job( project_id: str, library_id: str, @@ -209,6 +244,20 @@ async def complete_direct_upload( url=storage_service.get_url(normalized_key), ) + # 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材 + filename = normalized_key.rsplit("/", 1)[-1] + mime_type = _infer_mime_type_from_storage_key(normalized_key) + pending_asset = _create_pending_asset( + asset_repository=asset_repository, + project_id=request.project_id, + library_id=request.library_id, + storage_key=normalized_key, + filename=filename, + mime_type=mime_type, + user_id=authenticated_user.user.id, + file_hash=request.file_hash, + ) + job = _submit_ingest_job( project_id=request.project_id, library_id=request.library_id, @@ -216,7 +265,12 @@ async def complete_direct_upload( ingest_job_repository=ingest_job_repository, file_hash=request.file_hash, ) - return DirectUploadCompleteResponse(storage_key=normalized_key, ingest_job_id=job.id, url=storage_service.get_url(normalized_key)) + return DirectUploadCompleteResponse( + storage_key=normalized_key, + ingest_job_id=job.id, + asset_id=pending_asset.id, + url=storage_service.get_url(normalized_key), + ) @router.post( @@ -284,6 +338,18 @@ async def upload_asset( detail=f"Failed to upload file: {type(error).__name__}", ) from error + # 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材 + pending_asset = _create_pending_asset( + asset_repository=asset_repository, + project_id=project_id, + library_id=library_id, + storage_key=storage_key, + filename=safe_filename, + mime_type=validated_content_type, + user_id=authenticated_user.user.id, + file_hash=file_hash, + ) + job = _submit_ingest_job( project_id=project_id, library_id=library_id, @@ -295,5 +361,6 @@ async def upload_asset( return UploadAssetResponse( storage_key=storage_key, ingest_job_id=job.id, + asset_id=pending_asset.id, url=file_url, ) diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index c9ac3f79e..7a9f362de 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -656,24 +656,55 @@ def ingest_asset(job_id: str) -> dict: "error": error_reason, } - # Create Asset - asset = Asset.create( - project_id=job.project_id, - library_id=job.library_id, - name=filename, - storage_key=job.storage_key, - mime_type=mime_type, - metadata=metadata, - file_size=int(metadata.get("size_bytes", 0)), - duration=float(metadata.get("duration", 0)), - width=int(metadata.get("width", 0)), - height=int(metadata.get("height", 0)), - codec=metadata.get("codec") or None, - status=AssetStatus.READY, - file_hash=job.file_hash, - thumbnail_url=thumbnail_url, - ) - asset_repo.create(asset) + # 查找已存在的 Asset 记录(由 API 端在上传完成时立即创建为 PROCESSING 状态) + existing_asset = None + try: + existing_asset = asset_repo.find_by_storage_key(job.storage_key) + except Exception: + logger.warning("find_by_storage_key not available, trying fallback lookup") + + if existing_asset is None: + # 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录 + logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key) + asset = Asset.create( + project_id=job.project_id, + library_id=job.library_id, + name=filename, + storage_key=job.storage_key, + mime_type=mime_type, + metadata=metadata, + file_size=int(metadata.get("size_bytes", 0)), + duration=float(metadata.get("duration", 0)), + width=int(metadata.get("width", 0)), + height=int(metadata.get("height", 0)), + codec=metadata.get("codec") or None, + status=AssetStatus.READY, + file_hash=job.file_hash, + thumbnail_url=thumbnail_url, + ) + asset_repo.create(asset) + else: + # 更新已有的 Asset 记录,补充元数据并将状态改为 READY + asset = existing_asset + asset.mime_type = mime_type + asset.metadata = metadata + asset.file_size = int(metadata.get("size_bytes", 0)) + asset.duration = float(metadata.get("duration", 0)) + asset.width = int(metadata.get("width", 0)) + asset.height = int(metadata.get("height", 0)) + codec_val = metadata.get("codec") + if codec_val: + asset.codec = str(codec_val) + fps_val = metadata.get("fps") + if fps_val: + try: + asset.fps = float(fps_val) + except (ValueError, TypeError): + pass + asset.status = AssetStatus.READY + asset.thumbnail_url = thumbnail_url + asset.updated_at = datetime.now(timezone.utc) + asset_repo.update(asset) # Update job status to COMPLETED job.status = IngestJobStatus.COMPLETED diff --git a/packages/adapters/in_memory/asset_repository.py b/packages/adapters/in_memory/asset_repository.py index c3fc13e5c..cc5bd76b2 100755 --- a/packages/adapters/in_memory/asset_repository.py +++ b/packages/adapters/in_memory/asset_repository.py @@ -127,6 +127,13 @@ class InMemoryAssetRepository: items = [a for a in self._assets.values() if tag_set.issubset(set(a.tag_ids))] return items[skip : skip + limit] + def find_by_storage_key(self, storage_key: str) -> Asset | None: + """按 storage_key 查找素材。""" + for asset in self._assets.values(): + if asset.storage_key == storage_key: + return asset + return None + def find_by_library_and_file_hash( self, library_id: str, diff --git a/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py index b0c5d03b9..146ab706d 100755 --- a/packages/adapters/sqlalchemy_impl/asset_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_repository.py @@ -426,6 +426,13 @@ class SQLAlchemyAssetRepository: models = self.session.query(AssetModel).filter(AssetModel.id.in_(ids)).offset(skip).limit(limit).all() return [self._to_domain(m) for m in models] + def find_by_storage_key(self, storage_key: str) -> Asset | None: + """按 storage_key(对应 DB 中的 file_url)查找素材。""" + model = self.session.query(AssetModel).filter(AssetModel.file_url == storage_key).first() + if model is None: + return None + return self._to_domain(model) + def find_by_library_and_file_hash( self, library_id: str, diff --git a/packages/ports/asset_repository.py b/packages/ports/asset_repository.py index 5122968d9..9a9c830ad 100755 --- a/packages/ports/asset_repository.py +++ b/packages/ports/asset_repository.py @@ -112,6 +112,11 @@ class AssetRepository(ABC): """查找包含所有指定标签的素材。""" pass + @abstractmethod + def find_by_storage_key(self, storage_key: str) -> Asset | None: + """按 storage_key 查找素材(用于异步处理时更新已创建的记录)。""" + pass + @abstractmethod def find_by_library_and_file_hash( self, diff --git a/tests/unit/test_form_upload_routes.py b/tests/unit/test_form_upload_routes.py index cc4c2239f..8b686c877 100644 --- a/tests/unit/test_form_upload_routes.py +++ b/tests/unit/test_form_upload_routes.py @@ -104,11 +104,31 @@ def _make_library( return AssetLibrary(id=id, name="Test Library", project_id=project_id, kind=kind) +class StubAssetRepository: + """Minimal asset repository stub for upload tests.""" + def __init__(self): + self._assets = {} + + def create(self, asset): + self._assets[asset.id] = asset + return asset + + def find_by_storage_key(self, storage_key): + for a in self._assets.values(): + if a.storage_key == storage_key: + return a + return None + + def find_by_library_and_file_hash(self, library_id, file_hash): + return None + + def _build_app( project_repo: StubProjectRepository | None = None, library_repo: StubAssetLibraryRepository | None = None, storage: MagicMock | None = None, ingest_repo: StubIngestJobRepository | None = None, + asset_repo: StubAssetRepository | None = None, ) -> FastAPI: """构建一个最小化的 FastAPI app,只注册 upload 路由。""" from app.api.routes.upload import router @@ -116,6 +136,7 @@ def _build_app( from app.core.storage import get_storage_service from app.dependencies import ( get_asset_library_repository, + get_asset_repository, get_ingest_job_repository, get_project_repository, ) @@ -129,17 +150,21 @@ def _build_app( storage.is_configured = True storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/test.mp4" ingest_repo = ingest_repo or StubIngestJobRepository() + asset_repo = asset_repo or StubAssetRepository() # Mock auth mock_user = MagicMock(spec=AuthenticatedUser) mock_user.id = "user-1" mock_user.email = "test@example.com" + mock_user.user = MagicMock() + mock_user.user.id = "user-1" app.dependency_overrides[get_current_user] = lambda: mock_user app.dependency_overrides[get_project_repository] = lambda: project_repo app.dependency_overrides[get_asset_library_repository] = lambda: library_repo app.dependency_overrides[get_storage_service] = lambda: storage app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo + app.dependency_overrides[get_asset_repository] = lambda: asset_repo return app