Files
quantum_backend/src/crud/team_crud.py

283 lines
7.5 KiB
Python

from typing import List, Optional
from sql_models.models import Permission, Team, TeamMember
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlalchemy.sql.expression import func
from starlette.exceptions import HTTPException
async def create_team(
db: AsyncSession,
name: str,
creator_id: str,
description: Optional[str] = None,
) -> Team:
team = Team(
name=name,
description=description,
creator_id=creator_id,
)
db.add(team)
await db.flush() # Get team.id
# Get all available permissions
all_permissions_result = await db.execute(select(Permission))
all_permissions = all_permissions_result.scalars().all()
# Add creator as a team member with all permissions
team_member = TeamMember(
team_id=team.id,
user_id=creator_id,
permissions=all_permissions,
)
db.add(team_member)
await db.commit()
await db.refresh(team, attribute_names=["team_memberships", "creator"])
return team
async def update_team(
db: AsyncSession,
team_id: int,
name: str,
description: Optional[str] = None,
) -> Team:
# Get the existing team
team_result = await db.execute(select(Team).where(Team.id == team_id))
team = team_result.scalar_one_or_none()
if not team:
raise HTTPException(status_code=404, detail="Team not found")
# Update fields
team.name = name
team.description = description
# Add updated_at if you have such a field
# team.updated_at = datetime.utcnow()
db.add(team)
await db.commit()
await db.refresh(team, attribute_names=["team_memberships", "creator"])
return team
async def get_team(db: AsyncSession, team_id: int) -> Optional[Team]:
"""Get team by ID with memberships loaded"""
result = await db.execute(
select(Team)
.where(Team.id == team_id)
.options(
selectinload(Team.team_memberships).selectinload(TeamMember.permissions)
)
)
return result.scalar_one_or_none()
async def get_user_teams(
db: AsyncSession, user_id: str, offset: int, limit: int
) -> tuple[List[Team], int | None]:
"""Get paginated teams a user belongs to and return total count"""
# Get paginated teams
result = await db.execute(
select(Team)
.join(TeamMember)
.where(TeamMember.user_id == user_id)
.options(selectinload(Team.team_memberships))
.order_by(Team.created_at.desc())
.offset(offset)
.limit(limit)
)
teams = list(result.scalars().all())
# Get total count of teams for this user
count_result = await db.execute(
select(func.count())
.select_from(Team)
.join(TeamMember)
.where(TeamMember.user_id == user_id)
)
total_count = count_result.scalar()
return teams, total_count
async def add_team_member(
db: AsyncSession,
team_id: int,
user_id: str,
permissions: List[str] = [],
) -> TeamMember:
"""Add or replace a team member with specified permissions (PUT semantics)"""
# Check if team member already exists
result = await db.execute(
select(TeamMember).where(
TeamMember.team_id == team_id, TeamMember.user_id == user_id
)
)
team_member = result.scalar_one_or_none()
# Get permission objects
permission_objs = []
if permissions:
result = await db.execute(
select(Permission).where(Permission.name.in_(permissions))
)
permission_objs = result.scalars().all()
if team_member:
# Replace existing: update permissions
team_member.permissions = list(permission_objs)
# If you have other fields to update, add them here
else:
# Create new team member
team_member = TeamMember(
team_id=team_id,
user_id=user_id,
permissions=permission_objs,
)
db.add(team_member)
await db.commit()
await db.refresh(team_member)
return team_member
async def get_team_members(db: AsyncSession, team_id: int) -> List[TeamMember]:
"""Get all members of a team with their permissions"""
result = await db.execute(
select(TeamMember)
.where(TeamMember.team_id == team_id)
.options(selectinload(TeamMember.permissions))
)
return list(result.scalars().all())
async def delete_team_member(db: AsyncSession, team_id: int, user_id: str) -> None:
"""Remove a user from a team"""
result = await db.execute(
select(TeamMember).where(
TeamMember.team_id == team_id, TeamMember.user_id == user_id
)
)
team_member = result.scalar_one_or_none()
if team_member:
await db.delete(team_member)
await db.commit()
return 1
return 0
async def update_team_member_permissions(
db: AsyncSession,
team_id: int,
user_id: str,
permissions: List[str],
) -> None:
"""Update a member's permissions"""
result = await db.execute(
select(TeamMember)
.where(TeamMember.team_id == team_id, TeamMember.user_id == user_id)
.options(selectinload(TeamMember.permissions))
)
team_member = result.scalar_one_or_none()
if team_member:
# Get new permission objects
perm_result = await db.execute(
select(Permission).where(Permission.name.in_(permissions))
)
new_permissions = perm_result.scalars().all()
# Update permissions
team_member.permissions = list(new_permissions)
await db.commit()
async def delete_team(
db: AsyncSession,
team_id: int,
) -> bool:
"""
Delete a team by ID.
Returns True if deleted, False if team not found.
"""
# Get the team
result = await db.execute(select(Team).where(Team.id == team_id))
team = result.scalar_one_or_none()
if not team:
return False
# Delete the team (cascade will delete team_memberships automatically)
await db.delete(team)
await db.commit()
return True
async def check_team_permission(
required_permissions: List[str],
team_id: int,
db: AsyncSession,
current_user_id,
) -> Team:
"""
Check if current user has required permissions for a team.
Usage:
@router.delete("/teams/{team_id}")
async def delete_team(
team: Team = Depends(check_team_permission(["delete_team"]))
):
# Permission already checked
pass
"""
# Fetch team with members and permissions
result = await db.execute(
select(Team)
.where(Team.id == team_id)
.options(
selectinload(Team.team_memberships).selectinload(TeamMember.permissions)
)
)
team = result.scalar_one_or_none()
if not team:
raise HTTPException(status_code=404, detail="Team not found")
# Team creator has all permissions
if team.creator_id == current_user_id:
return team
# Find user's membership
user_membership = None
for member in team.team_memberships:
if member.user_id == current_user_id:
user_membership = member
break
if not user_membership:
raise HTTPException(status_code=403, detail="You are not a member of this team")
# Check required permissions
user_permissions = {p.name for p in user_membership.permissions}
for perm in required_permissions:
if perm not in user_permissions:
raise HTTPException(
status_code=403, detail=f"Missing required permission: {perm}"
)
return team