diff --git a/apps/api/app/api/routes/tts.py b/apps/api/app/api/routes/tts.py index bfebc6ce3..78b8f8988 100755 --- a/apps/api/app/api/routes/tts.py +++ b/apps/api/app/api/routes/tts.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging -from typing import Optional +from typing import Any, Optional from app.auth import AuthenticatedUser, get_current_user from app.core.celery_app import celery_app @@ -309,7 +309,7 @@ def _find_or_create_voice_library( *, user_id: str, project_repository: ProjectRepository, - asset_library_repository: AssetLibraryRepository, + asset_library_repository: Any, # port Protocol 声明为 async,SQLAlchemy 实现为同步,与 upload/asset_libraries 路由惯例一致用 Any ) -> AssetLibrary: """在用户可访问的项目中找到(或自动创建)voice 素材库。 @@ -341,12 +341,9 @@ def _find_or_create_voice_library( ) try: return asset_library_repository.create(library) - except Exception as e: - if not isinstance(e, IntegrityError): - raise - session = getattr(asset_library_repository, "session", None) - if session is not None: - session.rollback() + except IntegrityError: + # 并发下另一个请求已抢先创建:SQLAlchemy commit 失败后 session 会自动回滚, + # 直接重查返回已存在的库即可(不依赖 repository 的内部 session 实现)。 for lib in asset_library_repository.find_by_project(project.id): kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind if kind == AssetLibraryKind.VOICE.value: @@ -354,7 +351,7 @@ def _find_or_create_voice_library( raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="配音素材库创建失败,请重试", - ) from e + ) @router.post( diff --git a/tests/unit/test_tts_save_to_library_assets.py b/tests/unit/test_tts_save_to_library_assets.py index f2dd51b69..453a14d98 100644 --- a/tests/unit/test_tts_save_to_library_assets.py +++ b/tests/unit/test_tts_save_to_library_assets.py @@ -88,7 +88,6 @@ class FakeAssetLibraryRepo: def __init__(self, libs=None, fail_integrity=False): self._libs = list(libs or []) self.fail_integrity = fail_integrity - self.session = MagicMock() def find_by_project(self, project_id): return [lib for lib in self._libs if lib.project_id == project_id]