283 lines
7.5 KiB
Python
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
|