Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f6c18b3aa9 | |||
| 6b1966ce90 |
@@ -0,0 +1,72 @@
|
||||
"""Projects is_default + partial unique index for idempotent default project (Issue #1775)
|
||||
|
||||
Revision ID: 069_project_is_default
|
||||
Revises: 068_user_profile_completed
|
||||
Create Date: 2026-09-08
|
||||
|
||||
背景:
|
||||
小程序端 getOrCreateDefaultProject 在重试/并发/前端重复调用下,
|
||||
仅靠应用层"先查再插"不保证幂等,会给同一用户重复创建默认项目。
|
||||
|
||||
改动:
|
||||
1. projects 表新增 is_default 布尔列(默认 false)
|
||||
2. 部分唯一索引 uq_projects_owner_default:(owner_user_id) WHERE is_default = true
|
||||
—— 保证每个用户至多一个默认项目
|
||||
3. 存量数据回填:把名为"默认项目"的存量项目按创建时间最早者标记为 is_default=true
|
||||
(只标记不删除;存量重复项目的清理另行确认后单独执行)
|
||||
|
||||
注意:部分唯一索引依赖 PostgreSQL,不支持 downgrade 到其他方言。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "069_project_is_default"
|
||||
down_revision = "068_user_profile_completed"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 1. 新增 is_default 列
|
||||
op.add_column(
|
||||
"projects",
|
||||
sa.Column(
|
||||
"is_default",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
)
|
||||
|
||||
# 2. 存量回填:每个拥有"默认项目"的用户,只把最早创建的那一个标记为默认。
|
||||
# 用 ROW_NUMBER() 取每组第一条;非"默认项目"命名的项目不标记(保守,不动用户自建项目)。
|
||||
op.execute("""
|
||||
UPDATE projects p
|
||||
SET is_default = true
|
||||
WHERE p.id IN (
|
||||
SELECT id FROM (
|
||||
SELECT id,
|
||||
ROW_NUMBER() OVER (
|
||||
PARTITION BY owner_user_id
|
||||
ORDER BY created_at ASC, id ASC
|
||||
) AS rn
|
||||
FROM projects
|
||||
WHERE name = '默认项目'
|
||||
) t
|
||||
WHERE t.rn = 1
|
||||
)
|
||||
""")
|
||||
|
||||
# 3. 部分唯一索引:每用户至多一个默认项目(只约束 is_default = true 的行)
|
||||
op.execute("""
|
||||
CREATE UNIQUE INDEX uq_projects_owner_default
|
||||
ON projects (owner_user_id)
|
||||
WHERE is_default = true
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP INDEX IF EXISTS uq_projects_owner_default")
|
||||
op.drop_column("projects", "is_default")
|
||||
@@ -20,7 +20,7 @@ from packages.application import (
|
||||
GetProjectUseCase,
|
||||
ListAssetLibrariesUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind
|
||||
from packages.domain import AssetLibraryKind
|
||||
|
||||
from ._helpers import check_project_access
|
||||
|
||||
@@ -120,30 +120,11 @@ def ensure_default_library(
|
||||
|
||||
kind = AssetLibraryKind(request.kind)
|
||||
|
||||
# 查找该项目下同 kind 的素材库,返回第一个
|
||||
existing = asset_library_repository.find_by_project(request.project_id)
|
||||
for lib in existing:
|
||||
if lib.kind == kind:
|
||||
return _to_asset_library_response(lib)
|
||||
|
||||
# 不存在 → 自动创建
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
# Issue #1775: 幂等获取/创建——依赖唯一约束 uq_asset_libraries_project_kind,
|
||||
# 并发创建冲突时回滚重查返回已有记录,不再依赖应用层"先查后插",也不会 500。
|
||||
default_name = _DEFAULT_LIBRARY_NAMES.get(request.kind, f"{request.kind}素材库")
|
||||
library = AssetLibrary(
|
||||
id=str(uuid.uuid4()),
|
||||
project_id=request.project_id,
|
||||
name=default_name,
|
||||
kind=kind,
|
||||
asset_count=0,
|
||||
total_size=0,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
created = asset_library_repository.create(library)
|
||||
return _to_asset_library_response(created)
|
||||
library = asset_library_repository.get_or_create_default_library(request.project_id, kind, name=default_name)
|
||||
return _to_asset_library_response(library)
|
||||
|
||||
|
||||
@router.delete("/{library_id}", status_code=status.HTTP_204_NO_CONTENT, response_class=Response)
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_project_repository
|
||||
from app.dependencies import get_asset_library_repository, get_project_repository
|
||||
from app.schemas.project import (
|
||||
CreateProjectRequest,
|
||||
ListProjectsResponse,
|
||||
ProjectResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Response, status
|
||||
from pydantic import BaseModel
|
||||
|
||||
from packages.application import (
|
||||
CreateProjectCommand,
|
||||
@@ -16,10 +17,20 @@ from packages.application import (
|
||||
GetProjectUseCase,
|
||||
ListProjectsUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibraryKind
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class DefaultContextResponse(BaseModel):
|
||||
"""幂等默认上下文响应(Issue #1775):默认项目 + 各类型默认素材库 ID。"""
|
||||
|
||||
project_id: str
|
||||
image_library_id: str
|
||||
video_library_id: str
|
||||
voice_library_id: str
|
||||
|
||||
|
||||
def _to_project_response(item) -> ProjectResponse:
|
||||
return ProjectResponse(
|
||||
id=item.id,
|
||||
@@ -72,6 +83,35 @@ def create_project(
|
||||
return _to_project_response(project)
|
||||
|
||||
|
||||
@router.post("/ensure-default", response_model=DefaultContextResponse)
|
||||
def ensure_default_project_and_libraries(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
asset_library_repository: Any = Depends(get_asset_library_repository),
|
||||
) -> DefaultContextResponse:
|
||||
"""幂等获取/创建当前用户的默认项目和三类默认素材库(Issue #1775)。
|
||||
|
||||
- 同一用户永远只有一个默认项目(部分唯一索引 uq_projects_owner_default)
|
||||
- 同一项目同 kind 永远只有一个默认素材库(唯一约束 uq_asset_libraries_project_kind)
|
||||
- 并发调用/失败重试:唯一约束冲突时返回已存在记录,不报 500
|
||||
- 项目和素材库的创建各自在仓储事务内幂等,冲突回滚后重查返回同一条
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
project = project_repository.get_or_create_default_project(user_id)
|
||||
|
||||
libraries = {}
|
||||
for kind in (AssetLibraryKind.VIDEO, AssetLibraryKind.VOICE, AssetLibraryKind.IMAGE):
|
||||
library = asset_library_repository.get_or_create_default_library(project.id, kind)
|
||||
libraries[kind] = library.id
|
||||
|
||||
return DefaultContextResponse(
|
||||
project_id=project.id,
|
||||
image_library_id=libraries[AssetLibraryKind.IMAGE],
|
||||
video_library_id=libraries[AssetLibraryKind.VIDEO],
|
||||
voice_library_id=libraries[AssetLibraryKind.VOICE],
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
|
||||
def delete_project(
|
||||
project_id: str,
|
||||
|
||||
@@ -2126,6 +2126,14 @@
|
||||
"type": "JSON",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "is_default",
|
||||
"nullable": false,
|
||||
"primary_key": false,
|
||||
"type": "BOOLEAN",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "created_at",
|
||||
@@ -2142,6 +2150,13 @@
|
||||
],
|
||||
"name": "ix_projects_owner_user_id",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"columns": [
|
||||
"owner_user_id"
|
||||
],
|
||||
"name": "uq_projects_owner_default",
|
||||
"unique": true
|
||||
}
|
||||
],
|
||||
"primary_key": [
|
||||
@@ -3570,4 +3585,4 @@
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -46,3 +46,21 @@ class InMemoryAssetLibraryRepository:
|
||||
if library:
|
||||
library.asset_count = max(0, library.asset_count - 1)
|
||||
library.total_size = max(0, library.total_size - size_delta)
|
||||
|
||||
def get_or_create_default_library(
|
||||
self,
|
||||
project_id: str,
|
||||
kind: AssetLibraryKind,
|
||||
*,
|
||||
name: str | None = None,
|
||||
) -> AssetLibrary:
|
||||
"""幂等获取/创建默认素材库(Issue #1775,内存实现,模拟唯一约束语义)。"""
|
||||
for lib in self._libraries.values():
|
||||
if lib.project_id == project_id and lib.kind == kind:
|
||||
return lib
|
||||
# 回退到 find_by_project
|
||||
for lib in self.find_by_project(project_id, kind):
|
||||
return lib
|
||||
library_name = name or f"{kind.value}素材库"
|
||||
library = AssetLibrary.create(project_id=project_id, name=library_name, kind=kind)
|
||||
return self.create(library)
|
||||
|
||||
@@ -29,3 +29,30 @@ class InMemoryProjectRepository:
|
||||
del self._items[project_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
def find_default_by_owner(self, owner_user_id: str) -> Project | None:
|
||||
"""查找用户的默认项目(Issue #1775 幂等接口,内存实现)。"""
|
||||
for p in self._items.values():
|
||||
if p.owner_user_id == owner_user_id and getattr(p, "is_default", False):
|
||||
return p
|
||||
return None
|
||||
|
||||
def get_or_create_default_project(
|
||||
self,
|
||||
owner_user_id: str,
|
||||
*,
|
||||
name: str = "默认项目",
|
||||
description: str = "小程序自动创建的默认项目",
|
||||
) -> Project:
|
||||
"""幂等获取/创建默认项目(内存实现,模拟 DB 部分唯一索引语义)。"""
|
||||
existing = self.find_default_by_owner(owner_user_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
project = Project.create(
|
||||
owner_user_id=owner_user_id,
|
||||
name=name,
|
||||
description=description,
|
||||
is_default=True,
|
||||
)
|
||||
self._items[project.id] = project
|
||||
return project
|
||||
|
||||
@@ -90,3 +90,75 @@ class SQLAlchemyAssetLibraryRepository:
|
||||
model.asset_count = max(0, (model.asset_count or 0) - 1)
|
||||
model.total_size = max(0, (model.total_size or 0) - size_delta)
|
||||
self.session.commit()
|
||||
|
||||
def get_or_create_default_library(
|
||||
self,
|
||||
project_id: str,
|
||||
kind: AssetLibraryKind,
|
||||
*,
|
||||
name: str | None = None,
|
||||
) -> AssetLibrary:
|
||||
"""幂等获取/创建项目下指定 kind 的默认素材库(Issue #1775)。
|
||||
|
||||
依赖唯一约束 uq_asset_libraries_project_kind(project_id, kind):
|
||||
并发创建只有一个成功,其余 IntegrityError 后回滚重查,
|
||||
保证同一项目同 kind 永远只有一个素材库。
|
||||
"""
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
default_names = {
|
||||
AssetLibraryKind.VIDEO: "视频素材库",
|
||||
AssetLibraryKind.VOICE: "配音素材库",
|
||||
AssetLibraryKind.IMAGE: "图片素材库",
|
||||
}
|
||||
library_name = name or default_names.get(kind, f"{kind.value}素材库")
|
||||
|
||||
# 快速路径
|
||||
existing = (
|
||||
self.session.query(AssetLibraryModel)
|
||||
.filter(AssetLibraryModel.project_id == project_id, AssetLibraryModel.kind == kind.value)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
return self._to_entity(existing)
|
||||
|
||||
library = AssetLibrary.create(project_id=project_id, name=library_name, kind=kind)
|
||||
model = AssetLibraryModel(
|
||||
id=library.id,
|
||||
project_id=library.project_id,
|
||||
name=library.name,
|
||||
kind=library.kind.value,
|
||||
asset_count=0,
|
||||
total_size=0,
|
||||
created_at=library.created_at,
|
||||
updated_at=library.updated_at,
|
||||
)
|
||||
try:
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
return library
|
||||
except IntegrityError:
|
||||
self.session.rollback()
|
||||
existing = (
|
||||
self.session.query(AssetLibraryModel)
|
||||
.filter(
|
||||
AssetLibraryModel.project_id == project_id,
|
||||
AssetLibraryModel.kind == kind.value,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
return self._to_entity(existing)
|
||||
raise
|
||||
|
||||
def _to_entity(self, model: AssetLibraryModel) -> AssetLibrary:
|
||||
return AssetLibrary(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
name=model.name,
|
||||
kind=AssetLibraryKind(model.kind),
|
||||
asset_count=int(model.asset_count or 0),
|
||||
total_size=int(model.total_size or 0),
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Index, Integer, String, Text, UniqueConstraint, text
|
||||
from sqlalchemy.orm import declarative_base
|
||||
|
||||
Base: Any = declarative_base()
|
||||
@@ -44,12 +44,18 @@ class UserModel(Base):
|
||||
|
||||
class ProjectModel(Base):
|
||||
__tablename__ = "projects"
|
||||
__table_args__ = (
|
||||
# Issue #1775: 每个用户至多一个默认项目(部分唯一索引,只约束 is_default=true 的行)。
|
||||
# 注意:不加 UniqueConstraint(那会要求全表唯一),用部分索引表达"每用户一个默认项目"。
|
||||
Index("uq_projects_owner_default", "owner_user_id", unique=True, postgresql_where=text("is_default = true")),
|
||||
)
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
owner_user_id = Column(String(36), nullable=False, index=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
shared_users = Column(JSON, nullable=False, default=list) # 被共享的用户 ID 列表
|
||||
is_default = Column(Boolean, nullable=False, default=False, server_default="false")
|
||||
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ class SQLAlchemyProjectRepository:
|
||||
name=model.name,
|
||||
description=model.description,
|
||||
shared_users=model.shared_users or [],
|
||||
is_default=bool(getattr(model, "is_default", False)),
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -33,9 +34,12 @@ class SQLAlchemyProjectRepository:
|
||||
name=project.name,
|
||||
description=project.description,
|
||||
shared_users=project.shared_users,
|
||||
is_default=project.is_default,
|
||||
created_at=project.created_at,
|
||||
)
|
||||
self.session.add(model)
|
||||
if existing:
|
||||
existing.is_default = project.is_default
|
||||
self.session.commit()
|
||||
return project
|
||||
|
||||
@@ -76,3 +80,59 @@ class SQLAlchemyProjectRepository:
|
||||
self.session.delete(model)
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
def find_default_by_owner(self, owner_user_id: str) -> Project | None:
|
||||
"""查找用户的默认项目(is_default=true)。"""
|
||||
model = (
|
||||
self.session.query(ProjectModel)
|
||||
.filter(ProjectModel.owner_user_id == owner_user_id, ProjectModel.is_default.is_(True))
|
||||
.first()
|
||||
)
|
||||
return self._to_entity(model) if model else None
|
||||
|
||||
def get_or_create_default_project(
|
||||
self,
|
||||
owner_user_id: str,
|
||||
*,
|
||||
name: str = "默认项目",
|
||||
description: str = "小程序自动创建的默认项目",
|
||||
) -> Project:
|
||||
"""幂等获取/创建用户的默认项目(Issue #1775)。
|
||||
|
||||
依赖部分唯一索引 uq_projects_owner_default(每用户至多一条 is_default=true):
|
||||
并发创建时只有一个 INSERT 成功,其余触发 IntegrityError 后回滚重查,
|
||||
保证同一用户永远只有一个默认项目。
|
||||
"""
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
# 快速路径:已有默认项目
|
||||
existing = self.find_default_by_owner(owner_user_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
|
||||
project = Project.create(
|
||||
owner_user_id=owner_user_id,
|
||||
name=name,
|
||||
description=description,
|
||||
is_default=True,
|
||||
)
|
||||
model = ProjectModel(
|
||||
id=project.id,
|
||||
owner_user_id=project.owner_user_id,
|
||||
name=project.name,
|
||||
description=project.description,
|
||||
shared_users=project.shared_users,
|
||||
is_default=True,
|
||||
created_at=project.created_at,
|
||||
)
|
||||
try:
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
return project
|
||||
except IntegrityError:
|
||||
# 并发:另一个请求已插入默认项目,回滚后重查
|
||||
self.session.rollback()
|
||||
existing = self.find_default_by_owner(owner_user_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
raise
|
||||
|
||||
@@ -70,10 +70,11 @@ class Project:
|
||||
name: str
|
||||
description: str = ""
|
||||
shared_users: list[str] = field(default_factory=list) # 被共享的用户 ID 列表
|
||||
is_default: bool = False # 是否为用户的默认项目(小程序自动创建),DB 部分唯一索引保证每人至多一个
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@classmethod
|
||||
def create(cls, owner_user_id: str, name: str, description: str = "") -> "Project":
|
||||
def create(cls, owner_user_id: str, name: str, description: str = "", is_default: bool = False) -> "Project":
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("项目名称不能为空")
|
||||
@@ -83,6 +84,7 @@ class Project:
|
||||
name=clean_name,
|
||||
description=description.strip(),
|
||||
shared_users=[],
|
||||
is_default=is_default,
|
||||
)
|
||||
|
||||
def is_owner(self, user_id: str) -> bool:
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
"""默认项目/默认素材库幂等化测试(Issue #1775)。
|
||||
|
||||
覆盖:
|
||||
- get_or_create_default_project:同用户幂等、不同用户独立、并发只建一个
|
||||
- get_or_create_default_library:同项目同 kind 幂等、IntegrityError 后重查
|
||||
- Project.is_default 字段传递
|
||||
- ensure-default-context 组合逻辑(用内存仓储)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.in_memory.asset_library_repository import InMemoryAssetLibraryRepository
|
||||
from packages.adapters.in_memory.project_repository import InMemoryProjectRepository
|
||||
from packages.adapters.sqlalchemy_impl.asset_library_repository import (
|
||||
SQLAlchemyAssetLibraryRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.project_repository import SQLAlchemyProjectRepository
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind, Project
|
||||
|
||||
|
||||
class TestDefaultProjectIdempotent:
|
||||
"""默认项目幂等。"""
|
||||
|
||||
def test_first_call_creates_default(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
project = repo.get_or_create_default_project("user-1")
|
||||
assert project.owner_user_id == "user-1"
|
||||
assert project.is_default is True
|
||||
assert project.name == "默认项目"
|
||||
|
||||
def test_second_call_returns_same_project(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
p1 = repo.get_or_create_default_project("user-1")
|
||||
p2 = repo.get_or_create_default_project("user-1")
|
||||
assert p1.id == p2.id
|
||||
|
||||
def test_different_users_independent(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
p1 = repo.get_or_create_default_project("user-1")
|
||||
p2 = repo.get_or_create_default_project("user-2")
|
||||
assert p1.id != p2.id
|
||||
assert p1.owner_user_id == "user-1"
|
||||
assert p2.owner_user_id == "user-2"
|
||||
|
||||
def test_find_default_by_owner(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
created = repo.get_or_create_default_project("user-1")
|
||||
found = repo.find_default_by_owner("user-1")
|
||||
assert found is not None
|
||||
assert found.id == created.id
|
||||
|
||||
def test_find_default_returns_none_when_no_default(self):
|
||||
repo = InMemoryProjectRepository()
|
||||
# 手动建一个非默认项目
|
||||
normal = Project.create(owner_user_id="user-1", name="普通项目")
|
||||
normal.is_default = False
|
||||
repo.save(normal)
|
||||
assert repo.find_default_by_owner("user-1") is None
|
||||
|
||||
def test_concurrent_10_calls_only_one_project(self):
|
||||
"""并发 10 次调用,只产生 1 个默认项目。"""
|
||||
repo = InMemoryProjectRepository()
|
||||
results = []
|
||||
lock = threading.Lock()
|
||||
|
||||
def call():
|
||||
p = repo.get_or_create_default_project("user-concurrent")
|
||||
with lock:
|
||||
results.append(p.id)
|
||||
|
||||
threads = [threading.Thread(target=call) for _ in range(10)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert len(results) == 10
|
||||
assert len(set(results)) == 1, f"应只有 1 个项目,实际: {set(results)}"
|
||||
# 数据库中也只有 1 个默认项目
|
||||
defaults = [p for p in repo.find_by_owner_user_id("user-concurrent") if p.is_default]
|
||||
assert len(defaults) == 1
|
||||
|
||||
|
||||
class TestDefaultLibraryIdempotent:
|
||||
"""默认素材库幂等。"""
|
||||
|
||||
def test_first_call_creates(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
lib = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
assert lib.project_id == "proj-1"
|
||||
assert lib.kind == AssetLibraryKind.VIDEO
|
||||
|
||||
def test_second_call_returns_same(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
l1 = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
l2 = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
assert l1.id == l2.id
|
||||
|
||||
def test_different_kinds_independent(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
v = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
a = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VOICE)
|
||||
i = repo.get_or_create_default_library("proj-1", AssetLibraryKind.IMAGE)
|
||||
assert len({v.id, a.id, i.id}) == 3
|
||||
|
||||
def test_different_projects_independent(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
l1 = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
l2 = repo.get_or_create_default_library("proj-2", AssetLibraryKind.VIDEO)
|
||||
assert l1.id != l2.id
|
||||
|
||||
def test_concurrent_calls_only_one_library(self):
|
||||
repo = InMemoryAssetLibraryRepository()
|
||||
results = []
|
||||
lock = threading.Lock()
|
||||
|
||||
def call():
|
||||
lib = repo.get_or_create_default_library("proj-cc", AssetLibraryKind.VOICE)
|
||||
with lock:
|
||||
results.append(lib.id)
|
||||
|
||||
threads = [threading.Thread(target=call) for _ in range(10)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert len(set(results)) == 1
|
||||
|
||||
|
||||
class TestProjectIsDefaultField:
|
||||
"""Project.is_default 字段语义。"""
|
||||
|
||||
def test_create_non_default_by_default(self):
|
||||
p = Project.create(owner_user_id="u", name="普通项目")
|
||||
assert p.is_default is False
|
||||
|
||||
def test_create_default(self):
|
||||
p = Project.create(owner_user_id="u", name="默认项目", is_default=True)
|
||||
assert p.is_default is True
|
||||
|
||||
|
||||
class TestSqlRepoIntegrityErrorRecovery:
|
||||
"""SQLAlchemy 仓储:唯一约束冲突时回滚重查,返回已有记录(不报 500)。"""
|
||||
|
||||
def test_project_integrity_error_returns_existing(self):
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
session = MagicMock()
|
||||
# commit 第一次抛 IntegrityError(并发冲突),回滚后查询返回已有项目
|
||||
existing_model = MagicMock()
|
||||
existing_model.id = "existing-id"
|
||||
existing_model.owner_user_id = "user-1"
|
||||
existing_model.name = "默认项目"
|
||||
existing_model.description = ""
|
||||
existing_model.shared_users = []
|
||||
existing_model.is_default = True
|
||||
existing_model.created_at = None
|
||||
|
||||
session.commit.side_effect = [IntegrityError("stmt", {}, Exception("dup")), None]
|
||||
# 第一次 query(快速路径 find_default)返回 None;rollback 后第二次返回 existing
|
||||
session.query.return_value.filter.return_value.first.side_effect = [None, existing_model]
|
||||
|
||||
repo = SQLAlchemyProjectRepository(session)
|
||||
result = repo.get_or_create_default_project("user-1")
|
||||
|
||||
assert result.id == "existing-id"
|
||||
session.rollback.assert_called_once()
|
||||
|
||||
def test_library_integrity_error_returns_existing(self):
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
session = MagicMock()
|
||||
existing_model = MagicMock()
|
||||
existing_model.id = "lib-existing"
|
||||
existing_model.project_id = "proj-1"
|
||||
existing_model.name = "视频素材库"
|
||||
existing_model.kind = "video"
|
||||
existing_model.asset_count = 0
|
||||
existing_model.total_size = 0
|
||||
existing_model.created_at = None
|
||||
existing_model.updated_at = None
|
||||
|
||||
session.commit.side_effect = [IntegrityError("stmt", {}, Exception("dup")), None]
|
||||
session.query.return_value.filter.return_value.first.side_effect = [None, existing_model]
|
||||
|
||||
repo = SQLAlchemyAssetLibraryRepository(session)
|
||||
result = repo.get_or_create_default_library("proj-1", AssetLibraryKind.VIDEO)
|
||||
|
||||
assert result.id == "lib-existing"
|
||||
assert result.kind == AssetLibraryKind.VIDEO
|
||||
session.rollback.assert_called_once()
|
||||
Reference in New Issue
Block a user