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