169 lines
4.6 KiB
Python
169 lines
4.6 KiB
Python
"""User management router endpoints."""
|
|
from fastapi import APIRouter, Depends, HTTPException, status
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlalchemy import select, update
|
|
from typing import Optional
|
|
|
|
from app.core.database import get_db
|
|
from app.core.security import get_current_user, hash_password
|
|
from app.models.models import User
|
|
from app.schemas.schemas import UserUpdate, UserResponse
|
|
|
|
router = APIRouter()
|
|
|
|
|
|
@router.get("/{user_id}", response_model=UserResponse)
|
|
async def get_user(
|
|
user_id: int,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
):
|
|
"""Get user by ID."""
|
|
stmt = select(User).where(User.id == user_id)
|
|
result = await db.execute(stmt)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail="User not found",
|
|
)
|
|
|
|
return UserResponse.model_validate(user)
|
|
|
|
|
|
@router.patch("/{user_id}", response_model=UserResponse)
|
|
async def update_user(
|
|
user_id: int,
|
|
request: UserUpdate,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
):
|
|
"""Update user profile."""
|
|
# Users can only update their own profile unless admin
|
|
if int(current_user["user_id"]) != user_id:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="Cannot update another user's profile",
|
|
)
|
|
|
|
stmt = select(User).where(User.id == user_id)
|
|
result = await db.execute(stmt)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail="User not found",
|
|
)
|
|
|
|
# Update fields
|
|
update_data = request.model_dump(exclude_unset=True)
|
|
|
|
# Hash new password if provided
|
|
if "new_password" in update_data:
|
|
update_data["password_hash"] = hash_password(update_data.pop("new_password"))
|
|
|
|
# Update user
|
|
stmt = (
|
|
update(User)
|
|
.where(User.id == user_id)
|
|
.values(**update_data)
|
|
)
|
|
await db.execute(stmt)
|
|
await db.flush()
|
|
|
|
# Fetch updated user
|
|
stmt = select(User).where(User.id == user_id)
|
|
result = await db.execute(stmt)
|
|
updated_user = result.scalar_one_or_none()
|
|
|
|
return UserResponse.model_validate(updated_user)
|
|
|
|
|
|
@router.post("/{user_id}/ban")
|
|
async def ban_user(
|
|
user_id: int,
|
|
reason: str,
|
|
expiration_time: Optional[str] = None,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
):
|
|
"""Ban a user (admin only)."""
|
|
# Check if current user is admin
|
|
if current_user.get("privlevel") not in ["Admin", "Judge"]:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="Admin privileges required",
|
|
)
|
|
|
|
stmt = select(User).where(User.id == user_id)
|
|
result = await db.execute(stmt)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail="User not found",
|
|
)
|
|
|
|
# Update user ban status
|
|
from datetime import datetime
|
|
ban_ends = None
|
|
if expiration_time:
|
|
ban_ends = datetime.fromisoformat(expiration_time)
|
|
|
|
stmt = (
|
|
update(User)
|
|
.where(User.id == user_id)
|
|
.values(
|
|
is_banned=True,
|
|
ban_reason=reason,
|
|
ban_ends=ban_ends,
|
|
)
|
|
)
|
|
await db.execute(stmt)
|
|
await db.flush()
|
|
|
|
return {"message": f"User {user_id} has been banned"}
|
|
|
|
|
|
@router.post("/{user_id}/unban")
|
|
async def unban_user(
|
|
user_id: int,
|
|
db: AsyncSession = Depends(get_db),
|
|
current_user: dict = Depends(get_current_user),
|
|
):
|
|
"""Unban a user (admin only)."""
|
|
# Check if current user is admin
|
|
if current_user.get("privlevel") not in ["Admin", "Judge"]:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="Admin privileges required",
|
|
)
|
|
|
|
stmt = select(User).where(User.id == user_id)
|
|
result = await db.execute(stmt)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail="User not found",
|
|
)
|
|
|
|
# Update user ban status
|
|
stmt = (
|
|
update(User)
|
|
.where(User.id == user_id)
|
|
.values(
|
|
is_banned=False,
|
|
ban_reason=None,
|
|
ban_ends=None,
|
|
)
|
|
)
|
|
await db.execute(stmt)
|
|
await db.flush()
|
|
|
|
return {"message": f"User {user_id} has been unbanned"}
|