from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import ProjectModel from packages.domain import Project class SQLAlchemyProjectRepository: def __init__(self, session: Session): self.session = session def _to_entity(self, model: ProjectModel) -> Project: return Project( id=model.id, owner_user_id=model.owner_user_id, 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, ) def save(self, project: Project) -> Project: """保存项目(创建或更新)""" existing = self.session.query(ProjectModel).filter(ProjectModel.id == project.id).first() if existing: existing.owner_user_id = project.owner_user_id existing.name = project.name existing.description = project.description existing.shared_users = project.shared_users else: 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=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 def find_by_id(self, project_id: str) -> Project | None: model = self.session.query(ProjectModel).filter(ProjectModel.id == project_id).first() if model is None: return None return self._to_entity(model) def find_by_owner_user_id(self, owner_user_id: str) -> list[Project]: """根据所有者用户 ID 查找项目""" models = self.session.query(ProjectModel).filter(ProjectModel.owner_user_id == owner_user_id).all() return [self._to_entity(model) for model in models] def find_accessible_projects(self, user_id: str) -> list[Project]: """查找用户可访问的所有项目(自己拥有的 + 被共享的)""" from sqlalchemy import cast, or_ from sqlalchemy.dialects.postgresql import JSONB models = ( self.session.query(ProjectModel) .filter( or_(ProjectModel.owner_user_id == user_id, cast(ProjectModel.shared_users, JSONB).contains([user_id])) ) .all() ) return [self._to_entity(model) for model in models] def count_by_owner(self, owner_user_id: str) -> int: """统计用户的项目数量""" return self.session.query(ProjectModel).filter(ProjectModel.owner_user_id == owner_user_id).count() def delete(self, project_id: str) -> bool: """删除项目""" model = self.session.query(ProjectModel).filter(ProjectModel.id == project_id).first() if model is None: return False 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