diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 4ef1edd..3b09424 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -20,9 +20,15 @@ repos: - repo: local hooks: + - id: prettier-frontend + name: prettier (frontend) + entry: bash -c 'cd frontend && args=(); for f in "$@"; do args+=("${f#frontend/}"); done; npx prettier --write --ignore-unknown "${args[@]}"' -- + language: system + files: ^frontend/.*\.(js|jsx|json|css|md)$ + pass_filenames: true - id: eslint-frontend name: eslint (frontend) - entry: bash -c 'cd frontend && npm run lint' + entry: bash -c 'cd frontend && args=(); for f in "$@"; do args+=("${f#frontend/}"); done; npx eslint --fix "${args[@]}" && npx eslint "${args[@]}"' -- language: system - files: ^frontend/src/.*\.(js|jsx)$ - pass_filenames: false + files: ^frontend/(src/.*\.(js|jsx)|eslint\.config\.js)$ + pass_filenames: true diff --git a/README.md b/README.md index 0ce5759..8dc93ac 100755 --- a/README.md +++ b/README.md @@ -70,6 +70,8 @@ This project is a ready-to-use fullstack template that leverages Docker Compose ./scripts/setup-pre-commit.sh ``` + Hooks auto-fix on commit: backend (Ruff format + lint fix), frontend (Prettier + ESLint fix), then re-stage if files changed. + re-enable hooks: ```bash diff --git a/backend/README.md b/backend/README.md index 9906a65..d58046b 100644 --- a/backend/README.md +++ b/backend/README.md @@ -29,19 +29,63 @@ This backend project is built with modern Python technologies to provide a robus - 🐳 Easy containerization with Docker - ✅ Async Testing & coverage with a fully isolated test environment -## Lint +## Lint & format -Requires [uv](https://docs.astral.sh/uv/) and dev dependencies (`uv sync`). +### Standards + +| Item | Value | +|------|--------| +| Config | [`pyproject.toml`](./pyproject.toml) — `[tool.ruff]`, `[tool.ruff.lint]`, `[tool.ruff.format]` | +| Formatter | [Ruff format](https://docs.astral.sh/ruff/formatter/) (Black-compatible) | +| Line length | 100 | +| Indent | 4 spaces; tabs are rewritten on format | +| Quotes | Double quotes | +| Target Python | 3.14 | + +| Rule set | Source | Purpose | +|----------|--------|---------| +| `E` | pycodestyle | PEP 8 style | +| `F` | Pyflakes | Unused imports, syntax issues | +| `I` | isort | Import order | +| `UP` | pyupgrade | Modern Python syntax | +| `B` | flake8-bugbear | Common bug patterns | + +| Ignored | Reason | +|---------|--------| +| `B008` | FastAPI `Depends()` in default arguments | +| `B904` | Exception chaining in FastAPI handlers | +| `E712` | SQLAlchemy boolean checks; auto-fix breaks ORM queries | + +| Per-file ignored | Path | Reason | +|------------------|------|--------| +| `B` | `tests/**` | Relaxed bugbear rules in tests | +| `E402` | `main.py`, `core/config.py`, `migrations/env.py` | Imports after bootstrap / env setup | +| `E501` | `utils/email_templates.py` | Long HTML email template lines | + +### Manual commands + +> Run from the `backend/` directory. Requires [uv](https://docs.astral.sh/uv/) and dev dependencies (`uv sync`). + +Lint the project; report issues without changing files: ```bash -cd backend uv run ruff check . +``` + +Check formatting only; report mismatches without writing: + +```bash uv run ruff format --check . ``` -Auto-fix: +Lint and auto-fix what Ruff can (imports, safe rewrites): ```bash uv run ruff check --fix . +``` + +Apply formatting to all Python files (including tab → spaces): + +```bash uv run ruff format . ``` diff --git a/backend/api/__init__.py b/backend/api/__init__.py index 4136ed9..fdd2903 100644 --- a/backend/api/__init__.py +++ b/backend/api/__init__.py @@ -1,18 +1,21 @@ from fastapi import APIRouter + from core.config import settings -from .auth.controller import router as auth_router + from .account.controller import router as account_router -from .users.controller import router as users_router +from .auth.controller import router as auth_router from .roles.controller import router as roles_router +from .users.controller import router as users_router api_router = APIRouter() if settings.DEBUG_MODE: from .debug.controller import router as debug_router + api_router.include_router(debug_router, prefix="/debug") # Add new API modules below. api_router.include_router(auth_router, prefix="/auth") api_router.include_router(account_router, prefix="/account") api_router.include_router(users_router, prefix="/users") -api_router.include_router(roles_router, prefix="/roles") \ No newline at end of file +api_router.include_router(roles_router, prefix="/roles") diff --git a/backend/api/account/controller.py b/backend/api/account/controller.py index 61e8e7d..ccfaa1e 100644 --- a/backend/api/account/controller.py +++ b/backend/api/account/controller.py @@ -1,16 +1,19 @@ -from core.redis import get_redis +from fastapi import APIRouter, Depends, HTTPException, Response +from sqlalchemy.ext.asyncio import AsyncSession + from core.dependencies import get_db +from core.redis import get_redis from core.security import verify_token -from sqlalchemy.ext.asyncio import AsyncSession -from fastapi import APIRouter, Depends, HTTPException, Response -from .schema import UserProfile, UserUpdate, PasswordChange -from utils.response import APIResponse, parse_responses, common_responses, make_error_examples -from .services import get_user_by_id, update_user_profile, change_password +from extensions.smtp import SMTPMailer, get_mailer from utils.custom_exception import AuthenticationException, NotFoundException -from extensions.smtp import get_mailer, SMTPMailer +from utils.response import APIResponse, common_responses, make_error_examples, parse_responses + +from .schema import PasswordChange, UserProfile, UserUpdate +from .services import change_password, get_user_by_id, update_user_profile router = APIRouter(tags=["Account"]) + def _to_user_profile(user) -> UserProfile: return UserProfile( id=user.id, @@ -23,17 +26,17 @@ def _to_user_profile(user) -> UserProfile: created_at=user.created_at, ) + @router.get( "/profile", response_model=APIResponse[UserProfile], summary="Get current user profile", - responses=parse_responses({ - 200: ("User profile retrieved successfully", UserProfile) - }, common_responses) + responses=parse_responses( + {200: ("User profile retrieved successfully", UserProfile)}, common_responses + ), ) async def get_user_profile_api( - token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db) ): """ Get the current authenticated user's profile information. @@ -41,56 +44,58 @@ async def get_user_profile_api( try: user_id = token.get("sub") user = await get_user_by_id(db, user_id) - + if not user: raise NotFoundException("User not found") - + user_data = _to_user_profile(user) - + return APIResponse(code=200, message="User profile retrieved successfully", data=user_data) except NotFoundException: raise HTTPException(status_code=404, detail="User not found") except Exception: raise HTTPException(status_code=500) + @router.put( "/profile", response_model=APIResponse[UserProfile], response_model_exclude_unset=True, summary="Update current user profile", - responses=parse_responses({ - 200: ("User profile updated successfully", UserProfile), - 202: ("Email verification required", UserProfile) - }, common_responses) + responses=parse_responses( + { + 200: ("User profile updated successfully", UserProfile), + 202: ("Email verification required", UserProfile), + }, + common_responses, + ), ) async def update_user_profile_api( user_update: UserUpdate, response: Response, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis), - mailer: SMTPMailer = Depends(get_mailer) + redis_client=Depends(get_redis), + mailer: SMTPMailer = Depends(get_mailer), ): """ Update the current authenticated user's profile information (excluding password). """ try: user_id = token.get("sub") - result = await update_user_profile( - db, user_id, user_update, mailer, redis_client - ) - + result = await update_user_profile(db, user_id, user_update, mailer, redis_client) + if not result: raise NotFoundException("User not found") - + user, email_change_requested = result - + user_data = _to_user_profile(user) if email_change_requested: response.status_code = 202 return APIResponse(code=202, message="Email verification required", data=user_data) - + return APIResponse(code=200, message="User profile updated successfully", data=user_data) except NotFoundException: raise HTTPException(status_code=404, detail="User not found") @@ -99,24 +104,35 @@ async def update_user_profile_api( except Exception: raise HTTPException(status_code=500) + @router.put( "/password", response_model=APIResponse[None], response_model_exclude_unset=True, summary="Change current user password", - responses=parse_responses({ - 200: ("Password changed successfully", None), - 401: ("Unauthorized", None, make_error_examples(401, { - "invalidToken": "Invalid or expired token", - "incorrectPassword": "Current password is incorrect", - })), - }, common_responses) + responses=parse_responses( + { + 200: ("Password changed successfully", None), + 401: ( + "Unauthorized", + None, + make_error_examples( + 401, + { + "invalidToken": "Invalid or expired token", + "incorrectPassword": "Current password is incorrect", + }, + ), + ), + }, + common_responses, + ), ) async def change_user_password_api( password_change: PasswordChange, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis) + redis_client=Depends(get_redis), ): """ Change the current authenticated user's password. @@ -131,11 +147,11 @@ async def change_user_password_api( redis_client, current_session_id=current_session_id, ) - + if success: return APIResponse(code=200, message="Password changed successfully") - + except AuthenticationException: raise HTTPException(status_code=401, detail="Current password is incorrect") except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) \ No newline at end of file + raise HTTPException(status_code=500, detail=str(e)) diff --git a/backend/api/account/schema.py b/backend/api/account/schema.py index f047a42..bcfde57 100644 --- a/backend/api/account/schema.py +++ b/backend/api/account/schema.py @@ -1,33 +1,45 @@ -from typing import Optional from datetime import datetime -from core.config import settings + from pydantic import BaseModel, EmailStr, Field, model_validator +from core.config import settings + + class UserProfile(BaseModel): id: str = Field(..., description="User ID") first_name: str = Field(..., description="First name") last_name: str = Field(..., description="Last name") email: EmailStr = Field(..., description="User email address") - pending_email: Optional[EmailStr] = Field(None, description="Pending email awaiting verification") + pending_email: EmailStr | None = Field(None, description="Pending email awaiting verification") phone: str = Field(..., description="Phone number") status: bool = Field(..., description="User status") created_at: datetime = Field(..., description="User creation time") + class UserUpdate(BaseModel): - first_name: Optional[str] = Field(None, min_length=1, max_length=50, description="First name") - last_name: Optional[str] = Field(None, min_length=1, max_length=50, description="Last name") - email: Optional[EmailStr] = Field(None, description="User email address") - phone: Optional[str] = Field(None, min_length=1, max_length=20, description="Phone number") + first_name: str | None = Field(None, min_length=1, max_length=50, description="First name") + last_name: str | None = Field(None, min_length=1, max_length=50, description="Last name") + email: EmailStr | None = Field(None, description="User email address") + phone: str | None = Field(None, min_length=1, max_length=20, description="Phone number") + class PasswordChange(BaseModel): - current_password: str = Field(..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="Current password") - new_password: str = Field(..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="New password") + current_password: str = Field( + ..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="Current password" + ) + new_password: str = Field( + ..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="New password" + ) logout_other_devices: bool = Field(True, description="Logout other devices") @model_validator(mode="before") @classmethod def migrate_legacy_logout_field(cls, data): - if isinstance(data, dict) and "logout_other_devices" not in data and "logout_all_devices" in data: + if ( + isinstance(data, dict) + and "logout_other_devices" not in data + and "logout_all_devices" in data + ): data = data.copy() data["logout_other_devices"] = data.pop("logout_all_devices") - return data \ No newline at end of file + return data diff --git a/backend/api/account/services.py b/backend/api/account/services.py index bb66a27..74fdac2 100644 --- a/backend/api/account/services.py +++ b/backend/api/account/services.py @@ -1,75 +1,82 @@ -import redis import logging -from typing import Optional, Tuple from urllib.parse import quote -from sqlalchemy import select, or_ -from models.users import Users + +import redis +from sqlalchemy import or_, select +from sqlalchemy.ext.asyncio import AsyncSession + +from api.auth.services import _request_email_change_verification_email from core.config import settings +from core.security import clear_user_all_sessions, hash_password, verify_password from extensions.smtp import SMTPMailer -from .schema import UserUpdate, PasswordChange -from sqlalchemy.ext.asyncio import AsyncSession -from utils.custom_exception import AuthenticationException, ServerException, SMTPNotConfiguredException -from core.security import hash_password, verify_password, clear_user_all_sessions +from models.users import Users +from utils.custom_exception import ( + AuthenticationException, + ServerException, + SMTPNotConfiguredException, +) from utils.email_templates import EMAIL_VERIFICATION_TEMPLATE -from api.auth.services import _request_email_change_verification_email + +from .schema import PasswordChange, UserUpdate logger = logging.getLogger("account") -async def get_user_by_id(db: AsyncSession, user_id: str) -> Optional[Users]: + +async def get_user_by_id(db: AsyncSession, user_id: str) -> Users | None: """Get user info by id""" - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + result = await db.execute(select(Users).where(Users.id == user_id)) return result.scalar_one_or_none() + async def update_user_profile( db: AsyncSession, user_id: str, user_update: UserUpdate, - mailer: Optional[SMTPMailer] = None, - redis_client: Optional[redis.Redis] = None -) -> Optional[Tuple[Users, bool]]: + mailer: SMTPMailer | None = None, + redis_client: redis.Redis | None = None, +) -> tuple[Users, bool] | None: """Update user info (excluding password)""" user = await get_user_by_id(db, user_id) if not user: return None - + email_change_requested = False new_email = user_update.email if new_email and new_email != user.email: result = await db.execute( select(Users).where( - or_( - Users.email == new_email, - Users.pending_email == new_email - ), - Users.id != user_id + or_(Users.email == new_email, Users.pending_email == new_email), Users.id != user_id ) ) if result.scalar_one_or_none(): raise ValueError("Email already exists") - + # Defer email change until verification completes. user.pending_email = new_email email_change_requested = True - + update_data = user_update.model_dump(exclude_unset=True) update_data.pop("email", None) for field, value in update_data.items(): setattr(user, field, value) - - if email_change_requested and settings.SMTP_ENABLE and mailer and getattr(mailer, "enabled", False): + + if ( + email_change_requested + and settings.SMTP_ENABLE + and mailer + and getattr(mailer, "enabled", False) + ): should_send = True if redis_client: cooldown_key = f"email_verification_cooldown:{new_email}" remaining_seconds = await redis_client.ttl(cooldown_key) try: remaining_seconds = int(remaining_seconds) - except (TypeError, ValueError): + except TypeError, ValueError: remaining_seconds = 0 if remaining_seconds > 0: should_send = False - + if should_send: try: token_meta = await _request_email_change_verification_email(db, user, new_email) @@ -78,56 +85,55 @@ async def update_user_profile( f"{settings.HOSTNAME}:{settings.FRONTEND_PORT}" f"/auth/verify-email?token={quote(token_meta['verification_token'], safe='')}" ) - + user_name = f"{user.first_name} {user.last_name}".strip() app_name = settings.PROJECT_NAME - + email_content = EMAIL_VERIFICATION_TEMPLATE.render( verification_url=verification_url, user_name=user_name, app_name=app_name, expire_minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES, ) - + mailer.send_text( to_emails=[new_email], subject=email_content["subject"], body=email_content["body"], html_body=email_content.get("html_body"), ) - + if redis_client: await redis_client.setex( - cooldown_key, - settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, - "1" + cooldown_key, settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, "1" ) except SMTPNotConfiguredException as exc: logger.warning("Skip email change verification send: %s", exc) - + await db.commit() await db.refresh(user) return user, email_change_requested + async def change_password( db: AsyncSession, user_id: str, password_change: PasswordChange, redis_client=None, - current_session_id: Optional[str] = None, + current_session_id: str | None = None, ) -> bool: """Change user password""" try: user = await get_user_by_id(db, user_id) if not user: return False - + if not await verify_password(password_change.current_password, user.hash_password): raise AuthenticationException("Current password is incorrect") - + user.hash_password = await hash_password(password_change.new_password) user.password_reset_required = False - + if password_change.logout_other_devices and redis_client: await clear_user_all_sessions( db, @@ -135,11 +141,11 @@ async def change_password( user_id, exclude_session_id=current_session_id, ) - + await db.commit() - + return True except AuthenticationException: raise except Exception: - raise ServerException("Failed to change password") \ No newline at end of file + raise ServerException("Failed to change password") diff --git a/backend/api/auth/controller.py b/backend/api/auth/controller.py index 50e1aff..c03e184 100644 --- a/backend/api/auth/controller.py +++ b/backend/api/auth/controller.py @@ -1,85 +1,92 @@ import logging -from core.redis import get_redis +from datetime import datetime, timedelta + +from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + from core.config import settings -from .schema import UserResponse from core.dependencies import get_db -from datetime import datetime, timedelta +from core.redis import get_redis +from core.security import verify_email_verification_token, verify_password_reset_token, verify_token +from extensions.smtp import SMTPMailer, get_mailer +from utils.custom_exception import ( + AuthenticationException, + ConflictException, + EmailVerificationRequiredException, + NotFoundException, + PasswordResetRequiredException, + RegistrationDisabledException, + SMTPNotConfiguredException, + ValidationException, +) from utils.get_real_ip import get_real_ip -from sqlalchemy.ext.asyncio import AsyncSession -from core.security import verify_password_reset_token, verify_token, verify_email_verification_token -from fastapi import APIRouter, Depends, HTTPException, Request, Response, Query -from utils.response import APIResponse, parse_responses, common_responses, make_error_examples -from extensions.smtp import get_mailer, SMTPMailer -from utils.custom_exception import SMTPNotConfiguredException +from utils.response import APIResponse, common_responses, make_error_examples, parse_responses + from .schema import ( - UserRegister, - UserLogin, - UserLoginResponse, - TokenResponse, + ActionRequiredResponse, CsrfTokenResponse, - ResetPasswordRequest, - TokenValidationResponse, - LogoutRequest, ForgotPasswordRequest, + LogoutRequest, PasswordResetCooldownResponse, ResendVerificationRequest, - ActionRequiredResponse, - action_required_response_examples + ResetPasswordRequest, + TokenResponse, + TokenValidationResponse, + UserLogin, + UserLoginResponse, + UserRegister, + UserResponse, + action_required_response_examples, ) from .services import ( - register, + forgot_password, + get_or_create_csrf_token, + get_password_reset_cooldown, login, logout, - token, logout_all_devices, + register, + resend_verification_email, reset_password, + token, validate_password_reset_token, - forgot_password, - get_password_reset_cooldown, - verify_email, - resend_verification_email, - get_or_create_csrf_token, verify_csrf_token, -) -from utils.custom_exception import ( - ConflictException, - AuthenticationException, - PasswordResetRequiredException, - NotFoundException, - ValidationException, - EmailVerificationRequiredException, - RegistrationDisabledException, + verify_email, ) logger = logging.getLogger("auth") router = APIRouter(tags=["Auth"]) + @router.post( - "/register", - response_model=APIResponse[UserLoginResponse], + "/register", + response_model=APIResponse[UserLoginResponse], response_model_exclude_none=True, summary="Register account", - responses=parse_responses({ - 200: ("User registered successfully", UserLoginResponse), - 202: ("Email verification required", None), - 409: ("Email already exists", None), - 503: ("Registration is disabled", None) - }, common_responses) + responses=parse_responses( + { + 200: ("User registered successfully", UserLoginResponse), + 202: ("Email verification required", None), + 409: ("Email already exists", None), + 503: ("Registration is disabled", None), + }, + common_responses, + ), ) async def register_api( user_data: UserRegister, request: Request, response: Response, db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis), - mailer: SMTPMailer = Depends(get_mailer) + redis_client=Depends(get_redis), + mailer: SMTPMailer = Depends(get_mailer), ): try: client_ip = get_real_ip(request) user_agent = request.headers.get("user-agent", "Registration") - + result = await register(db, redis_client, user_data, client_ip, user_agent, mailer) - + user = result["user"] session_id = result["session_id"] access_token = result["access_token"] @@ -87,22 +94,23 @@ async def register_api( _set_session_cookie(response, session_id) _set_csrf_cookie(response, csrf_token) - + user_response = UserResponse( id=user.id, first_name=user.first_name, last_name=user.last_name, email=user.email, - phone=user.phone + phone=user.phone, ) - + response_data = UserLoginResponse( access_token=access_token, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), - user=user_response + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), + user=user_response, ) return APIResponse(code=200, message="User registered successfully", data=response_data) - except EmailVerificationRequiredException as e: + except EmailVerificationRequiredException: resp = APIResponse(code=202, message="Email verification required") raise HTTPException(status_code=202, detail=resp.dict(exclude_none=True)) except ConflictException: @@ -112,31 +120,35 @@ async def register_api( except Exception: raise HTTPException(status_code=500) + @router.post( - "/login", + "/login", response_model=APIResponse[UserLoginResponse], response_model_exclude_none=True, summary="Login account", - responses=parse_responses({ - 200: ("User logged in successfully", UserLoginResponse), - 202: ("Action required", ActionRequiredResponse, action_required_response_examples), - 401: ("Invalid email or password", None) - }, common_responses) + responses=parse_responses( + { + 200: ("User logged in successfully", UserLoginResponse), + 202: ("Action required", ActionRequiredResponse, action_required_response_examples), + 401: ("Invalid email or password", None), + }, + common_responses, + ), ) async def login_api( user_data: UserLogin, request: Request, response: Response, db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis), - mailer: SMTPMailer = Depends(get_mailer) + redis_client=Depends(get_redis), + mailer: SMTPMailer = Depends(get_mailer), ): try: client_ip = get_real_ip(request) user_agent = request.headers.get("user-agent", "") - + result = await login(db, redis_client, user_data, client_ip, user_agent, mailer) - + user = result["user"] session_id = result["session_id"] access_token = result["access_token"] @@ -147,16 +159,17 @@ async def login_api( first_name=user.first_name, last_name=user.last_name, email=user.email, - phone=user.phone + phone=user.phone, ) _set_session_cookie(response, session_id) _set_csrf_cookie(response, csrf_token) - + response_data = UserLoginResponse( access_token=access_token, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), - user=user_response + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), + user=user_response, ) return APIResponse(code=200, message="User logged in successfully", data=response_data) except PasswordResetRequiredException as e: @@ -165,41 +178,52 @@ async def login_api( except EmailVerificationRequiredException as e: resp = APIResponse(code=202, message="Email verification required", data=e.details) raise HTTPException(status_code=202, detail=resp.dict(exclude_none=True)) - except AuthenticationException as e: + except AuthenticationException: raise HTTPException(status_code=401, detail="Invalid email or password") except Exception: raise HTTPException(status_code=500) + @router.post( - "/logout", + "/logout", response_model=APIResponse[None], response_model_exclude_none=True, summary="Logout account", - responses=parse_responses({ - 200: ("User logged out successfully", None), - 401: ("Unauthorized", None, make_error_examples(401, { - "invalidSession": "Invalid or expired session", - "invalidToken": "Invalid or expired token", - })), - }, common_responses) + responses=parse_responses( + { + 200: ("User logged out successfully", None), + 401: ( + "Unauthorized", + None, + make_error_examples( + 401, + { + "invalidSession": "Invalid or expired session", + "invalidToken": "Invalid or expired token", + }, + ), + ), + }, + common_responses, + ), ) async def logout_api( logout_data: LogoutRequest, token: dict = Depends(verify_token), response: Response = None, db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis) + redis_client=Depends(get_redis), ): """ Logout user from current device or all devices - + Args: logout_data: Contains logout_all flag to determine logout scope """ try: user_id = token.get("sub") session_id = token.get("sid") - + if logout_data.logout_all: # Logout from all devices if await logout_all_devices(db, redis_client, user_id): @@ -210,7 +234,7 @@ async def logout_api( # Logout from current device only if not session_id: raise AuthenticationException("Invalid or expired session") - + if await logout(db, redis_client, user_id, session_id): if response: _clear_auth_cookies(response) @@ -220,23 +244,32 @@ async def logout_api( except Exception: raise HTTPException(status_code=500) + @router.post( "/token", response_model=APIResponse[TokenResponse], response_model_exclude_unset=True, summary="Refresh token", - responses=parse_responses({ - 200: ("Token refreshed successfully", TokenResponse), - 401: ("Unauthorized", None, make_error_examples(401, { - "invalidSession": "Invalid or expired session", - "invalidCsrf": "Invalid or expired CSRF token", - })), - }, common_responses) + responses=parse_responses( + { + 200: ("Token refreshed successfully", TokenResponse), + 401: ( + "Unauthorized", + None, + make_error_examples( + 401, + { + "invalidSession": "Invalid or expired session", + "invalidCsrf": "Invalid or expired CSRF token", + }, + ), + ), + }, + common_responses, + ), ) async def token_api( - request: Request, - db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis) + request: Request, db: AsyncSession = Depends(get_db), redis_client=Depends(get_redis) ): """Validate CSRF token then use session_id cookie to issue new access_token""" try: @@ -250,7 +283,8 @@ async def token_api( new_access_token = await token(db, redis_client, session_id) response_data = TokenResponse( access_token=new_access_token, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), ) return APIResponse(code=200, message="Token refreshed successfully", data=response_data) except AuthenticationException as e: @@ -260,15 +294,19 @@ async def token_api( except Exception: raise HTTPException(status_code=500) + @router.post( "/csrf-token", response_model=APIResponse[CsrfTokenResponse], response_model_exclude_none=True, summary="Get CSRF token", - responses=parse_responses({ - 200: ("CSRF token retrieved successfully", CsrfTokenResponse), - 401: ("Invalid or expired session", None), - }, common_responses), + responses=parse_responses( + { + 200: ("CSRF token retrieved successfully", CsrfTokenResponse), + 401: ("Invalid or expired session", None), + }, + common_responses, + ), ) async def csrf_token_api( request: Request, @@ -289,22 +327,28 @@ async def csrf_token_api( expires_at=datetime.now().astimezone() + timedelta(minutes=settings.CSRF_TOKEN_EXPIRE_MINUTES), ) - return APIResponse(code=200, message="CSRF token retrieved successfully", data=response_data) + return APIResponse( + code=200, message="CSRF token retrieved successfully", data=response_data + ) except AuthenticationException: raise HTTPException(status_code=401, detail="Invalid or expired session") except Exception: raise HTTPException(status_code=500) + @router.post( "/reset-password", response_model=APIResponse[UserLoginResponse], response_model_exclude_none=True, summary="Reset password with token", - responses=parse_responses({ - 200: ("Password reset successfully", UserLoginResponse), - 401: ("Invalid or expired token", None), - 404: ("User not found", None) - }, common_responses) + responses=parse_responses( + { + 200: ("Password reset successfully", UserLoginResponse), + 401: ("Invalid or expired token", None), + 404: ("User not found", None), + }, + common_responses, + ), ) async def reset_password_api( request: Request, @@ -312,37 +356,40 @@ async def reset_password_api( request_data: ResetPasswordRequest, token: dict = Depends(verify_password_reset_token), db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis) + redis_client=Depends(get_redis), ): """Reset password using token""" try: client_ip = get_real_ip(request) user_agent = request.headers.get("user-agent", "") - - result = await reset_password(db, redis_client, token, request_data.new_password, client_ip, user_agent) - + + result = await reset_password( + db, redis_client, token, request_data.new_password, client_ip, user_agent + ) + user = result["user"] session_id = result["session_id"] access_token = result["access_token"] csrf_token = result["csrf_token"] - + user_response = UserResponse( id=user.id, first_name=user.first_name, last_name=user.last_name, email=user.email, - phone=user.phone + phone=user.phone, ) - + response_data = UserLoginResponse( access_token=access_token, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), - user=user_response + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), + user=user_response, ) _set_session_cookie(response, session_id) _set_csrf_cookie(response, csrf_token) - + return APIResponse(code=200, message="Password reset successfully", data=response_data) except AuthenticationException: raise HTTPException(status_code=401, detail="Invalid or expired token") @@ -351,18 +398,16 @@ async def reset_password_api( except Exception: raise HTTPException(status_code=500) + @router.get( "/validate-reset-token", response_model=APIResponse[TokenValidationResponse], response_model_exclude_none=True, summary="Validate password reset token", - responses=parse_responses({ - 200: ("Token is valid", TokenValidationResponse) - }, common_responses) + responses=parse_responses({200: ("Token is valid", TokenValidationResponse)}, common_responses), ) async def validate_reset_token_api( - token: dict = Depends(verify_password_reset_token), - db: AsyncSession = Depends(get_db) + token: dict = Depends(verify_password_reset_token), db: AsyncSession = Depends(get_db) ): """Validate password reset token without consuming it""" try: @@ -379,29 +424,34 @@ async def validate_reset_token_api( response_model=APIResponse[None], response_model_exclude_none=True, summary="Send reset password email", - responses=parse_responses({ - 200: ("Reset password email sent", None), - 400: ("Please wait before requesting another password reset email", None), - 403: ("Account is disabled", None), - 404: ("User not registered", None), - 503: ("SMTP is disabled", None), - }, common_responses), + responses=parse_responses( + { + 200: ("Reset password email sent", None), + 400: ("Please wait before requesting another password reset email", None), + 403: ("Account is disabled", None), + 404: ("User not registered", None), + 503: ("SMTP is disabled", None), + }, + common_responses, + ), ) async def forgot_password_api( request: Request, request_data: ForgotPasswordRequest, db: AsyncSession = Depends(get_db), mailer: SMTPMailer = Depends(get_mailer), - redis_client = Depends(get_redis), + redis_client=Depends(get_redis), ): """ Send password reset email based on input email. """ try: - await forgot_password(db, request_data.email, mailer, redis_client) + await forgot_password(db, request_data.email, mailer, redis_client) return APIResponse(code=200, message="Reset password email sent") except ValidationException: - raise HTTPException(status_code=400, detail="Please wait before requesting another password reset email") + raise HTTPException( + status_code=400, detail="Please wait before requesting another password reset email" + ) except AuthenticationException: raise HTTPException(status_code=403, detail="Account is disabled") except NotFoundException: @@ -411,18 +461,22 @@ async def forgot_password_api( except Exception: raise HTTPException(status_code=500) + @router.get( "/forgot-password/cooldown", response_model=APIResponse[PasswordResetCooldownResponse], response_model_exclude_none=True, summary="Get password reset email cooldown status", - responses=parse_responses({ - 200: ("Cooldown status retrieved", PasswordResetCooldownResponse), - }, common_responses), + responses=parse_responses( + { + 200: ("Cooldown status retrieved", PasswordResetCooldownResponse), + }, + common_responses, + ), ) async def get_password_reset_cooldown_api( email: str = Query(..., description="Email address to check cooldown for"), - redis_client = Depends(get_redis), + redis_client=Depends(get_redis), ): """ Get remaining cooldown time for password reset email. @@ -430,61 +484,64 @@ async def get_password_reset_cooldown_api( """ try: result = await get_password_reset_cooldown(email, redis_client) - response_data = PasswordResetCooldownResponse( - cooldown_seconds=result["cooldown_seconds"] - ) + response_data = PasswordResetCooldownResponse(cooldown_seconds=result["cooldown_seconds"]) return APIResponse(code=200, message="Cooldown status retrieved", data=response_data) except Exception: raise HTTPException(status_code=500) + @router.get( "/verify-email", response_model=APIResponse[UserLoginResponse], response_model_exclude_none=True, summary="Verify email address", - responses=parse_responses({ - 200: ("Email verified successfully", UserLoginResponse), - 401: ("Invalid or expired token", None), - 404: ("User not found", None), - 409: ("Email already exists", None) - }, common_responses) + responses=parse_responses( + { + 200: ("Email verified successfully", UserLoginResponse), + 401: ("Invalid or expired token", None), + 404: ("User not found", None), + 409: ("Email already exists", None), + }, + common_responses, + ), ) async def verify_email_api( request: Request, response: Response, token: dict = Depends(verify_email_verification_token), db: AsyncSession = Depends(get_db), - redis_client = Depends(get_redis) + redis_client=Depends(get_redis), ): """Verify email address using token and create session""" try: client_ip = get_real_ip(request) user_agent = request.headers.get("user-agent", "") - + result = await verify_email(db, redis_client, token, client_ip, user_agent) - + user = result["user"] session_id = result["session_id"] access_token = result["access_token"] csrf_token = result["csrf_token"] - + user_response = UserResponse( id=user.id, first_name=user.first_name, last_name=user.last_name, email=user.email, - phone=user.phone + phone=user.phone, ) - + response_data = UserLoginResponse( access_token=access_token, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), - user=user_response + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES), + user=user_response, ) _set_session_cookie(response, session_id) _set_csrf_cookie(response, csrf_token) - + return APIResponse(code=200, message="Email verified successfully", data=response_data) except AuthenticationException: raise HTTPException(status_code=401, detail="Invalid or expired token") @@ -495,32 +552,38 @@ async def verify_email_api( except Exception: raise HTTPException(status_code=500) + @router.post( "/resend-verification", response_model=APIResponse[None], response_model_exclude_none=True, summary="Resend email verification", - responses=parse_responses({ - 200: ("Verification email sent", None), - 400: ("Please wait before requesting another verification email", None), - 403: ("Account is disabled", None), - 404: ("User not registered", None), - 503: ("SMTP is disabled", None), - }, common_responses) + responses=parse_responses( + { + 200: ("Verification email sent", None), + 400: ("Please wait before requesting another verification email", None), + 403: ("Account is disabled", None), + 404: ("User not registered", None), + 503: ("SMTP is disabled", None), + }, + common_responses, + ), ) async def resend_verification_api( request: Request, request_data: ResendVerificationRequest, db: AsyncSession = Depends(get_db), mailer: SMTPMailer = Depends(get_mailer), - redis_client = Depends(get_redis), + redis_client=Depends(get_redis), ): """Resend email verification email""" try: await resend_verification_email(db, request_data.email, mailer, redis_client) return APIResponse(code=200, message="Verification email sent") except ValidationException: - raise HTTPException(status_code=400, detail="Please wait before requesting another verification email") + raise HTTPException( + status_code=400, detail="Please wait before requesting another verification email" + ) except AuthenticationException: raise HTTPException(status_code=403, detail="Account is disabled") except NotFoundException: @@ -530,18 +593,22 @@ async def resend_verification_api( except Exception: raise HTTPException(status_code=500) + @router.get( "/resend-verification/cooldown", response_model=APIResponse[PasswordResetCooldownResponse], response_model_exclude_none=True, summary="Get email verification cooldown status", - responses=parse_responses({ - 200: ("Cooldown status retrieved", PasswordResetCooldownResponse), - }, common_responses), + responses=parse_responses( + { + 200: ("Cooldown status retrieved", PasswordResetCooldownResponse), + }, + common_responses, + ), ) async def get_email_verification_cooldown_api( email: str = Query(..., description="Email address to check cooldown for"), - redis_client = Depends(get_redis), + redis_client=Depends(get_redis), ): """ Get remaining cooldown time for email verification email. @@ -550,18 +617,17 @@ async def get_email_verification_cooldown_api( try: cooldown_key = f"email_verification_cooldown:{email}" remaining_seconds = await redis_client.ttl(cooldown_key) - + # TTL returns -1 if key exists but has no expiry, -2 if key doesn't exist if remaining_seconds < 0: remaining_seconds = 0 - - response_data = PasswordResetCooldownResponse( - cooldown_seconds=remaining_seconds - ) + + response_data = PasswordResetCooldownResponse(cooldown_seconds=remaining_seconds) return APIResponse(code=200, message="Cooldown status retrieved", data=response_data) except Exception: raise HTTPException(status_code=500) + def _set_session_cookie(response: Response, session_id: str) -> None: response.set_cookie( key="session_id", @@ -586,4 +652,4 @@ def _set_csrf_cookie(response: Response, csrf_token: str) -> None: def _clear_auth_cookies(response: Response) -> None: response.delete_cookie("session_id") - response.delete_cookie("csrf_token") \ No newline at end of file + response.delete_cookie("csrf_token") diff --git a/backend/api/auth/schema.py b/backend/api/auth/schema.py index bb29c01..82dd677 100644 --- a/backend/api/auth/schema.py +++ b/backend/api/auth/schema.py @@ -1,30 +1,39 @@ -from typing import TypedDict, Optional from datetime import datetime -from core.config import settings +from typing import TypedDict + from pydantic import BaseModel, EmailStr, Field +from core.config import settings + + class LoginResult(TypedDict): - user: "UserResponse" - session_id: str = Field(..., description="Session ID") + user: UserResponse + session_id: str = Field(..., description="Session ID") access_token: str = Field(..., description="JWT access token") csrf_token: str = Field(..., description="CSRF token") + class SessionResult(TypedDict): session_id: str = Field(..., description="Session ID") access_token: str = Field(..., description="JWT access token") csrf_token: str = Field(..., description="CSRF token") + class UserRegister(BaseModel): first_name: str = Field(..., min_length=1, max_length=50, description="First name") last_name: str = Field(..., min_length=1, max_length=50, description="Last name") email: EmailStr = Field(..., description="User email address") phone: str = Field(..., min_length=1, max_length=20, description="Phone number") - password: str = Field(..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="Password") + password: str = Field( + ..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="Password" + ) + class UserLogin(BaseModel): email: EmailStr = Field(..., description="User email address") password: str = Field(..., min_length=1, description="Password") + class UserResponse(BaseModel): id: str = Field(..., description="User ID") first_name: str = Field(..., description="First name") @@ -32,52 +41,71 @@ class UserResponse(BaseModel): email: str = Field(..., description="User email address") phone: str = Field(..., description="Phone number") + class UserLoginResponse(BaseModel): access_token: str = Field(..., description="JWT access token") expires_at: datetime = Field(..., description="Token expiration time") user: UserResponse = Field(..., description="User information") + class TokenResponse(BaseModel): access_token: str = Field(..., description="JWT access token") expires_at: datetime = Field(..., description="Token expiration time") + class CsrfTokenResponse(BaseModel): csrf_token: str = Field(..., description="CSRF token") expires_at: datetime = Field(..., description="CSRF token expiration time") + class ActionRequiredResponse(BaseModel): - action_type: str = Field(..., description="Action type for frontend routing: 'password_reset' or 'email_verification'") - token: Optional[str] = Field(default=None, description="Token for the password reset") - expires_at: Optional[str] = Field(default=None, description="Token expiration time (ISO format)") + action_type: str = Field( + ..., + description="Action type for frontend routing: 'password_reset' or 'email_verification'", + ) + token: str | None = Field(default=None, description="Token for the password reset") + expires_at: str | None = Field(default=None, description="Token expiration time (ISO format)") + class LogoutRequest(BaseModel): logout_all: bool = Field(False, description="Whether to logout from all devices") + class ResetPasswordRequest(BaseModel): - new_password: str = Field(..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="New password") + new_password: str = Field( + ..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="New password" + ) + class TokenValidationResponse(BaseModel): is_valid: bool = Field(..., description="Whether the token is valid") + class ForgotPasswordRequest(BaseModel): email: EmailStr = Field(..., description="User email address") + class PasswordResetCooldownResponse(BaseModel): cooldown_seconds: int = Field(..., description="Remaining cooldown time in seconds") + class EmailVerificationResponse(BaseModel): message: str = Field(..., description="Verification result message") + class EmailVerificationRequiredResponse(BaseModel): - expires_at: Optional[str] = Field(default=None, description="Token expiration time (ISO format)") + expires_at: str | None = Field(default=None, description="Token expiration time (ISO format)") + class PasswordResetRequiredResponse(BaseModel): reset_token: str = Field(..., description="Password reset token") expires_at: str = Field(..., description="Token expiration time (ISO format)") + class ResendVerificationRequest(BaseModel): email: EmailStr = Field(..., description="Email address to resend verification") + action_required_response_examples = { "passwordReset": { "summary": "Password reset required", @@ -87,9 +115,9 @@ class ResendVerificationRequest(BaseModel): "data": { "action_type": "password_reset", "token": "password_reset_token", - "expires_at": "2024-01-01T12:00:00+00:00" - } - } + "expires_at": "2024-01-01T12:00:00+00:00", + }, + }, }, "emailVerification": { "summary": "Email verification required", @@ -99,8 +127,8 @@ class ResendVerificationRequest(BaseModel): "data": { "action_type": "email_verification", "token": None, - "expires_at": "2024-01-01T12:00:00+00:00" - } - } - } -} \ No newline at end of file + "expires_at": "2024-01-01T12:00:00+00:00", + }, + }, + }, +} diff --git a/backend/api/auth/services.py b/backend/api/auth/services.py index 4c03c6a..aec12fd 100644 --- a/backend/api/auth/services.py +++ b/backend/api/auth/services.py @@ -1,118 +1,117 @@ import ast -import redis -from jose import jwt, JWTError -from typing import Optional -from sqlalchemy import select, update, or_ -from urllib.parse import quote -from models.users import Users -from core.config import settings -from extensions.smtp import SMTPMailer -from models.login_logs import LoginLogs from datetime import datetime, timedelta -from models.user_sessions import UserSessions +from urllib.parse import quote + +import redis +from jose import JWTError, jwt +from sqlalchemy import func, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession -from utils.email_templates import ( - PASSWORD_RESET_TEMPLATE, - EMAIL_VERIFICATION_TEMPLATE, -) -from models.password_reset_tokens import PasswordResetTokens -from models.email_verification_tokens import EmailVerificationTokens -from .schema import ( - UserRegister, - UserLogin, - LoginResult, - SessionResult, - TokenValidationResponse, - ActionRequiredResponse, -) + +from core.config import settings from core.security import ( - verify_password, + clear_user_all_sessions, create_access_token, create_csrf_token, - hash_password, - extend_session_ttl, - clear_user_all_sessions, + create_email_verification_token, create_password_reset_token, - create_email_verification_token + extend_session_ttl, + hash_password, + verify_password, ) +from extensions.smtp import SMTPMailer +from models.email_verification_tokens import EmailVerificationTokens +from models.login_logs import LoginLogs +from models.password_reset_tokens import PasswordResetTokens +from models.user_sessions import UserSessions +from models.users import Users from utils.custom_exception import ( - ConflictException, AuthenticationException, - PasswordResetRequiredException, + ConflictException, + EmailVerificationRequiredException, NotFoundException, - SMTPNotConfiguredException, + PasswordResetRequiredException, + RegistrationDisabledException, ServerException, + SMTPNotConfiguredException, ValidationException, - EmailVerificationRequiredException, - RegistrationDisabledException, +) +from utils.email_templates import ( + EMAIL_VERIFICATION_TEMPLATE, + PASSWORD_RESET_TEMPLATE, +) + +from .schema import ( + ActionRequiredResponse, + LoginResult, + SessionResult, + TokenValidationResponse, + UserLogin, + UserRegister, ) async def register( - db: AsyncSession, + db: AsyncSession, redis_client: redis.Redis, - user_data: UserRegister, - ip_address: str, + user_data: UserRegister, + ip_address: str, user_agent: str, - mailer: Optional[SMTPMailer] = None + mailer: SMTPMailer | None = None, ) -> LoginResult: """User register""" if not settings.REGISTRATION_ENABLE: raise RegistrationDisabledException("Registration is disabled") - + user = await _create_user(db, user_data) - + # Check if email verification is required if settings.EMAIL_VERIFICATION_ENABLE and settings.SMTP_ENABLE and mailer: # Check cooldown cooldown_key = f"email_verification_cooldown:{user.email}" remaining_seconds = await redis_client.ttl(cooldown_key) - + if remaining_seconds > 0: # In cooldown, return 202 without data await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) raise EmailVerificationRequiredException( - message="Email verification required", - details=None + message="Email verification required", details=None ) else: # Not in cooldown, send verification email await _send_registration_verification_email(db, mailer, user) - + # Set cooldown await redis_client.setex( - cooldown_key, - settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, - "1" + cooldown_key, settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, "1" ) - + await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) raise EmailVerificationRequiredException( - message="Email verification required", - details=None + message="Email verification required", details=None ) - - session_result = await _create_user_session( - db, redis_client, user, ip_address, user_agent - ) + + session_result = await _create_user_session(db, redis_client, user, ip_address, user_agent) await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) return { "user": user, @@ -121,145 +120,151 @@ async def register( "csrf_token": session_result["csrf_token"], } + async def login( db: AsyncSession, redis_client: redis.Redis, - login_data: UserLogin, - ip_address: str, + login_data: UserLogin, + ip_address: str, user_agent: str, - mailer: Optional[SMTPMailer] = None + mailer: SMTPMailer | None = None, ) -> LoginResult: """User login""" - result = await db.execute( - select(Users).where(Users.email == login_data.email) - ) + result = await db.execute(select(Users).where(Users.email == login_data.email)) user = result.scalar_one_or_none() - + if not user: await _log_login_attempt( - db, email=login_data.email, + db, + email=login_data.email, ip_address=ip_address, user_agent=user_agent, is_success=False, - failure_reason="User not found" + failure_reason="User not found", ) raise AuthenticationException("Invalid email or password") - + # Check if user account is disabled if not user.status: await _log_login_attempt( - db, email=login_data.email, + db, + email=login_data.email, ip_address=ip_address, user_agent=user_agent, is_success=False, - failure_reason="Account disabled" + failure_reason="Account disabled", ) raise AuthenticationException("Account is disabled") - + # Now verify password if not await verify_password(login_data.password, user.hash_password): await _log_login_attempt( - db, email=login_data.email, + db, + email=login_data.email, ip_address=ip_address, user_agent=user_agent, is_success=False, - failure_reason="Invalid password" + failure_reason="Invalid password", ) raise AuthenticationException("Invalid email or password") - + # Check if password reset is required if user.password_reset_required: reset_token = await create_password_reset_token(user.id, user.email) - + reset_token_record = PasswordResetTokens( user_id=user.id, token=reset_token, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.PASSWORD_RESET_TOKEN_EXPIRE_MINUTES) + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.PASSWORD_RESET_TOKEN_EXPIRE_MINUTES), ) db.add(reset_token_record) await db.commit() - + await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) - + raise PasswordResetRequiredException( message="Password reset required", details=ActionRequiredResponse( action_type="password_reset", token=reset_token, - expires_at=reset_token_record.expires_at.isoformat() if reset_token_record.expires_at else None - ) + expires_at=reset_token_record.expires_at.isoformat() + if reset_token_record.expires_at + else None, + ), ) - + # Check if email verification is required if settings.EMAIL_VERIFICATION_ENABLE and settings.SMTP_ENABLE and mailer: if not user.email_verified: # Check cooldown cooldown_key = f"email_verification_cooldown:{user.email}" remaining_seconds = await redis_client.ttl(cooldown_key) - + if remaining_seconds > 0: # In cooldown, return 202 with cooldown time await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) # Calculate expires_at from cooldown - expires_at = (datetime.now().astimezone() + timedelta(seconds=remaining_seconds)).isoformat() + expires_at = ( + datetime.now().astimezone() + timedelta(seconds=remaining_seconds) + ).isoformat() raise EmailVerificationRequiredException( message="Email verification required", details=ActionRequiredResponse( - action_type="email_verification", - token=None, - expires_at=expires_at - ) + action_type="email_verification", token=None, expires_at=expires_at + ), ) else: # Not in cooldown, send verification email await _send_registration_verification_email(db, mailer, user) - + # Set cooldown await redis_client.setex( - cooldown_key, - settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, - "1" + cooldown_key, settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, "1" ) - + await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) # Calculate expires_at from cooldown - expires_at = (datetime.now().astimezone() + timedelta(seconds=settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS)).isoformat() + expires_at = ( + datetime.now().astimezone() + + timedelta(seconds=settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS) + ).isoformat() raise EmailVerificationRequiredException( message="Email verification required", details=ActionRequiredResponse( - action_type="email_verification", - token=None, - expires_at=expires_at - ) + action_type="email_verification", token=None, expires_at=expires_at + ), ) - - session_result = await _create_user_session( - db, redis_client, user, ip_address, user_agent - ) + + session_result = await _create_user_session(db, redis_client, user, ip_address, user_agent) await _log_login_attempt( - db, email=user.email, + db, + email=user.email, ip_address=ip_address, user_agent=user_agent, is_success=True, - user_id=user.id + user_id=user.id, ) return { "user": user, @@ -268,43 +273,38 @@ async def login( "csrf_token": session_result["csrf_token"], } + async def logout( - db: AsyncSession, - redis_client: redis.Redis, - user_id: str, - session_id: str + db: AsyncSession, redis_client: redis.Redis, user_id: str, session_id: str ) -> bool: """User logout""" try: redis_key = f"session:{session_id}" await redis_client.delete(redis_key, f"csrf:{session_id}") - + result = await db.execute( select(UserSessions).where( - UserSessions.user_id == user_id, - UserSessions.id == session_id + UserSessions.user_id == user_id, UserSessions.id == session_id ) ) session = result.scalar_one_or_none() if session: session.is_active = False await db.commit() - + return True except Exception: raise ServerException("Logout failed") -async def logout_all_devices( - db: AsyncSession, - redis_client: redis.Redis, - user_id: str -) -> bool: + +async def logout_all_devices(db: AsyncSession, redis_client: redis.Redis, user_id: str) -> bool: """Logout user from all devices""" try: return await clear_user_all_sessions(db, redis_client, user_id) except Exception: raise ServerException("Failed to logout all devices") + async def get_or_create_csrf_token( redis_client: redis.Redis, session_id: str, @@ -323,7 +323,7 @@ async def get_or_create_csrf_token( async def verify_csrf_token( redis_client: redis.Redis, - csrf_token: Optional[str], + csrf_token: str | None, ) -> str: """Validate CSRF token and return session_id.""" if not csrf_token: @@ -356,11 +356,7 @@ async def verify_csrf_token( return session_id -async def token( - db: AsyncSession, - redis_client: redis.Redis, - session_id: str -) -> str: +async def token(db: AsyncSession, redis_client: redis.Redis, session_id: str) -> str: """Use session_id (Cookie) to issue new access_token and refresh session""" raw = await redis_client.get(f"session:{session_id}") if not raw: @@ -378,122 +374,112 @@ async def token( user = await _get_user_by_id(db, user_id) if not user: raise NotFoundException("User not found") - + if not user.status: raise AuthenticationException("Account is disabled") - new_access_token = await create_access_token(data={ - "sub": user_id, - "email": user.email, - "sid": session_id - }) - + new_access_token = await create_access_token( + data={"sub": user_id, "email": user.email, "sid": session_id} + ) + data["access_token"] = new_access_token await extend_session_ttl(redis_client, session_id, data) await _update_session_expiry(db, session_id) - + return new_access_token + async def reset_password( db: AsyncSession, redis_client: redis.Redis, token: dict, new_password: str, ip_address: str, - user_agent: str + user_agent: str, ) -> LoginResult: """Reset password using token""" try: user_id = token.get("sub") token_string = token.get("token") - + result = await db.execute( select(PasswordResetTokens).where( PasswordResetTokens.token == token_string, PasswordResetTokens.user_id == user_id, - PasswordResetTokens.is_used == False, - PasswordResetTokens.expires_at > datetime.now().astimezone() + PasswordResetTokens.is_used.is_(False), + PasswordResetTokens.expires_at > func.now(), ) ) token_record = result.scalar_one_or_none() - + if not token_record: raise AuthenticationException("Invalid or expired token") - - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + + result = await db.execute(select(Users).where(Users.id == user_id)) user = result.scalar_one_or_none() - + if not user: raise NotFoundException("User not found") - + user.hash_password = await hash_password(new_password) user.password_reset_required = False - + token_record.is_used = True - + # Force logout all devices await clear_user_all_sessions(db, redis_client, user_id) - + # Create new session - session_result = await _create_user_session( - db, redis_client, user, ip_address, user_agent - ) - + session_result = await _create_user_session(db, redis_client, user, ip_address, user_agent) + await db.commit() - + return { "user": user, "session_id": session_result["session_id"], "access_token": session_result["access_token"], "csrf_token": session_result["csrf_token"], } - - except (AuthenticationException, NotFoundException): + + except AuthenticationException, NotFoundException: raise except Exception as e: raise ServerException(f"Failed to reset password: {str(e)}") -async def validate_password_reset_token( - db: AsyncSession, - token: dict -) -> TokenValidationResponse: + +async def validate_password_reset_token(db: AsyncSession, token: dict) -> TokenValidationResponse: """Validate password reset token without consuming it""" try: user_id = token.get("sub") token_string = token.get("token") - + result = await db.execute( select(PasswordResetTokens).where( PasswordResetTokens.token == token_string, PasswordResetTokens.user_id == user_id, - PasswordResetTokens.is_used == False, - PasswordResetTokens.expires_at > datetime.now().astimezone() + PasswordResetTokens.is_used.is_(False), + PasswordResetTokens.expires_at > func.now(), ) ) token_record = result.scalar_one_or_none() - + if not token_record: raise AuthenticationException("Invalid or expired token") - - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + + result = await db.execute(select(Users).where(Users.id == user_id)) user = result.scalar_one_or_none() - + if not user or not user.status: raise AuthenticationException("User not found or account disabled") - - return TokenValidationResponse( - is_valid=True - ) - + + return TokenValidationResponse(is_valid=True) + except AuthenticationException: raise except Exception as e: raise ServerException(f"Token validation failed: {str(e)}") + async def forgot_password( db: AsyncSession, email: str, @@ -511,11 +497,14 @@ async def forgot_password( cooldown_key = f"password_reset_cooldown:{email}" remaining_seconds = await redis_client.ttl(cooldown_key) - + if remaining_seconds > 0: raise ValidationException( - f"Please wait {remaining_seconds} seconds before requesting another password reset email", - details={"cooldown_seconds": remaining_seconds} + ( + f"Please wait {remaining_seconds} seconds before requesting " + "another password reset email" + ), + details={"cooldown_seconds": remaining_seconds}, ) token_meta = await _request_password_reset_email(db, user) @@ -529,13 +518,13 @@ async def forgot_password( # Render email template with user name and app name user_name = f"{user.first_name} {user.last_name}".strip() app_name = settings.PROJECT_NAME - + email_content = PASSWORD_RESET_TEMPLATE.render( reset_url=reset_url, user_name=user_name, app_name=app_name, ) - + mailer.send_text( to_emails=[email], subject=email_content["subject"], @@ -544,15 +533,16 @@ async def forgot_password( ) # Set cooldown period in Redis - await redis_client.setex( - cooldown_key, - settings.PASSWORD_RESET_EMAIL_COOLDOWN_SECONDS, - "1" - ) + await redis_client.setex(cooldown_key, settings.PASSWORD_RESET_EMAIL_COOLDOWN_SECONDS, "1") return {**token_meta, "reset_url": reset_url} - - except (NotFoundException, AuthenticationException, SMTPNotConfiguredException, ValidationException): + + except ( + NotFoundException, + AuthenticationException, + SMTPNotConfiguredException, + ValidationException, + ): raise except Exception as e: raise ServerException(f"Failed to send password reset email: {str(e)}") @@ -564,52 +554,51 @@ async def get_password_reset_cooldown( ) -> dict: """ Get remaining cooldown time for password reset email. - + Returns: Dict with 'cooldown_seconds' (0 if no cooldown active) """ cooldown_key = f"password_reset_cooldown:{email}" remaining_seconds = await redis_client.ttl(cooldown_key) - + # TTL returns -1 if key exists but has no expiry, -2 if key doesn't exist if remaining_seconds < 0: remaining_seconds = 0 - + return {"cooldown_seconds": remaining_seconds} + async def _update_session_expiry(db: AsyncSession, session_id: str) -> None: """Update session expiry time in database""" try: - result = await db.execute( - select(UserSessions).where(UserSessions.id == session_id) - ) + result = await db.execute(select(UserSessions).where(UserSessions.id == session_id)) session = result.scalar_one_or_none() if session: - session.expires_at = datetime.now().astimezone() + timedelta(minutes=settings.SESSION_EXPIRE_MINUTES) + session.expires_at = datetime.now().astimezone() + timedelta( + minutes=settings.SESSION_EXPIRE_MINUTES + ) await db.commit() except Exception as e: raise ServerException(f"Failed to update session expiry in database: {e}") + async def _create_user(db: AsyncSession, user_data: UserRegister) -> Users: try: result = await db.execute( select(Users).where( - or_( - Users.email == user_data.email, - Users.pending_email == user_data.email - ) + or_(Users.email == user_data.email, Users.pending_email == user_data.email) ) ) existing_user = result.scalar_one_or_none() if existing_user: raise ConflictException("Email already exists") - + user = Users( first_name=user_data.first_name, last_name=user_data.last_name, email=user_data.email, phone=user_data.phone, - hash_password=await hash_password(user_data.password) + hash_password=await hash_password(user_data.password), ) db.add(user) await db.commit() @@ -620,18 +609,14 @@ async def _create_user(db: AsyncSession, user_data: UserRegister) -> Users: except Exception as e: raise ServerException(f"Failed to create user: {str(e)}") -async def _get_user_by_id(db: AsyncSession, user_id: str) -> Optional[Users]: - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + +async def _get_user_by_id(db: AsyncSession, user_id: str) -> Users | None: + result = await db.execute(select(Users).where(Users.id == user_id)) return result.scalar_one_or_none() + async def _create_user_session( - db: AsyncSession, - redis_client: redis.Redis, - user: Users, - ip_address: str, - user_agent: str + db: AsyncSession, redis_client: redis.Redis, user: Users, ip_address: str, user_agent: str ) -> SessionResult: try: session = UserSessions( @@ -639,18 +624,17 @@ async def _create_user_session( jwt_access_token="", ip_address=ip_address, user_agent=user_agent, - expires_at=datetime.now().astimezone() + timedelta(minutes=settings.SESSION_EXPIRE_MINUTES) + expires_at=datetime.now().astimezone() + + timedelta(minutes=settings.SESSION_EXPIRE_MINUTES), ) db.add(session) await db.commit() await db.refresh(session) session_id = session.id - access_token = await create_access_token(data={ - "sub": user.id, - "email": user.email, - "sid": session_id - }) + access_token = await create_access_token( + data={"sub": user.id, "email": user.email, "sid": session_id} + ) session.jwt_access_token = access_token await db.commit() @@ -665,15 +649,11 @@ async def _create_user_session( "created_at": datetime.now().astimezone().isoformat(), "last_activity": datetime.now().astimezone().isoformat(), } - - await redis_client.setex( - redis_key, - settings.SESSION_EXPIRE_MINUTES * 60, - str(session_data) - ) + + await redis_client.setex(redis_key, settings.SESSION_EXPIRE_MINUTES * 60, str(session_data)) csrf_token = await _create_csrf_token_for_session(redis_client, session_id) - + return { "session_id": session_id, "access_token": access_token, @@ -682,14 +662,15 @@ async def _create_user_session( except Exception as e: raise ServerException(f"Failed to create user session: {str(e)}") + async def _log_login_attempt( db: AsyncSession, - email: str, - ip_address: str, - user_agent: str, - is_success: bool, - user_id: Optional[str] = None, - failure_reason: Optional[str] = None + email: str, + ip_address: str, + user_agent: str, + is_success: bool, + user_id: str | None = None, + failure_reason: str | None = None, ) -> None: log = LoginLogs( user_id=user_id, @@ -697,11 +678,12 @@ async def _log_login_attempt( ip_address=ip_address, user_agent=user_agent, is_success=is_success, - failure_reason=failure_reason + failure_reason=failure_reason, ) db.add(log) await db.commit() + async def _get_user_by_email_for_password_reset(db: AsyncSession, email: str) -> Users: result = await db.execute(select(Users).where(Users.email == email)) user = result.scalar_one_or_none() @@ -711,6 +693,7 @@ async def _get_user_by_email_for_password_reset(db: AsyncSession, email: str) -> raise AuthenticationException("Account is disabled") return user + async def _request_password_reset_email( db: AsyncSession, user: Users, @@ -725,10 +708,7 @@ async def _request_password_reset_email( # Invalidate all previous unused tokens for this user await db.execute( update(PasswordResetTokens) - .where( - PasswordResetTokens.user_id == user.id, - PasswordResetTokens.is_used == False - ) + .where(PasswordResetTokens.user_id == user.id, PasswordResetTokens.is_used.is_(False)) .values(is_used=True) ) @@ -747,12 +727,9 @@ async def _request_password_reset_email( "user_id": user.id, } + async def verify_email( - db: AsyncSession, - redis_client: redis.Redis, - token: dict, - ip_address: str, - user_agent: str + db: AsyncSession, redis_client: redis.Redis, token: dict, ip_address: str, user_agent: str ) -> LoginResult: """Verify email using token and create session""" try: @@ -760,7 +737,7 @@ async def verify_email( email = token.get("email") verification_type = token.get("verification_type") token_string = token.get("token") - + # Verify token record exists and is valid result = await db.execute( select(EmailVerificationTokens).where( @@ -768,30 +745,28 @@ async def verify_email( EmailVerificationTokens.user_id == user_id, EmailVerificationTokens.email == email, EmailVerificationTokens.token_type == verification_type, - EmailVerificationTokens.is_used == False, - EmailVerificationTokens.expires_at > datetime.now().astimezone() + EmailVerificationTokens.is_used.is_(False), + EmailVerificationTokens.expires_at > func.now(), ) ) token_record = result.scalar_one_or_none() - + if not token_record: raise AuthenticationException("Invalid or expired token") - + # Get user - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + result = await db.execute(select(Users).where(Users.id == user_id)) user = result.scalar_one_or_none() - + if not user: raise NotFoundException("User not found") - + if not user.status: raise AuthenticationException("Account is disabled") - + # Mark token as used token_record.is_used = True - + if verification_type == "registration": # Mark email as verified user.email_verified = True @@ -799,37 +774,36 @@ async def verify_email( # Update email from pending_email to email if user.pending_email != email: raise AuthenticationException("Email mismatch") - + # Check if new email already exists result = await db.execute( select(Users).where(Users.email == email, Users.id != user_id) ) if result.scalar_one_or_none(): raise ConflictException("Email already exists") - + user.email = email user.pending_email = None user.email_verified = True - + # Create session for verified user - session_result = await _create_user_session( - db, redis_client, user, ip_address, user_agent - ) - + session_result = await _create_user_session(db, redis_client, user, ip_address, user_agent) + await db.commit() - + return { "user": user, "session_id": session_result["session_id"], "access_token": session_result["access_token"], "csrf_token": session_result["csrf_token"], } - - except (AuthenticationException, NotFoundException, ConflictException): + + except AuthenticationException, NotFoundException, ConflictException: raise except Exception as e: raise ServerException(f"Failed to verify email: {str(e)}") + async def resend_verification_email( db: AsyncSession, email: str, @@ -839,41 +813,47 @@ async def resend_verification_email( """Resend email verification""" try: user = await _get_user_by_email_for_password_reset(db, email) - + if not settings.SMTP_ENABLE or not getattr(mailer, "enabled", False): raise SMTPNotConfiguredException("SMTP is disabled") - + # Check cooldown cooldown_key = f"email_verification_cooldown:{email}" remaining_seconds = await redis_client.ttl(cooldown_key) - + if remaining_seconds > 0: raise ValidationException( - f"Please wait {remaining_seconds} seconds before requesting another verification email", - details={"cooldown_seconds": remaining_seconds} + ( + f"Please wait {remaining_seconds} seconds before requesting " + "another verification email" + ), + details={"cooldown_seconds": remaining_seconds}, ) - + # Determine verification type if user.email_verified: - # If email is already verified but there's a pending email, resend email change verification + # If email is already verified but there's a pending email, resend email change + # verification if user.pending_email: - token_meta = await _request_email_change_verification_email(db, user, user.pending_email) + token_meta = await _request_email_change_verification_email( + db, user, user.pending_email + ) verification_url = ( f"http{'s' if settings.SSL_ENABLE else ''}://" f"{settings.HOSTNAME}:{settings.FRONTEND_PORT}" f"/auth/verify-email?token={quote(token_meta['verification_token'], safe='')}" ) - + user_name = f"{user.first_name} {user.last_name}".strip() app_name = settings.PROJECT_NAME - + email_content = EMAIL_VERIFICATION_TEMPLATE.render( verification_url=verification_url, user_name=user_name, app_name=app_name, expire_minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES, ) - + mailer.send_text( to_emails=[user.pending_email], subject=email_content["subject"], @@ -890,34 +870,35 @@ async def resend_verification_email( f"{settings.HOSTNAME}:{settings.FRONTEND_PORT}" f"/auth/verify-email?token={quote(token_meta['verification_token'], safe='')}" ) - + user_name = f"{user.first_name} {user.last_name}".strip() app_name = settings.PROJECT_NAME - + email_content = EMAIL_VERIFICATION_TEMPLATE.render( verification_url=verification_url, user_name=user_name, app_name=app_name, expire_minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES, ) - + mailer.send_text( to_emails=[email], subject=email_content["subject"], body=email_content["body"], html_body=email_content.get("html_body"), ) - + # Set cooldown - await redis_client.setex( - cooldown_key, - settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, - "1" - ) - + await redis_client.setex(cooldown_key, settings.EMAIL_VERIFICATION_COOLDOWN_SECONDS, "1") + return {"message": "Verification email sent"} - - except (NotFoundException, AuthenticationException, SMTPNotConfiguredException, ValidationException): + + except ( + NotFoundException, + AuthenticationException, + SMTPNotConfiguredException, + ValidationException, + ): raise except Exception as e: raise ServerException(f"Failed to send verification email: {str(e)}") @@ -936,32 +917,31 @@ async def _create_csrf_token_for_session( await redis_client.setex(_csrf_redis_key(session_id), ttl, csrf_token) return csrf_token + async def _send_registration_verification_email( - db: AsyncSession, - mailer: SMTPMailer, - user: Users + db: AsyncSession, mailer: SMTPMailer, user: Users ) -> None: """Send registration verification email""" if not settings.SMTP_ENABLE or not getattr(mailer, "enabled", False): return - + token_meta = await _request_registration_verification_email(db, user) verification_url = ( f"http{'s' if settings.SSL_ENABLE else ''}://" f"{settings.HOSTNAME}:{settings.FRONTEND_PORT}" f"/auth/verify-email?token={quote(token_meta['verification_token'], safe='')}" ) - + user_name = f"{user.first_name} {user.last_name}".strip() app_name = settings.PROJECT_NAME - + email_content = EMAIL_VERIFICATION_TEMPLATE.render( verification_url=verification_url, user_name=user_name, app_name=app_name, expire_minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES, ) - + mailer.send_text( to_emails=[user.email], subject=email_content["subject"], @@ -969,6 +949,7 @@ async def _send_registration_verification_email( html_body=email_content.get("html_body"), ) + async def _request_registration_verification_email( db: AsyncSession, user: Users, @@ -976,18 +957,18 @@ async def _request_registration_verification_email( """Create a registration verification token record""" now = datetime.now().astimezone() expires_at = now + timedelta(minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES) - + # Invalidate all previous unused registration tokens for this user await db.execute( update(EmailVerificationTokens) .where( EmailVerificationTokens.user_id == user.id, EmailVerificationTokens.token_type == "registration", - EmailVerificationTokens.is_used == False + EmailVerificationTokens.is_used.is_(False), ) .values(is_used=True) ) - + verification_token = await create_email_verification_token(user.id, user.email, "registration") token_record = EmailVerificationTokens( user_id=user.id, @@ -998,13 +979,14 @@ async def _request_registration_verification_email( ) db.add(token_record) await db.commit() - + return { "verification_token": verification_token, "expires_at": expires_at, "user_id": user.id, } + async def _request_email_change_verification_email( db: AsyncSession, user: Users, @@ -1013,18 +995,18 @@ async def _request_email_change_verification_email( """Create an email change verification token record""" now = datetime.now().astimezone() expires_at = now + timedelta(minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES) - + # Invalidate all previous unused email_change tokens for this user await db.execute( update(EmailVerificationTokens) .where( EmailVerificationTokens.user_id == user.id, EmailVerificationTokens.token_type == "email_change", - EmailVerificationTokens.is_used == False + EmailVerificationTokens.is_used.is_(False), ) .values(is_used=True) ) - + verification_token = await create_email_verification_token(user.id, new_email, "email_change") token_record = EmailVerificationTokens( user_id=user.id, @@ -1035,9 +1017,9 @@ async def _request_email_change_verification_email( ) db.add(token_record) await db.commit() - + return { "verification_token": verification_token, "expires_at": expires_at, "user_id": user.id, - } \ No newline at end of file + } diff --git a/backend/api/debug/controller.py b/backend/api/debug/controller.py index f86fd0a..7e6c009 100644 --- a/backend/api/debug/controller.py +++ b/backend/api/debug/controller.py @@ -1,22 +1,29 @@ import logging -from utils.response import parse_responses, APIResponse -from .services import get_ip_debug_info, clear_blocked_ips -from .schema import IPDebugResponse, ClearBlockedIPsResponse -from fastapi import APIRouter, Request, HTTPException + +from fastapi import APIRouter, HTTPException, Request + +from utils.response import APIResponse, parse_responses + +from .schema import ClearBlockedIPsResponse, IPDebugResponse +from .services import clear_blocked_ips, get_ip_debug_info logger = logging.getLogger("debug") router = APIRouter(tags=["Debug"]) -@router.get("/test-ip", - response_model=APIResponse[IPDebugResponse], - summary="Test IP detection", - responses=parse_responses({ - 200: ("IP detection successful", IPDebugResponse), - 400: ("Invalid IP address", None), - 429: ("Too Many Requests", None), - 500: ("Internal Server Error", None), - }), + +@router.get( + "/test-ip", + response_model=APIResponse[IPDebugResponse], + summary="Test IP detection", + responses=parse_responses( + { + 200: ("IP detection successful", IPDebugResponse), + 400: ("Invalid IP address", None), + 429: ("Too Many Requests", None), + 500: ("Internal Server Error", None), + } + ), ) async def test_ip_detection(request: Request): try: @@ -25,22 +32,21 @@ async def test_ip_detection(request: Request): except Exception: raise HTTPException(status_code=500) + @router.delete( "/clear-blocked-ip", summary="Clear all blocked IPs in Redis", response_model=APIResponse[ClearBlockedIPsResponse], - responses=parse_responses({ - 200: ("Blocked IPs cleared successfully", ClearBlockedIPsResponse), - 500: ("Internal Server Error", None), - }) + responses=parse_responses( + { + 200: ("Blocked IPs cleared successfully", ClearBlockedIPsResponse), + 500: ("Internal Server Error", None), + } + ), ) async def clear_blocked_ips_api(): try: result = await clear_blocked_ips() - return APIResponse( - code=200, - message="Blocked IPs cleared successfully", - data=result - ) + return APIResponse(code=200, message="Blocked IPs cleared successfully", data=result) except Exception: raise HTTPException(status_code=500) diff --git a/backend/api/debug/schema.py b/backend/api/debug/schema.py index 7529757..70770b1 100644 --- a/backend/api/debug/schema.py +++ b/backend/api/debug/schema.py @@ -1,12 +1,13 @@ -from typing import Optional, List from pydantic import BaseModel, Field + class IPDebugResponse(BaseModel): - client_host: Optional[str] = Field(..., description="Client host") - x_forwarded_for: Optional[str] = Field(..., description="X-Forwarded-For") - x_real_ip: Optional[str] = Field(..., description="X-Real-IP") - detected_real_ip: Optional[str] = Field(..., description="Detected real IP") + client_host: str | None = Field(..., description="Client host") + x_forwarded_for: str | None = Field(..., description="X-Forwarded-For") + x_real_ip: str | None = Field(..., description="X-Real-IP") + detected_real_ip: str | None = Field(..., description="Detected real IP") + class ClearBlockedIPsResponse(BaseModel): - cleared_ips: List[str] = Field(..., description="Cleared blocked IPs") - count: int = Field(..., description="Number of cleared IPs") \ No newline at end of file + cleared_ips: list[str] = Field(..., description="Cleared blocked IPs") + count: int = Field(..., description="Number of cleared IPs") diff --git a/backend/api/debug/services.py b/backend/api/debug/services.py index 0b711bb..168e267 100644 --- a/backend/api/debug/services.py +++ b/backend/api/debug/services.py @@ -1,9 +1,12 @@ -from fastapi import Request -from utils import get_real_ip import redis.asyncio as aioredis +from fastapi import Request + from core.config import settings +from utils import get_real_ip from utils.custom_exception import ServerException -from .schema import IPDebugResponse, ClearBlockedIPsResponse + +from .schema import ClearBlockedIPsResponse, IPDebugResponse + async def get_ip_debug_info(request: Request) -> IPDebugResponse: try: @@ -11,11 +14,12 @@ async def get_ip_debug_info(request: Request) -> IPDebugResponse: client_host=request.client.host if request.client else None, x_forwarded_for=request.headers.get("x-forwarded-for"), x_real_ip=request.headers.get("x-real-ip"), - detected_real_ip=get_real_ip(request) + detected_real_ip=get_real_ip(request), ) except Exception as e: raise ServerException(f"Failed to get IP debug info: {e}") + async def clear_blocked_ips() -> ClearBlockedIPsResponse: try: redis = await aioredis.from_url(settings.REDIS_URL, encoding="utf-8", decode_responses=True) @@ -27,4 +31,4 @@ async def clear_blocked_ips() -> ClearBlockedIPsResponse: cleared.append(key.replace("block:", "")) return ClearBlockedIPsResponse(cleared_ips=cleared, count=len(cleared)) except Exception as e: - raise ServerException(f"Failed to clear blocked IPs: {e}") \ No newline at end of file + raise ServerException(f"Failed to clear blocked IPs: {e}") diff --git a/backend/api/roles/controller.py b/backend/api/roles/controller.py index 85cdb7c..0e5822f 100644 --- a/backend/api/roles/controller.py +++ b/backend/api/roles/controller.py @@ -1,41 +1,56 @@ +from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + from core.dependencies import get_db -from core.security import verify_token -from core.rbac import require_permission from core.permissions import Permission -from sqlalchemy.ext.asyncio import AsyncSession -from utils.response import APIResponse, parse_responses, common_responses -from fastapi import APIRouter, Depends, HTTPException, Request, Path, Response -from utils.custom_exception import NotFoundException, ConflictException, ServerException, AuthorizationException -from .services import ( - get_all_roles, create_role, update_role, delete_role, - get_role_attribute_mapping, update_role_attribute_mapping, - check_user_permissions +from core.rbac import require_permission +from core.security import verify_token +from utils.custom_exception import ( + AuthorizationException, + ConflictException, + NotFoundException, + ServerException, ) +from utils.response import APIResponse, common_responses, parse_responses + from .schema import ( - RoleResponse, RoleCreate, RoleUpdate, RolesListResponse, - RoleAttributesMapping, RoleAttributeMappingBatchResponse, - RoleAttributesGroupedResponse, PermissionCheckResponse, - role_attributes_success_response_example, role_attributes_partial_response_example, - role_attributes_failed_response_example + RoleAttributeMappingBatchResponse, + RoleAttributesGroupedResponse, + RoleAttributesMapping, + RoleCreate, + RoleResponse, + RolesListResponse, + RoleUpdate, + role_attributes_failed_response_example, + role_attributes_partial_response_example, + role_attributes_success_response_example, +) +from .services import ( + check_user_permissions, + create_role, + delete_role, + get_all_roles, + get_role_attribute_mapping, + update_role, + update_role_attribute_mapping, ) router = APIRouter(tags=["Roles"]) + @router.get( "", response_model=APIResponse[RolesListResponse], response_model_exclude_none=True, summary="Get all custom roles", - responses=parse_responses({ - 200: ("Successfully retrieved roles", RolesListResponse) - }, common_responses) + responses=parse_responses( + {200: ("Successfully retrieved roles", RolesListResponse)}, common_responses + ), ) @require_permission([Permission.VIEW_ROLES, Permission.MANAGE_ROLES]) async def get_roles( - request: Request, - token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + request: Request, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db) ): """Get all custom roles""" try: @@ -44,22 +59,23 @@ async def get_roles( except Exception: raise HTTPException(status_code=500) + @router.post( "", response_model=APIResponse[RoleResponse], response_model_exclude_none=True, summary="Create new role", - responses=parse_responses({ - 200: ("Role created successfully", RoleResponse), - 409: ("Role name already exists", None) - }, common_responses) + responses=parse_responses( + {200: ("Role created successfully", RoleResponse), 409: ("Role name already exists", None)}, + common_responses, + ), ) @require_permission([Permission.MANAGE_ROLES]) async def create_role_api( role_data: RoleCreate, request: Request, token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Create a new role""" try: @@ -72,16 +88,20 @@ async def create_role_api( except Exception: raise HTTPException(status_code=500) + @router.put( "/{role_id}", response_model=APIResponse[RoleResponse], response_model_exclude_none=True, summary="Update role info", - responses=parse_responses({ - 200: ("Role updated successfully", RoleResponse), - 404: ("Role not found", None), - 409: ("Role name already exists", None) - }, common_responses) + responses=parse_responses( + { + 200: ("Role updated successfully", RoleResponse), + 404: ("Role not found", None), + 409: ("Role name already exists", None), + }, + common_responses, + ), ) @require_permission([Permission.MANAGE_ROLES]) async def update_role_api( @@ -89,7 +109,7 @@ async def update_role_api( role_data: RoleUpdate = None, request: Request = None, token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Update role information""" try: @@ -104,23 +124,27 @@ async def update_role_api( except ServerException: raise HTTPException(status_code=500) + @router.delete( "/{role_id}", response_model=APIResponse[dict], response_model_exclude_none=True, summary="Delete role", - responses=parse_responses({ - 200: ("Role deleted successfully", None), - 404: ("Role not found", None), - 409: ("Cannot delete role that is assigned to users", None) - }, common_responses) + responses=parse_responses( + { + 200: ("Role deleted successfully", None), + 404: ("Role not found", None), + 409: ("Cannot delete role that is assigned to users", None), + }, + common_responses, + ), ) @require_permission([Permission.MANAGE_ROLES]) async def delete_role_api( role_id: str = Path(..., description="Role ID"), request: Request = None, token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Delete a role""" try: @@ -135,30 +159,38 @@ async def delete_role_api( except Exception: raise HTTPException(status_code=500) + @router.get( "/{role_id}/attributes", response_model=APIResponse[RoleAttributesGroupedResponse], response_model_exclude_none=True, summary="Get role attributes mapping", - responses=parse_responses({ - 200: ("Successfully retrieved role attributes mapping", RoleAttributesGroupedResponse, RoleAttributesGroupedResponse.get_example_response()), - 404: ("Role not found", None) - }, common_responses) + responses=parse_responses( + { + 200: ( + "Successfully retrieved role attributes mapping", + RoleAttributesGroupedResponse, + RoleAttributesGroupedResponse.get_example_response(), + ), + 404: ("Role not found", None), + }, + common_responses, + ), ) @require_permission([Permission.VIEW_ROLES, Permission.MANAGE_ROLES]) async def get_role_attribute_mapping_api( role_id: str = Path(..., description="Role ID"), request: Request = None, token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Get role attributes mapping with all available attributes""" try: attributes_mapping = await get_role_attribute_mapping(db, role_id) return APIResponse( - code=200, - message="Successfully retrieved role attributes mapping", - data=attributes_mapping + code=200, + message="Successfully retrieved role attributes mapping", + data=attributes_mapping, ) except NotFoundException: raise HTTPException(status_code=404, detail="Role not found") @@ -167,16 +199,32 @@ async def get_role_attribute_mapping_api( except Exception: raise HTTPException(status_code=500) + @router.put( "/{role_id}/attributes", response_model=APIResponse[RoleAttributeMappingBatchResponse], response_model_exclude_none=True, summary="Update role attributes", - responses=parse_responses({ - 200: ("All role attributes processed successfully", RoleAttributeMappingBatchResponse, role_attributes_success_response_example), - 207: ("Role attributes processed with partial success", RoleAttributeMappingBatchResponse, role_attributes_partial_response_example), - 400: ("All role attributes failed to process", RoleAttributeMappingBatchResponse, role_attributes_failed_response_example) - }, common_responses) + responses=parse_responses( + { + 200: ( + "All role attributes processed successfully", + RoleAttributeMappingBatchResponse, + role_attributes_success_response_example, + ), + 207: ( + "Role attributes processed with partial success", + RoleAttributeMappingBatchResponse, + role_attributes_partial_response_example, + ), + 400: ( + "All role attributes failed to process", + RoleAttributeMappingBatchResponse, + role_attributes_failed_response_example, + ), + }, + common_responses, + ), ) @require_permission([Permission.MANAGE_ROLES]) async def update_role_attribute_mapping_api( @@ -184,45 +232,37 @@ async def update_role_attribute_mapping_api( attributes_data: RoleAttributesMapping = None, request: Request = None, token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Update role attributes with batch processing results""" try: batch_result = await update_role_attribute_mapping( db, role_id, attributes_data.attributes, actor_user_id=token["sub"] ) - + # Determine response code based on results if batch_result.failed_count == 0: # All successful return APIResponse( - code=200, - message="All role attributes processed successfully", - data=batch_result + code=200, message="All role attributes processed successfully", data=batch_result ) elif batch_result.success_count == 0: # All failed - return 400 status code response = APIResponse( - code=400, - message="All role attributes failed to process", - data=batch_result + code=400, message="All role attributes failed to process", data=batch_result ) return Response( - content=response.model_dump_json(), - status_code=400, - media_type="application/json" + content=response.model_dump_json(), status_code=400, media_type="application/json" ) else: # Partial success - return 207 status code response = APIResponse( - code=207, - message="Role attributes processed with partial success", - data=batch_result + code=207, + message="Role attributes processed with partial success", + data=batch_result, ) return Response( - content=response.model_dump_json(), - status_code=207, - media_type="application/json" + content=response.model_dump_json(), status_code=207, media_type="application/json" ) except NotFoundException: raise HTTPException(status_code=404, detail="Role not found") @@ -231,27 +271,33 @@ async def update_role_attribute_mapping_api( except Exception: raise HTTPException(status_code=500) + @router.get( "/permissions", response_model=APIResponse[PermissionCheckResponse], response_model_exclude_none=True, summary="Get current user permissions", - responses=parse_responses({ - 200: ("User permissions retrieved", PermissionCheckResponse, PermissionCheckResponse.get_example_response()) - }, common_responses) + responses=parse_responses( + { + 200: ( + "User permissions retrieved", + PermissionCheckResponse, + PermissionCheckResponse.get_example_response(), + ) + }, + common_responses, + ), ) async def get_user_permissions_api( - request: Request = None, - token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + request: Request = None, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db) ): """Get all permissions for the current user""" try: user_id = token.get("sub") result = await check_user_permissions(db, user_id, None) - + return APIResponse(code=200, message="User permissions retrieved", data=result) except HTTPException: raise except Exception: - raise HTTPException(status_code=500) \ No newline at end of file + raise HTTPException(status_code=500) diff --git a/backend/api/roles/schema.py b/backend/api/roles/schema.py index b4ea529..de774fd 100644 --- a/backend/api/roles/schema.py +++ b/backend/api/roles/schema.py @@ -1,23 +1,26 @@ from pydantic import BaseModel, Field -from typing import Optional, Dict, List + from core.config import settings + class RoleResponse(BaseModel): id: str = Field(..., description="Role ID") name: str = Field(..., description="Role name") - description: Optional[str] = Field(None, description="Role description") + description: str | None = Field(None, description="Role description") level: int = Field(..., description="Role privilege level (higher is more privileged)") + class RolesListResponse(BaseModel): - roles: List[RoleResponse] = Field(..., description="List of roles") + roles: list[RoleResponse] = Field(..., description="List of roles") actor_level: int = Field(..., description="Current user's role level") - actor_role_id: Optional[str] = Field( + actor_role_id: str | None = Field( None, description="Current user's primary role ID (cannot self-edit/delete)" ) + class RoleCreate(BaseModel): name: str = Field(..., min_length=1, max_length=100, description="Role name") - description: Optional[str] = Field(None, max_length=500, description="Role description") + description: str | None = Field(None, max_length=500, description="Role description") level: int = Field( ..., ge=1, @@ -25,41 +28,34 @@ class RoleCreate(BaseModel): description="Role privilege level (must be lower than the actor's level)", ) + class RoleUpdate(BaseModel): - name: Optional[str] = Field(None, min_length=1, max_length=100, description="Role name") - description: Optional[str] = Field(None, max_length=500, description="Role description") - level: Optional[int] = Field( + name: str | None = Field(None, min_length=1, max_length=100, description="Role name") + description: str | None = Field(None, max_length=500, description="Role description") + level: int | None = Field( None, ge=1, le=settings.MAX_CUSTOM_ROLE_LEVEL, description="Role privilege level (must be lower than the actor's level)", ) + class RoleAttributesMapping(BaseModel): - attributes: Dict[str, bool] = Field( - ..., + attributes: dict[str, bool] = Field( + ..., description="Role attributes mapping (attribute_name: value)", - example={ - "view-users": True, - "manage-users": False, - "view-roles": True - } + example={"view-users": True, "manage-users": False, "view-roles": True}, ) - + @classmethod def get_example_response(cls): return { "code": 200, "message": "Role attributes mapping example", - "data": { - "attributes": { - "view-users": True, - "manage-users": False, - "view-roles": True - } - } + "data": {"attributes": {"view-users": True, "manage-users": False, "view-roles": True}}, } + class RoleAttributeDetail(BaseModel): name: str = Field(..., description="Attribute name", example="view-users") value: bool = Field(..., description="Whether the role has this attribute", example=True) @@ -67,11 +63,15 @@ class RoleAttributeDetail(BaseModel): class RoleAttributesGroup(BaseModel): group: str = Field(..., description="Top-level group key", example="user-role-management") - categories: Dict[str, List[RoleAttributeDetail]] = Field(..., description="Categories inside this group (category -> attributes)") + categories: dict[str, list[RoleAttributeDetail]] = Field( + ..., description="Categories inside this group (category -> attributes)" + ) class RoleAttributesGroupedResponse(BaseModel): - groups: List[RoleAttributesGroup] = Field(..., description="Role attributes grouped by group and category") + groups: list[RoleAttributesGroup] = Field( + ..., description="Role attributes grouped by group and category" + ) @classmethod def get_example_response(cls): @@ -84,58 +84,57 @@ def get_example_response(cls): "group": "user-role-management", "categories": { "user": [ - { - "name": "view-users", - "value": True - }, - { - "name": "manage-users", - "value": False - } + {"name": "view-users", "value": True}, + {"name": "manage-users", "value": False}, ], - "role": [ - { - "name": "view-roles", - "value": True - } - ] - } + "role": [{"name": "view-roles", "value": True}], + }, } ] - } + }, } + class AttributeMappingResult(BaseModel): attribute_id: str = Field(..., description="Attribute ID") - status: str = Field(..., description="Processing status: success, failed", pattern="^(success|failed)$") + status: str = Field( + ..., description="Processing status: success, failed", pattern="^(success|failed)$" + ) message: str = Field(..., description="Result message") + class RoleAttributeMappingBatchResponse(BaseModel): - results: List[AttributeMappingResult] = Field(..., description="Individual attribute mapping results") + results: list[AttributeMappingResult] = Field( + ..., description="Individual attribute mapping results" + ) total_attributes: int = Field(..., description="Total number of attributes processed") success_count: int = Field(..., description="Number of successfully processed attributes") failed_count: int = Field(..., description="Number of failed attributes") + class PermissionCheckRequest(BaseModel): - attributes: Optional[List[str]] = Field( - None, + attributes: list[str] | None = Field( + None, min_items=1, - description="List of permission attributes to check. If not provided, returns all user permissions.", - example=["view-users", "manage-roles"] + description=( + "List of permission attributes to check. If not provided, returns all user permissions." + ), + example=["view-users", "manage-roles"], ) + class PermissionCheckResponse(BaseModel): - permissions: Dict[str, bool] = Field( - ..., + permissions: dict[str, bool] = Field( + ..., description="Permission check results (attribute_name: has_permission)", example={ "view-users": True, "manage-users": False, "view-roles": True, - "manage-roles": False - } + "manage-roles": False, + }, ) - + @classmethod def get_example_response(cls): return { @@ -146,31 +145,24 @@ def get_example_response(cls): "view-users": True, "manage-users": False, "view-roles": True, - "manage-roles": False + "manage-roles": False, } - } + }, } + role_attributes_success_response_example = { "code": 200, "message": "All role attributes processed successfully", "data": { "results": [ - { - "attribute_id": "attr-001", - "status": "success", - "message": "Updated successfully" - }, - { - "attribute_id": "attr-002", - "status": "success", - "message": "Updated successfully" - } + {"attribute_id": "attr-001", "status": "success", "message": "Updated successfully"}, + {"attribute_id": "attr-002", "status": "success", "message": "Updated successfully"}, ], "total_attributes": 2, "success_count": 2, - "failed_count": 0 - } + "failed_count": 0, + }, } role_attributes_partial_response_example = { @@ -178,21 +170,13 @@ def get_example_response(cls): "message": "Role attributes processed with partial success", "data": { "results": [ - { - "attribute_id": "attr-001", - "status": "success", - "message": "Updated successfully" - }, - { - "attribute_id": "attr-002", - "status": "failed", - "message": "Invalid attribute ID" - } + {"attribute_id": "attr-001", "status": "success", "message": "Updated successfully"}, + {"attribute_id": "attr-002", "status": "failed", "message": "Invalid attribute ID"}, ], "total_attributes": 2, "success_count": 1, - "failed_count": 1 - } + "failed_count": 1, + }, } role_attributes_failed_response_example = { @@ -203,16 +187,16 @@ def get_example_response(cls): { "attribute_id": "invalid-attr-001", "status": "failed", - "message": "Invalid attribute ID" + "message": "Invalid attribute ID", }, { "attribute_id": "invalid-attr-002", - "status": "failed", - "message": "Invalid attribute ID" - } + "status": "failed", + "message": "Invalid attribute ID", + }, ], "total_attributes": 2, "success_count": 0, - "failed_count": 2 - } -} \ No newline at end of file + "failed_count": 2, + }, +} diff --git a/backend/api/roles/services.py b/backend/api/roles/services.py index 074fe9d..9f96a53 100644 --- a/backend/api/roles/services.py +++ b/backend/api/roles/services.py @@ -1,6 +1,6 @@ -from typing import Dict, List, Optional -from models.roles import Roles -from models.role_mapper import RoleMapper +from sqlalchemy import and_, delete, func, select +from sqlalchemy.ext.asyncio import AsyncSession + from core.config import settings from core.rbac import ( check_user_has_super_role, @@ -9,28 +9,37 @@ is_super_admin_role_name, user_has_role, ) -from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy import select, func, delete, and_ from models.role_attributes import RoleAttributes from models.role_attributes_mapper import RoleAttributesMapper +from models.role_mapper import RoleMapper +from models.roles import Roles from utils.custom_exception import ( - ServerException, + AuthorizationException, ConflictException, NotFoundException, - AuthorizationException, + ServerException, ) + from .schema import ( - RoleResponse, RoleCreate, RoleUpdate, RolesListResponse, - RoleAttributeMappingBatchResponse, AttributeMappingResult, - RoleAttributesGroupedResponse, RoleAttributesGroup, RoleAttributeDetail, - PermissionCheckResponse + AttributeMappingResult, + PermissionCheckResponse, + RoleAttributeDetail, + RoleAttributeMappingBatchResponse, + RoleAttributesGroup, + RoleAttributesGroupedResponse, + RoleCreate, + RoleResponse, + RolesListResponse, + RoleUpdate, ) + def _ensure_role_is_mutable(role: Roles) -> None: """Block mutations of the system super-admin role.""" if is_super_admin_role_name(role.name): raise AuthorizationException("Cannot modify the system super-admin role") + def _to_role_response(role: Roles) -> RoleResponse: return RoleResponse( id=role.id, @@ -39,12 +48,13 @@ def _to_role_response(role: Roles) -> RoleResponse: level=role.level, ) + async def _assert_can_manage_role_level( db: AsyncSession, actor_user_id: str, *, target_level: int, - existing_role_level: Optional[int] = None, + existing_role_level: int | None = None, ) -> int: """ Allow managing roles at the same level or lower. @@ -52,15 +62,12 @@ async def _assert_can_manage_role_level( """ actor_level = await get_user_role_level(actor_user_id, db) if existing_role_level is not None and existing_role_level > actor_level: - raise AuthorizationException( - "Cannot manage a role with higher level than your own" - ) + raise AuthorizationException("Cannot manage a role with higher level than your own") if target_level > actor_level: - raise AuthorizationException( - "Cannot assign a role level higher than your own" - ) + raise AuthorizationException("Cannot assign a role level higher than your own") return actor_level + async def _assert_not_own_role( db: AsyncSession, actor_user_id: str, @@ -70,6 +77,7 @@ async def _assert_not_own_role( if await user_has_role(actor_user_id, role_id, db): raise AuthorizationException("Cannot modify or delete your own role") + async def get_all_roles(db: AsyncSession, actor_user_id: str) -> RolesListResponse: """Get assignable roles (excludes the system super-admin role).""" try: @@ -93,6 +101,7 @@ async def get_all_roles(db: AsyncSession, actor_user_id: str) -> RolesListRespon except Exception as e: raise ServerException(f"Failed to retrieve roles: {str(e)}") + async def create_role( db: AsyncSession, role_data: RoleCreate, @@ -109,9 +118,7 @@ async def create_role( target_level=role_data.level, ) - existing_role = await db.execute( - select(Roles).where(Roles.name == role_data.name) - ) + existing_role = await db.execute(select(Roles).where(Roles.name == role_data.name)) if existing_role.scalar_one_or_none(): raise ConflictException("Role name already exists") @@ -126,11 +133,12 @@ async def create_role( return _to_role_response(role) - except (ConflictException, AuthorizationException): + except ConflictException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to create role: {str(e)}") + async def update_role( db: AsyncSession, role_id: str, @@ -139,9 +147,7 @@ async def update_role( ) -> RoleResponse: """Update role information""" try: - role_result = await db.execute( - select(Roles).where(Roles.id == role_id) - ) + role_result = await db.execute(select(Roles).where(Roles.id == role_id)) role = role_result.scalar_one_or_none() if not role: raise NotFoundException("Role not found") @@ -177,17 +183,16 @@ async def update_role( return _to_role_response(role) - except (ConflictException, NotFoundException, AuthorizationException): + except ConflictException, NotFoundException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to update role: {str(e)}") + async def delete_role(db: AsyncSession, role_id: str, actor_user_id: str) -> bool: """Delete a role""" try: - role_result = await db.execute( - select(Roles).where(Roles.id == role_id) - ) + role_result = await db.execute(select(Roles).where(Roles.id == role_id)) role = role_result.scalar_one_or_none() if not role: raise NotFoundException("Role not found") @@ -211,24 +216,23 @@ async def delete_role(db: AsyncSession, role_id: str, actor_user_id: str) -> boo delete(RoleAttributesMapper).where(RoleAttributesMapper.role_id == role_id) ) - await db.execute( - delete(Roles).where(Roles.id == role_id) - ) + await db.execute(delete(Roles).where(Roles.id == role_id)) await db.commit() return True - except (ConflictException, NotFoundException, AuthorizationException): + except ConflictException, NotFoundException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to delete role: {str(e)}") -async def get_role_attribute_mapping(db: AsyncSession, role_id: str) -> RoleAttributesGroupedResponse: + +async def get_role_attribute_mapping( + db: AsyncSession, role_id: str +) -> RoleAttributesGroupedResponse: """Get role attributes mapping grouped by group and category (left join).""" try: - role_result = await db.execute( - select(Roles).where(Roles.id == role_id) - ) + role_result = await db.execute(select(Roles).where(Roles.id == role_id)) role = role_result.scalar_one_or_none() if not role: raise NotFoundException("Role not found") @@ -236,25 +240,28 @@ async def get_role_attribute_mapping(db: AsyncSession, role_id: str) -> RoleAttr _ensure_role_is_mutable(role) # Use LEFT JOIN to get all attributes and their mappings - query = select( - RoleAttributes.name, - RoleAttributes.group, - RoleAttributes.category, - RoleAttributesMapper.value - ).select_from( - RoleAttributes - ).outerjoin( - RoleAttributesMapper, - and_( - RoleAttributes.id == RoleAttributesMapper.attributes_id, - RoleAttributesMapper.role_id == role_id + query = ( + select( + RoleAttributes.name, + RoleAttributes.group, + RoleAttributes.category, + RoleAttributesMapper.value, + ) + .select_from(RoleAttributes) + .outerjoin( + RoleAttributesMapper, + and_( + RoleAttributes.id == RoleAttributesMapper.attributes_id, + RoleAttributesMapper.role_id == role_id, + ), ) - ).order_by(RoleAttributes.id) + .order_by(RoleAttributes.id) + ) result = await db.execute(query) rows = result.all() - grouped: Dict[str, Dict[str, List[RoleAttributeDetail]]] = {} + grouped: dict[str, dict[str, list[RoleAttributeDetail]]] = {} for row in rows: group = row.group or "default" category = row.category or "uncategorized" @@ -276,22 +283,21 @@ async def get_role_attribute_mapping(db: AsyncSession, role_id: str) -> RoleAttr return RoleAttributesGroupedResponse(groups=groups) - except (NotFoundException, AuthorizationException): + except NotFoundException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to get role attributes: {str(e)}") + async def update_role_attribute_mapping( db: AsyncSession, role_id: str, - attributes_data: Dict[str, bool], + attributes_data: dict[str, bool], actor_user_id: str, ) -> RoleAttributeMappingBatchResponse: """Batch update role and attributes mapping with detailed results""" try: - role_result = await db.execute( - select(Roles).where(Roles.id == role_id) - ) + role_result = await db.execute(select(Roles).where(Roles.id == role_id)) role = role_result.scalar_one_or_none() if not role: raise NotFoundException("Role not found") @@ -316,7 +322,9 @@ async def update_role_attribute_mapping( if attribute_names: existing_attributes = await db.execute( - select(RoleAttributes.id, RoleAttributes.name).where(RoleAttributes.name.in_(attribute_names)) + select(RoleAttributes.id, RoleAttributes.name).where( + RoleAttributes.name.in_(attribute_names) + ) ) for row in existing_attributes: name_to_id_map[row.name] = row.id @@ -324,11 +332,13 @@ async def update_role_attribute_mapping( # Handle invalid attribute names invalid_names = set(attribute_names) - set(name_to_id_map.keys()) for invalid_name in invalid_names: - results.append(AttributeMappingResult( - attribute_id=invalid_name, # Keep name for error reporting - status="failed", - message="Invalid attribute name" - )) + results.append( + AttributeMappingResult( + attribute_id=invalid_name, # Keep name for error reporting + status="failed", + message="Invalid attribute name", + ) + ) failed_count += 1 # Process valid attributes @@ -345,7 +355,7 @@ async def update_role_attribute_mapping( select(RoleAttributesMapper).where( and_( RoleAttributesMapper.role_id == role_id, - RoleAttributesMapper.attributes_id == attribute_id + RoleAttributesMapper.attributes_id == attribute_id, ) ) ) @@ -357,25 +367,27 @@ async def update_role_attribute_mapping( else: # Create new mapping new_mapping = RoleAttributesMapper( - role_id=role_id, - attributes_id=attribute_id, - value=value + role_id=role_id, attributes_id=attribute_id, value=value ) db.add(new_mapping) - results.append(AttributeMappingResult( - attribute_id=attribute_name, - status="success", - message="Updated successfully" - )) + results.append( + AttributeMappingResult( + attribute_id=attribute_name, + status="success", + message="Updated successfully", + ) + ) success_count += 1 except Exception as e: - results.append(AttributeMappingResult( - attribute_id=attribute_name, - status="failed", - message=f"Failed to process: {str(e)}" - )) + results.append( + AttributeMappingResult( + attribute_id=attribute_name, + status="failed", + message=f"Failed to process: {str(e)}", + ) + ) failed_count += 1 await db.commit() @@ -384,18 +396,17 @@ async def update_role_attribute_mapping( results=results, total_attributes=len(attributes_data), success_count=success_count, - failed_count=failed_count + failed_count=failed_count, ) - except (NotFoundException, AuthorizationException): + except NotFoundException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to update role attributes mapping: {str(e)}") + async def check_user_permissions( - db: AsyncSession, - user_id: str, - required_attributes: List[str] = None + db: AsyncSession, user_id: str, required_attributes: list[str] = None ) -> PermissionCheckResponse: """Check if user has required permission attributes.""" try: @@ -416,13 +427,14 @@ async def check_user_permissions( user_attributes_set = set() if user_role_id: - attributes_query = select(RoleAttributes.name).join( - RoleAttributesMapper, - RoleAttributes.id == RoleAttributesMapper.attributes_id - ).where( - and_( - RoleAttributesMapper.role_id == user_role_id, - RoleAttributesMapper.value == True + attributes_query = ( + select(RoleAttributes.name) + .join(RoleAttributesMapper, RoleAttributes.id == RoleAttributesMapper.attributes_id) + .where( + and_( + RoleAttributesMapper.role_id == user_role_id, + RoleAttributesMapper.value, + ) ) ) attributes_result = await db.execute(attributes_query) diff --git a/backend/api/users/controller.py b/backend/api/users/controller.py index a143623..84b14f7 100644 --- a/backend/api/users/controller.py +++ b/backend/api/users/controller.py @@ -1,43 +1,61 @@ import redis -from typing import Optional -from core.redis import get_redis +from fastapi import APIRouter, Depends, HTTPException, Path, Query, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + from core.dependencies import get_db -from core.security import verify_token from core.permissions import Permission from core.rbac import require_permission -from sqlalchemy.ext.asyncio import AsyncSession -from utils.response import APIResponse, parse_responses, common_responses -from fastapi import APIRouter, Depends, HTTPException, Query, Request, Path, Response -from .services import get_all_users, create_user, update_user, delete_users, reset_user_password +from core.redis import get_redis +from core.security import verify_token +from utils.custom_exception import AuthorizationException, ConflictException, NotFoundException +from utils.response import APIResponse, common_responses, parse_responses + from .schema import ( - UserPagination, UserSortBy, UserCreate, UserUpdate, UserDelete, PasswordReset, UserResponse, - UserDeleteBatchResponse, user_delete_success_response_example, user_delete_partial_response_example, - user_delete_failed_response_example + PasswordReset, + UserCreate, + UserDelete, + UserDeleteBatchResponse, + UserPagination, + UserResponse, + UserSortBy, + UserUpdate, + user_delete_failed_response_example, + user_delete_partial_response_example, + user_delete_success_response_example, ) -from utils.custom_exception import NotFoundException, ConflictException, AuthorizationException +from .services import create_user, delete_users, get_all_users, reset_user_password, update_user router = APIRouter(tags=["Users"]) + @router.get( "", response_model=APIResponse[UserPagination], summary="Get all users", - responses=parse_responses({ - 200: ("Successfully retrieved users", UserPagination) - }, common_responses) + responses=parse_responses( + {200: ("Successfully retrieved users", UserPagination)}, common_responses + ), ) @require_permission([Permission.VIEW_USERS, Permission.MANAGE_USERS]) async def get_users( request: Request, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db), - keyword: Optional[str] = Query(None, description="Keyword to search for users"), - status: Optional[str] = Query(None, description="Filter user status (multiple values separated by commas, example: true,false)"), - role: Optional[str] = Query(None, description="Filter user role (multiple values separated by commas, example: admin,manager)"), + keyword: str | None = Query(None, description="Keyword to search for users"), + status: str | None = Query( + None, + description="Filter user status (multiple values separated by commas, example: true,false)", + ), + role: str | None = Query( + None, + description=( + "Filter user role (multiple values separated by commas, example: admin,manager)" + ), + ), page: int = Query(1, ge=1, description="Page number"), per_page: int = Query(10, ge=1, le=100, description="Number of users per page"), - sort_by: Optional[UserSortBy] = Query(None, description="Sort by field"), - desc: bool = Query(False, description="Sort order") + sort_by: UserSortBy | None = Query(None, description="Sort by field"), + desc: bool = Query(False, description="Sort order"), ): try: data = await get_all_users( @@ -48,27 +66,26 @@ async def get_users( page=page, per_page=per_page, sort_by=sort_by.value if sort_by else None, - desc=desc - ) + desc=desc, + ) return APIResponse(code=200, message="Successfully retrieved users", data=data) except Exception: raise HTTPException(status_code=500) + @router.post( "", response_model=APIResponse[UserResponse], response_model_exclude_none=True, summary="Create new user", - responses=parse_responses({ - 200: ("User created successfully", UserResponse) - }, common_responses) + responses=parse_responses({200: ("User created successfully", UserResponse)}, common_responses), ) @require_permission([Permission.MANAGE_USERS]) async def create_user_api( user_data: UserCreate, request: Request, token: dict = Depends(verify_token), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Create a new user account""" try: @@ -81,14 +98,13 @@ async def create_user_api( raise HTTPException(status_code=409, detail="Email already exists") raise HTTPException(status_code=500) + @router.put( "/{user_id}", response_model=APIResponse[UserResponse], response_model_exclude_none=True, summary="Update user info", - responses=parse_responses({ - 200: ("User updated successfully", UserResponse) - }, common_responses) + responses=parse_responses({200: ("User updated successfully", UserResponse)}, common_responses), ) @require_permission([Permission.MANAGE_USERS]) async def update_user_api( @@ -96,7 +112,7 @@ async def update_user_api( token: dict = Depends(verify_token), user_id: str = Path(..., description="User ID"), user_data: UserUpdate = None, - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Update user information""" try: @@ -111,16 +127,32 @@ async def update_user_api( except Exception: raise HTTPException(status_code=500) + @router.delete( "", response_model=APIResponse[UserDeleteBatchResponse], response_model_exclude_none=True, summary="Delete users", - responses=parse_responses({ - 200: ("All users deleted successfully", UserDeleteBatchResponse, user_delete_success_response_example), - 207: ("Users deleted with partial success", UserDeleteBatchResponse, user_delete_partial_response_example), - 400: ("All users failed to delete", UserDeleteBatchResponse, user_delete_failed_response_example) - }, common_responses) + responses=parse_responses( + { + 200: ( + "All users deleted successfully", + UserDeleteBatchResponse, + user_delete_success_response_example, + ), + 207: ( + "Users deleted with partial success", + UserDeleteBatchResponse, + user_delete_partial_response_example, + ), + 400: ( + "All users failed to delete", + UserDeleteBatchResponse, + user_delete_failed_response_example, + ), + }, + common_responses, + ), ) @require_permission([Permission.MANAGE_USERS]) async def delete_users_api( @@ -128,55 +160,44 @@ async def delete_users_api( request: Request, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db), - redis_client: redis.Redis = Depends(get_redis) + redis_client: redis.Redis = Depends(get_redis), ): """Delete multiple users""" try: batch_result = await delete_users(db, redis_client, delete_data.user_ids, token) - + # Determine response code based on results if batch_result.failed_count == 0: # All successful return APIResponse( - code=200, - message="All users deleted successfully", - data=batch_result + code=200, message="All users deleted successfully", data=batch_result ) elif batch_result.success_count == 0: # All failed - return 400 status code response = APIResponse( - code=400, - message="All users failed to delete", - data=batch_result + code=400, message="All users failed to delete", data=batch_result ) return Response( - content=response.model_dump_json(), - status_code=400, - media_type="application/json" + content=response.model_dump_json(), status_code=400, media_type="application/json" ) else: # Partial success - return 207 status code response = APIResponse( - code=207, - message="Users deleted with partial success", - data=batch_result + code=207, message="Users deleted with partial success", data=batch_result ) return Response( - content=response.model_dump_json(), - status_code=207, - media_type="application/json" + content=response.model_dump_json(), status_code=207, media_type="application/json" ) except Exception: raise HTTPException(status_code=500) + @router.post( "/{user_id}/reset-password", response_model=APIResponse[dict], response_model_exclude_none=True, summary="Reset user password", - responses=parse_responses({ - 200: ("Password reset successfully", dict) - }, common_responses) + responses=parse_responses({200: ("Password reset successfully", dict)}, common_responses), ) @require_permission([Permission.MANAGE_USERS]) async def reset_user_password_api( @@ -185,13 +206,15 @@ async def reset_user_password_api( request: Request = None, token: dict = Depends(verify_token), db: AsyncSession = Depends(get_db), - redis_client: redis.Redis = Depends(get_redis) + redis_client: redis.Redis = Depends(get_redis), ): """Reset user password and logout all devices""" try: await reset_user_password(db, redis_client, user_id, password_data.new_password) - return APIResponse(code=200, message="Password reset successfully and all devices logged out") + return APIResponse( + code=200, message="Password reset successfully and all devices logged out" + ) except NotFoundException: raise HTTPException(status_code=404, detail="User not found") except Exception: - raise HTTPException(status_code=500) \ No newline at end of file + raise HTTPException(status_code=500) diff --git a/backend/api/users/schema.py b/backend/api/users/schema.py index c194473..0ec5958 100644 --- a/backend/api/users/schema.py +++ b/backend/api/users/schema.py @@ -1,9 +1,11 @@ -from enum import Enum from datetime import datetime -from core.config import settings -from typing import List, Optional +from enum import StrEnum + from pydantic import BaseModel, EmailStr, Field +from core.config import settings + + class UserResponse(BaseModel): id: str = Field(..., description="User ID") email: str = Field(..., description="User email address") @@ -12,19 +14,19 @@ class UserResponse(BaseModel): phone: str = Field(..., description="Phone number") status: bool = Field(..., description="User status") created_at: datetime = Field(..., description="User creation time") - role: Optional[str] = Field(None, description="User role") - role_level: Optional[int] = Field( - None, description="Privilege level of the user's primary role" - ) + role: str | None = Field(None, description="User role") + role_level: int | None = Field(None, description="Privilege level of the user's primary role") + class UserPagination(BaseModel): - users: List[UserResponse] = Field(..., description="List of users") + users: list[UserResponse] = Field(..., description="List of users") total: int = Field(..., description="Total number of users") page: int = Field(..., description="Current page number") per_page: int = Field(..., description="Number of users per page") total_pages: int = Field(..., description="Total number of pages") -class UserSortBy(str, Enum): + +class UserSortBy(StrEnum): FIRST_NAME: str = "first_name" LAST_NAME: str = "last_name" EMAIL: str = "email" @@ -33,40 +35,56 @@ class UserSortBy(str, Enum): STATUS: str = "status" CREATED_AT: str = "created_at" + class UserCreate(BaseModel): first_name: str = Field(..., min_length=1, max_length=50, description="First name") last_name: str = Field(..., min_length=1, max_length=50, description="Last name") email: EmailStr = Field(..., description="User email address") phone: str = Field(..., min_length=1, max_length=20, description="Phone number") - password: str = Field(..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="Password") + password: str = Field( + ..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="Password" + ) status: bool = Field(True, description="User status") - role: Optional[str] = Field(None, description="User role") + role: str | None = Field(None, description="User role") + class UserUpdate(BaseModel): - first_name: Optional[str] = Field(None, min_length=1, max_length=50, description="First name") - last_name: Optional[str] = Field(None, min_length=1, max_length=50, description="Last name") - email: Optional[EmailStr] = Field(None, description="User email address") - phone: Optional[str] = Field(None, min_length=1, max_length=20, description="Phone number") - status: Optional[bool] = Field(None, description="User status") - role: Optional[str] = Field(None, description="User role") + first_name: str | None = Field(None, min_length=1, max_length=50, description="First name") + last_name: str | None = Field(None, min_length=1, max_length=50, description="Last name") + email: EmailStr | None = Field(None, description="User email address") + phone: str | None = Field(None, min_length=1, max_length=20, description="Phone number") + status: bool | None = Field(None, description="User status") + role: str | None = Field(None, description="User role") + class UserDelete(BaseModel): - user_ids: List[str] = Field(..., min_items=1, description="List of user IDs to delete") + user_ids: list[str] = Field(..., min_items=1, description="List of user IDs to delete") + class PasswordReset(BaseModel): - new_password: str = Field(..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="New password") + new_password: str = Field( + ..., min_length=settings.PASSWORD_MIN_LENGTH, max_length=50, description="New password" + ) + class UserDeleteResult(BaseModel): user_id: str = Field(..., description="User ID") - status: str = Field(..., description="Processing status: success, failed", example="success|failed", pattern="^(success|failed)$") + status: str = Field( + ..., + description="Processing status: success, failed", + example="success|failed", + pattern="^(success|failed)$", + ) message: str = Field(..., description="Result message") + class UserDeleteBatchResponse(BaseModel): - results: List[UserDeleteResult] = Field(..., description="Individual user deletion results") + results: list[UserDeleteResult] = Field(..., description="Individual user deletion results") total_users: int = Field(..., description="Total number of users processed") success_count: int = Field(..., description="Number of successfully deleted users") failed_count: int = Field(..., description="Number of failed deletions") + user_delete_success_response_example = { "code": 200, "message": "All users deleted successfully", @@ -75,18 +93,18 @@ class UserDeleteBatchResponse(BaseModel): { "user_id": "uuid-user-id-1", "status": "success", - "message": "User deleted successfully" + "message": "User deleted successfully", }, { "user_id": "uuid-user-id-2", "status": "success", - "message": "User deleted successfully" - } + "message": "User deleted successfully", + }, ], "total_users": 2, "success_count": 2, - "failed_count": 0 - } + "failed_count": 0, + }, } user_delete_partial_response_example = { @@ -97,33 +115,23 @@ class UserDeleteBatchResponse(BaseModel): { "user_id": "uuid-user-id-1", "status": "success", - "message": "User deleted successfully" + "message": "User deleted successfully", }, - { - "user_id": "uuid-user-id-2", - "status": "failed", - "message": "User not found" - } + {"user_id": "uuid-user-id-2", "status": "failed", "message": "User not found"}, ], "total_users": 2, "success_count": 1, - "failed_count": 1 - } + "failed_count": 1, + }, } user_delete_failed_response_example = { "code": 400, "message": "All users failed to delete", "data": { - "results": [ - { - "user_id": "uuid-user-id-1", - "status": "failed", - "message": "User not found" - } - ], + "results": [{"user_id": "uuid-user-id-1", "status": "failed", "message": "User not found"}], "total_users": 1, "success_count": 0, - "failed_count": 1 - } -} \ No newline at end of file + "failed_count": 1, + }, +} diff --git a/backend/api/users/services.py b/backend/api/users/services.py index 0ed6a41..59e6a4e 100644 --- a/backend/api/users/services.py +++ b/backend/api/users/services.py @@ -1,31 +1,40 @@ -import redis import logging -from models.users import Users -from models.roles import Roles -from models.role_mapper import RoleMapper -from models.login_logs import LoginLogs -from models.user_sessions import UserSessions -from models.password_reset_tokens import PasswordResetTokens -from models.email_verification_tokens import EmailVerificationTokens -from typing import Optional, List, Dict + +import redis +from sqlalchemy import case, delete, func, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy import select, func, or_, delete, case, update -from core.security import hash_password, clear_user_all_sessions -from .schema import UserResponse, UserPagination, UserCreate, UserUpdate, UserDeleteBatchResponse, UserDeleteResult -from utils.custom_exception import ( - ServerException, - ConflictException, - NotFoundException, - AuthorizationException, -) -from core.permissions import Permission + from core.config import settings +from core.permissions import Permission from core.rbac import ( check_user_has_super_role, get_user_attributes, get_user_role_level, is_super_admin_role_name, ) +from core.security import clear_user_all_sessions, hash_password +from models.email_verification_tokens import EmailVerificationTokens +from models.login_logs import LoginLogs +from models.password_reset_tokens import PasswordResetTokens +from models.role_mapper import RoleMapper +from models.roles import Roles +from models.user_sessions import UserSessions +from models.users import Users +from utils.custom_exception import ( + AuthorizationException, + ConflictException, + NotFoundException, + ServerException, +) + +from .schema import ( + UserCreate, + UserDeleteBatchResponse, + UserDeleteResult, + UserPagination, + UserResponse, + UserUpdate, +) logger = logging.getLogger("users") @@ -41,38 +50,38 @@ def _super_admin_user_ids_subquery(): async def get_all_users( db: AsyncSession, - keyword: Optional[str] = None, - status: Optional[str] = None, - role: Optional[str] = None, + keyword: str | None = None, + status: str | None = None, + role: str | None = None, page: int = 1, per_page: int = 10, - sort_by: Optional[str] = None, - desc: bool = False + sort_by: str | None = None, + desc: bool = False, ) -> UserPagination: """Get all users list""" try: query = select(Users) - + if keyword: query = query.where( or_( Users.first_name.ilike(f"%{keyword}%"), Users.last_name.ilike(f"%{keyword}%"), - Users.email.ilike(f"%{keyword}%") + Users.email.ilike(f"%{keyword}%"), ) ) - + if status: - status_list = [s.strip().lower() == 'true' for s in status.split(',')] + status_list = [s.strip().lower() == "true" for s in status.split(",")] if len(status_list) == 1: query = query.where(Users.status == status_list[0]) else: query = query.where(Users.status.in_(status_list)) - + has_role_join = False - + if role: - role_list = [r.strip() for r in role.split(',')] + role_list = [r.strip() for r in role.split(",")] query = query.join(RoleMapper, Users.id == RoleMapper.user_id) query = query.join(Roles, RoleMapper.role_id == Roles.id) query = query.where(Roles.name.in_(role_list)) @@ -80,7 +89,7 @@ async def get_all_users( if not settings.SHOW_SUPER_ADMIN: query = query.where(Users.id.not_in(_super_admin_user_ids_subquery())) - + if sort_by: if sort_by == "role": # For role sorting, use LEFT JOIN with distinct to avoid duplicates @@ -106,51 +115,43 @@ async def get_all_users( query = query.order_by(Users.created_at.desc()) else: query = query.order_by(Users.id.asc()) - + count_query = select(func.count(Users.id)) if keyword: count_query = count_query.where( or_( Users.first_name.ilike(f"%{keyword}%"), Users.last_name.ilike(f"%{keyword}%"), - Users.email.ilike(f"%{keyword}%") + Users.email.ilike(f"%{keyword}%"), ) ) if status: - status_list = [s.strip().lower() == 'true' for s in status.split(',')] + status_list = [s.strip().lower() == "true" for s in status.split(",")] if len(status_list) == 1: count_query = count_query.where(Users.status == status_list[0]) else: count_query = count_query.where(Users.status.in_(status_list)) if role: - role_list = [r.strip() for r in role.split(',')] + role_list = [r.strip() for r in role.split(",")] count_query = count_query.join(RoleMapper, Users.id == RoleMapper.user_id) count_query = count_query.join(Roles, RoleMapper.role_id == Roles.id) count_query = count_query.where(Roles.name.in_(role_list)) if not settings.SHOW_SUPER_ADMIN: - count_query = count_query.where( - Users.id.not_in(_super_admin_user_ids_subquery()) - ) - + count_query = count_query.where(Users.id.not_in(_super_admin_user_ids_subquery())) + total_result = await db.execute(count_query) total = total_result.scalar() - + offset = (page - 1) * per_page query = query.offset(offset).limit(per_page) - + result = await db.execute(query) users = result.scalars().all() - + if not users: - return UserPagination( - users=[], - total=0, - page=page, - per_page=per_page, - total_pages=0 - ) - + return UserPagination(users=[], total=0, page=page, per_page=per_page, total_pages=0) + user_roles = await _get_user_roles_map(db, [user.id for user in users]) user_responses = [] @@ -168,20 +169,17 @@ async def get_all_users( role_level=role_level, ) user_responses.append(user_response) - + total_pages = (total + per_page - 1) // per_page - + return UserPagination( - users=user_responses, - total=total, - page=page, - per_page=per_page, - total_pages=total_pages + users=user_responses, total=total, page=page, per_page=per_page, total_pages=total_pages ) - + except Exception as e: raise ServerException(f"Failed to retrieve users: {str(e)}") + async def create_user( db: AsyncSession, user_data: UserCreate, @@ -192,10 +190,7 @@ async def create_user( # Check if the email already exists result = await db.execute( select(Users).where( - or_( - Users.email == user_data.email, - Users.pending_email == user_data.email - ) + or_(Users.email == user_data.email, Users.pending_email == user_data.email) ) ) existing_user = result.scalar_one_or_none() @@ -204,27 +199,27 @@ async def create_user( if user_data.role: await _assert_can_manage_user_role(db, actor_user_id, user_data.role) - + user = Users( first_name=user_data.first_name, last_name=user_data.last_name, email=user_data.email, phone=user_data.phone, hash_password=await hash_password(user_data.password), - status=user_data.status + status=user_data.status, ) - + db.add(user) await db.commit() await db.refresh(user) - + user_role = None user_role_level = None if user_data.role: await _assign_user_role(db, user.id, user_data.role) user_role = user_data.role user_role_level = await _get_role_level_by_name(db, user_data.role) - + return UserResponse( id=user.id, email=user.email, @@ -236,12 +231,13 @@ async def create_user( role=user_role, role_level=user_role_level, ) - - except (ConflictException, NotFoundException, AuthorizationException): + + except ConflictException, NotFoundException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to create user: {str(e)}") + async def update_user( db: AsyncSession, user_id: str, @@ -250,22 +246,17 @@ async def update_user( ) -> UserResponse: """Update user information""" try: - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + result = await db.execute(select(Users).where(Users.id == user_id)) user = result.scalar_one_or_none() if not user: raise NotFoundException("User not found") - + # Check if the email is already used by another user if user_data.email and user_data.email != user.email: result = await db.execute( select(Users).where( - or_( - Users.email == user_data.email, - Users.pending_email == user_data.email - ), - Users.id != user_id + or_(Users.email == user_data.email, Users.pending_email == user_data.email), + Users.id != user_id, ) ) if result.scalar_one_or_none(): @@ -285,8 +276,8 @@ async def update_user( user_data.role, target_user_id=user_id, ) - - update_data = user_data.model_dump(exclude_unset=True, exclude={'role'}) + + update_data = user_data.model_dump(exclude_unset=True, exclude={"role"}) email_changed = "email" in update_data and update_data["email"] != user.email for field, value in update_data.items(): @@ -301,17 +292,17 @@ async def update_user( .where( EmailVerificationTokens.user_id == user.id, EmailVerificationTokens.token_type == "email_change", - EmailVerificationTokens.is_used == False + EmailVerificationTokens.is_used.is_(False), ) .values(is_used=True) ) - + await db.commit() await db.refresh(user) - + if role_update_requested: await _update_user_role(db, user_id, user_data.role) - + role_query = ( select(Roles.name, Roles.level) .join(RoleMapper, Roles.id == RoleMapper.role_id) @@ -323,7 +314,7 @@ async def update_user( role_row = role_result.one_or_none() user_role = role_row[0] if role_row else None user_role_level = role_row[1] if role_row else None - + return UserResponse( id=user.id, email=user.email, @@ -335,19 +326,22 @@ async def update_user( role=user_role, role_level=user_role_level, ) - - except (ConflictException, NotFoundException, AuthorizationException): + + except ConflictException, NotFoundException, AuthorizationException: raise except Exception as e: raise ServerException(f"Failed to update user: {str(e)}") -async def delete_users(db: AsyncSession, redis_client: redis.Redis, user_ids: List[str], token: Optional[dict] = None) -> UserDeleteBatchResponse: + +async def delete_users( + db: AsyncSession, redis_client: redis.Redis, user_ids: list[str], token: dict | None = None +) -> UserDeleteBatchResponse: """Delete multiple users with detailed batch processing results""" try: results = [] success_count = 0 failed_count = 0 - + # Get current user ID from token current_user_id = token.get("sub") if token else None actor_is_super = False @@ -356,125 +350,125 @@ async def delete_users(db: AsyncSession, redis_client: redis.Redis, user_ids: Li actor_is_super = await check_user_has_super_role(current_user_id, db) if not actor_is_super: actor_level = await get_user_role_level(current_user_id, db) - + # Check which users exist - result = await db.execute( - select(Users.id).where(Users.id.in_(user_ids)) - ) + result = await db.execute(select(Users.id).where(Users.id.in_(user_ids))) existing_ids = set(result.scalars().all()) - + # Process each user ID for user_id in user_ids: try: # Skip if trying to delete own account if current_user_id and user_id == current_user_id: - results.append(UserDeleteResult( - user_id=user_id, - status="failed", - message="Cannot delete your own account" - )) + results.append( + UserDeleteResult( + user_id=user_id, + status="failed", + message="Cannot delete your own account", + ) + ) failed_count += 1 continue - + if user_id in existing_ids: if await check_user_has_super_role(user_id, db): - results.append(UserDeleteResult( - user_id=user_id, - status="failed", - message="Cannot delete a system super-admin user" - )) + results.append( + UserDeleteResult( + user_id=user_id, + status="failed", + message="Cannot delete a system super-admin user", + ) + ) failed_count += 1 continue if not actor_is_super: target_level = await get_user_role_level(user_id, db) if target_level > actor_level: - results.append(UserDeleteResult( - user_id=user_id, - status="failed", - message="Cannot delete a user with a higher role level than your own" - )) + results.append( + UserDeleteResult( + user_id=user_id, + status="failed", + message=( + "Cannot delete a user with a higher role level " + "than your own" + ), + ) + ) failed_count += 1 continue # Clear user sessions and tokens before deletion await clear_user_all_sessions(db, redis_client, user_id) - + # Delete related records first to avoid foreign key constraints await _delete_user_related_records(db, user_id) - + # Delete the user - await db.execute( - delete(Users).where(Users.id == user_id) + await db.execute(delete(Users).where(Users.id == user_id)) + + results.append( + UserDeleteResult( + user_id=user_id, status="success", message="User deleted successfully" + ) ) - - results.append(UserDeleteResult( - user_id=user_id, - status="success", - message="User deleted successfully" - )) success_count += 1 else: - results.append(UserDeleteResult( - user_id=user_id, - status="failed", - message="User not found" - )) + results.append( + UserDeleteResult(user_id=user_id, status="failed", message="User not found") + ) failed_count += 1 - + except Exception as e: - results.append(UserDeleteResult( - user_id=user_id, - status="failed", - message=f"Failed to delete user: {str(e)}" - )) + results.append( + UserDeleteResult( + user_id=user_id, status="failed", message=f"Failed to delete user: {str(e)}" + ) + ) failed_count += 1 - + # Commit all successful deletions if success_count > 0: await db.commit() - + return UserDeleteBatchResponse( results=results, total_users=len(user_ids), success_count=success_count, - failed_count=failed_count + failed_count=failed_count, ) - + except Exception as e: raise ServerException(f"Failed to delete users: {str(e)}") + async def reset_user_password( - db: AsyncSession, - redis_client: redis.Redis, - user_id: str, - new_password: str + db: AsyncSession, redis_client: redis.Redis, user_id: str, new_password: str ) -> bool: """Reset user password and logout all devices""" try: - result = await db.execute( - select(Users).where(Users.id == user_id) - ) + result = await db.execute(select(Users).where(Users.id == user_id)) user = result.scalar_one_or_none() if not user: raise NotFoundException("User not found") - + user.hash_password = await hash_password(new_password) user.password_reset_required = True await db.commit() - + await clear_user_all_sessions(db, redis_client, user_id) - + return True - + except NotFoundException: raise except Exception as e: raise ServerException(f"Failed to reset password: {str(e)}") + async def _get_user_roles_map( - db: AsyncSession, user_ids: List[str] -) -> Dict[str, tuple[Optional[str], Optional[int]]]: + db: AsyncSession, user_ids: list[str] +) -> dict[str, tuple[str | None, int | None]]: """Batch-load primary role name/level per user (highest level wins).""" if not user_ids: return {} @@ -486,13 +480,14 @@ async def _get_user_roles_map( .order_by(Roles.level.desc(), Roles.name.asc()) ) role_result = await db.execute(roles_query) - user_roles: Dict[str, tuple[Optional[str], Optional[int]]] = {} + user_roles: dict[str, tuple[str | None, int | None]] = {} for user_id, role_name, role_level in role_result.all(): if user_id not in user_roles: user_roles[user_id] = (role_name, role_level) return user_roles -async def _get_user_role_name(db: AsyncSession, user_id: str) -> Optional[str]: + +async def _get_user_role_name(db: AsyncSession, user_id: str) -> str | None: result = await db.execute( select(Roles.name) .join(RoleMapper, Roles.id == RoleMapper.role_id) @@ -503,10 +498,8 @@ async def _get_user_role_name(db: AsyncSession, user_id: str) -> Optional[str]: return result.scalar_one_or_none() -async def _get_role_level_by_name(db: AsyncSession, role_name: str) -> Optional[int]: - result = await db.execute( - select(Roles.level).where(Roles.name == role_name).limit(1) - ) +async def _get_role_level_by_name(db: AsyncSession, role_name: str) -> int | None: + result = await db.execute(select(Roles.level).where(Roles.name == role_name).limit(1)) level = result.scalar_one_or_none() return int(level) if level is not None else None @@ -525,17 +518,15 @@ async def _assert_can_manage_target_user( actor_level = await get_user_role_level(actor_user_id, db) target_level = await get_user_role_level(target_user_id, db) if target_level > actor_level: - raise AuthorizationException( - "Cannot manage a user with a higher role level than your own" - ) + raise AuthorizationException("Cannot manage a user with a higher role level than your own") async def _assert_can_manage_user_role( db: AsyncSession, actor_user_id: str, - role_name: Optional[str], + role_name: str | None, *, - target_user_id: Optional[str] = None, + target_user_id: str | None = None, ) -> None: """ Require manage-roles (or super-admin) to change roles. @@ -571,84 +562,66 @@ async def _assert_can_manage_user_role( if new_role_level is None: raise NotFoundException(f"Role '{role_name}' not found") if new_role_level > actor_level: - raise AuthorizationException( - "Cannot assign a role with a higher level than your own" - ) + raise AuthorizationException("Cannot assign a role with a higher level than your own") async def _assign_user_role(db: AsyncSession, user_id: str, role_name: str) -> None: """Assign a role to a user""" try: - role_result = await db.execute( - select(Roles).where(Roles.name == role_name) - ) + role_result = await db.execute(select(Roles).where(Roles.name == role_name)) role = role_result.scalar_one_or_none() if not role: raise NotFoundException(f"Role '{role_name}' not found") - + existing_mapping = await db.execute( - select(RoleMapper).where( - RoleMapper.user_id == user_id, - RoleMapper.role_id == role.id - ) + select(RoleMapper).where(RoleMapper.user_id == user_id, RoleMapper.role_id == role.id) ) if existing_mapping.scalar_one_or_none(): return - - role_mapping = RoleMapper( - user_id=user_id, - role_id=role.id - ) + + role_mapping = RoleMapper(user_id=user_id, role_id=role.id) db.add(role_mapping) await db.commit() - + except NotFoundException: raise except Exception as e: raise ServerException(f"Failed to assign role: {str(e)}") -async def _update_user_role(db: AsyncSession, user_id: str, role_name: Optional[str]) -> None: + +async def _update_user_role(db: AsyncSession, user_id: str, role_name: str | None) -> None: """Update user role (remove existing and assign new one)""" try: - await db.execute( - delete(RoleMapper).where(RoleMapper.user_id == user_id) - ) - + await db.execute(delete(RoleMapper).where(RoleMapper.user_id == user_id)) + if role_name: await _assign_user_role(db, user_id, role_name) - + await db.commit() - + except Exception as e: raise ServerException(f"Failed to update user role: {str(e)}") + async def _delete_user_related_records(db: AsyncSession, user_id: str) -> None: """Delete all records related to a user to avoid foreign key constraints""" try: # Delete login logs - await db.execute( - delete(LoginLogs).where(LoginLogs.user_id == user_id) - ) - + await db.execute(delete(LoginLogs).where(LoginLogs.user_id == user_id)) + # Delete user sessions - await db.execute( - delete(UserSessions).where(UserSessions.user_id == user_id) - ) - + await db.execute(delete(UserSessions).where(UserSessions.user_id == user_id)) + # Delete role mappings - await db.execute( - delete(RoleMapper).where(RoleMapper.user_id == user_id) - ) - + await db.execute(delete(RoleMapper).where(RoleMapper.user_id == user_id)) + # Delete password reset tokens - await db.execute( - delete(PasswordResetTokens).where(PasswordResetTokens.user_id == user_id) - ) + await db.execute(delete(PasswordResetTokens).where(PasswordResetTokens.user_id == user_id)) # Delete email verification tokens await db.execute( delete(EmailVerificationTokens).where(EmailVerificationTokens.user_id == user_id) ) - + except Exception as e: - raise ServerException(f"Failed to delete user related records: {str(e)}") \ No newline at end of file + raise ServerException(f"Failed to delete user related records: {str(e)}") diff --git a/backend/core/config.py b/backend/core/config.py index 6212a7c..fee2f1b 100644 --- a/backend/core/config.py +++ b/backend/core/config.py @@ -1,17 +1,17 @@ from dotenv import load_dotenv + load_dotenv() # Load .env +import logging.config import os import re + import yaml -import logging.config from pydantic_settings import BaseSettings # Docs / health probes — skip logging, rate limit, and tracing. # Not an env setting: these paths are part of the app, not deployment. -SKIP_PATHS: frozenset[str] = frozenset( - {"/", "/docs", "/redoc", "/openapi.json", "/healthz"} -) +SKIP_PATHS: frozenset[str] = frozenset({"/", "/docs", "/redoc", "/openapi.json", "/healthz"}) # CORS preflight — skip logging, rate limit, and tracing (URL exclude cannot filter method). SKIP_METHODS: frozenset[str] = frozenset({"OPTIONS"}) @@ -30,6 +30,7 @@ def otel_excluded_urls(paths: frozenset[str] = SKIP_PATHS) -> str: patterns.append(re.escape(path)) return ",".join(patterns) + class Settings(BaseSettings): # Project settings PROJECT_NAME: str = "Backend API Docs" @@ -82,7 +83,7 @@ class Settings(BaseSettings): # Session settings SESSION_EXPIRE_MINUTES: int = 10080 # 7 days CSRF_TOKEN_EXPIRE_MINUTES: int = 30 - + # Cookie settings COOKIE_SECURE: bool = SSL_ENABLE COOKIE_HTTPONLY: bool = True @@ -130,12 +131,14 @@ def rate_limit_whitelist_ips(self) -> set[str]: # When false, users with the system super-admin role are hidden from user list for everyone SHOW_SUPER_ADMIN: bool = False + # Create a settings instance to be imported elsewhere settings = Settings() + def setup_logging(yaml_path="logging_config.yaml"): os.makedirs("logs", exist_ok=True) - with open(yaml_path, "r") as f: + with open(yaml_path) as f: config = yaml.safe_load(f) # Override the root logger or specified logger's level with LOG_LEVEL from environment log_level = settings.LOG_LEVEL @@ -147,4 +150,4 @@ def setup_logging(yaml_path="logging_config.yaml"): logger["level"] = log_level if "handlers" in config and "file" in config["handlers"]: config["handlers"]["file"]["backupCount"] = settings.LOG_LOCAL_RETENTION_DAYS - logging.config.dictConfig(config) \ No newline at end of file + logging.config.dictConfig(config) diff --git a/backend/core/database.py b/backend/core/database.py index 8f993d4..6188a71 100644 --- a/backend/core/database.py +++ b/backend/core/database.py @@ -1,11 +1,14 @@ import logging -from core.config import settings + from sqlalchemy import create_engine +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.orm import declarative_base, sessionmaker -from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker + +from core.config import settings logger = logging.getLogger("database") + def make_async_url(url: str) -> str: if url.startswith("mysql://"): return url.replace("mysql://", "mysql+aiomysql://", 1) @@ -13,6 +16,7 @@ def make_async_url(url: str) -> str: return url.replace("mysql+pymysql://", "mysql+aiomysql://", 1) return url + # Async engine/session for API async_engine = create_async_engine( make_async_url(settings.DATABASE_URL), @@ -25,10 +29,14 @@ def make_async_url(url: str) -> str: max_overflow=settings.DB_MAX_OVERFLOW, connect_args={ "charset": "utf8mb4", - } + }, ) AsyncSessionLocal = async_sessionmaker( - bind=async_engine, class_=AsyncSession, expire_on_commit=False, autoflush=False, autocommit=False + bind=async_engine, + class_=AsyncSession, + expire_on_commit=False, + autoflush=False, + autocommit=False, ) # Sync engine/session for migration and schedule @@ -46,20 +54,22 @@ def make_async_url(url: str) -> str: "connect_timeout": settings.DB_CONNECT_TIMEOUT, "read_timeout": settings.DB_READ_TIMEOUT, "write_timeout": settings.DB_WRITE_TIMEOUT, - } + }, ) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) # Base class for declarative models Base = declarative_base() + async def init_db(): """Initialize the database with default data""" try: # Import here to avoid circular import from core.init_db import init_database + await init_database() logger.info("Database initialization completed") except Exception as e: logger.error(f"Database initialization failed: {str(e)}") - raise \ No newline at end of file + raise diff --git a/backend/core/dependencies.py b/backend/core/dependencies.py index 1c3798b..7a8abfc 100644 --- a/backend/core/dependencies.py +++ b/backend/core/dependencies.py @@ -1,16 +1,19 @@ import logging + from fastapi import HTTPException from sqlalchemy.ext.asyncio import AsyncSession + from core.database import AsyncSessionLocal, SessionLocal logger = logging.getLogger("dependencies") + # Async DB dependency async def get_db() -> AsyncSession: async with AsyncSessionLocal() as db: try: yield db - except (HTTPException): + except HTTPException: await db.rollback() raise except Exception as e: @@ -18,6 +21,7 @@ async def get_db() -> AsyncSession: await db.rollback() raise e + # Sync DB dependency (for migration/schedule) def get_sync_db(): db = SessionLocal() @@ -28,4 +32,4 @@ def get_sync_db(): db.rollback() raise e finally: - db.close() \ No newline at end of file + db.close() diff --git a/backend/core/init_db.py b/backend/core/init_db.py index 080adb0..0f07eeb 100644 --- a/backend/core/init_db.py +++ b/backend/core/init_db.py @@ -1,16 +1,19 @@ import logging + from sqlalchemy import delete, select, text -from models.users import Users -from models.roles import Roles + from core.config import settings -from core.security import hash_password -from models.role_mapper import RoleMapper from core.database import AsyncSessionLocal from core.permissions import get_attributes +from core.security import hash_password from models.role_attributes import RoleAttributes +from models.role_mapper import RoleMapper +from models.roles import Roles +from models.users import Users logger = logging.getLogger("init_db") + async def is_already_initialized() -> bool: """Return True if super admin role and user already exist (seed done).""" async with AsyncSessionLocal() as db: @@ -38,6 +41,7 @@ async def is_already_initialized() -> bool: return False + async def init_database(): """Initialize database with default data: role attributes, roles and admin account""" lock_name = "init_db_lock" @@ -77,6 +81,7 @@ async def init_database(): await lock_session.execute(text("SELECT RELEASE_LOCK(:name)"), {"name": lock_name}) await lock_session.commit() + async def create_role_attributes(): """Create role attributes""" async with AsyncSessionLocal() as db: @@ -84,21 +89,27 @@ async def create_role_attributes(): attributes = get_attributes() created_count = 0 updated_count = 0 - + for attr_config in attributes: existing_attr = await db.execute( select(RoleAttributes).where(RoleAttributes.name == attr_config["name"]) ) existing = existing_attr.scalar_one_or_none() if existing: - if getattr(existing, "group", None) is None and attr_config.get("group") is not None: + if ( + getattr(existing, "group", None) is None + and attr_config.get("group") is not None + ): existing.group = attr_config.get("group") updated_count += 1 - if getattr(existing, "category", None) is None and attr_config.get("category") is not None: + if ( + getattr(existing, "category", None) is None + and attr_config.get("category") is not None + ): existing.category = attr_config.get("category") updated_count += 1 continue - + attribute = RoleAttributes( name=attr_config["name"], group=attr_config.get("group"), @@ -106,14 +117,15 @@ async def create_role_attributes(): ) db.add(attribute) created_count += 1 - + await db.commit() - + except Exception as e: logger.error(f"Failed to create role attributes: {str(e)}") await db.rollback() raise + async def create_default_roles(): """Create system super-admin role and a basic user role.""" async with AsyncSessionLocal() as db: @@ -128,18 +140,18 @@ async def create_default_roles(): "name": "user", "description": "Regular user role with basic permissions", "level": settings.DEFAULT_USER_ROLE_LEVEL, - } + }, ] - + created_count = 0 - + for role_config in default_roles: existing_role = await db.execute( select(Roles).where(Roles.name == role_config["name"]) ) if existing_role.scalar_one_or_none(): continue - + role = Roles( name=role_config["name"], description=role_config["description"], @@ -147,14 +159,15 @@ async def create_default_roles(): ) db.add(role) created_count += 1 - + await db.commit() - + except Exception as e: logger.error(f"Failed to create roles: {str(e)}") await db.rollback() raise + async def create_default_admin(): """Create default super-admin account from ENV settings.""" async with AsyncSessionLocal() as db: @@ -163,11 +176,11 @@ async def create_default_admin(): select(Roles).where(Roles.name == settings.DEFAULT_SUPER_ADMIN_ROLE) ) super_role = super_role_result.scalar_one_or_none() - + if not super_role: logger.error("Super admin role not found, please run role initialization first") return - + super_users_result = await db.execute( select(Users) .join(RoleMapper, Users.id == RoleMapper.user_id) @@ -175,31 +188,32 @@ async def create_default_admin(): .where(Roles.name == settings.DEFAULT_SUPER_ADMIN_ROLE) ) existing_super_users = super_users_result.scalars().all() - + if existing_super_users: - logger.info(f"Super admin users already exist: {[user.email for user in existing_super_users]}") + logger.info( + "Super admin users already exist: " + f"{[user.email for user in existing_super_users]}" + ) await db.commit() return - + existing_user_result = await db.execute( select(Users).where(Users.email == settings.DEFAULT_ADMIN_EMAIL) ) existing_user = existing_user_result.scalar_one_or_none() - + if existing_user: existing_role_mapping = await db.execute( select(RoleMapper).where( - RoleMapper.user_id == existing_user.id, - RoleMapper.role_id == super_role.id + RoleMapper.user_id == existing_user.id, RoleMapper.role_id == super_role.id ) ) if not existing_role_mapping.scalar_one_or_none(): - role_mapping = RoleMapper( - user_id=existing_user.id, - role_id=super_role.id - ) + role_mapping = RoleMapper(user_id=existing_user.id, role_id=super_role.id) db.add(role_mapping) - logger.info(f"Assigned super admin role to existing user: {existing_user.email}") + logger.info( + f"Assigned super admin role to existing user: {existing_user.email}" + ) await db.execute( delete(RoleMapper).where( RoleMapper.user_id == existing_user.id, @@ -217,21 +231,18 @@ async def create_default_admin(): password_reset_required=False, email_verified=True, ) - + db.add(admin_user) await db.commit() await db.refresh(admin_user) - - role_mapping = RoleMapper( - user_id=admin_user.id, - role_id=super_role.id - ) + + role_mapping = RoleMapper(user_id=admin_user.id, role_id=super_role.id) db.add(role_mapping) - + await db.commit() logger.info("Admin account initialization completed") - + except Exception as e: logger.error(f"Failed to create admin account: {str(e)}") await db.rollback() - raise \ No newline at end of file + raise diff --git a/backend/core/permissions.py b/backend/core/permissions.py index d7a6ce9..d8dd532 100644 --- a/backend/core/permissions.py +++ b/backend/core/permissions.py @@ -1,5 +1,5 @@ from enum import Enum -from typing import List, Dict, Optional + class Permission(Enum): """Attributes for roles""" @@ -7,32 +7,33 @@ class Permission(Enum): # User management attributes VIEW_USERS = ("view-users", "system-management", "user-management") MANAGE_USERS = ("manage-users", "system-management", "user-management") - + # Role management attributes VIEW_ROLES = ("view-roles", "system-management", "role-management") MANAGE_ROLES = ("manage-roles", "system-management", "role-management") - + def __init__( self, value: str, - group: Optional[str] = None, - category: Optional[str] = None, + group: str | None = None, + category: str | None = None, ): self._value_ = value self.group = group self.category = category - + def __str__(self) -> str: """Return string value of Permission.VIEW_USERS""" return self.value -def get_attributes() -> List[Dict[str, str]]: + +def get_attributes() -> list[dict[str, str]]: """Get default role attributes configuration""" return [ { "name": permission.value, "group": getattr(permission, "group", None), - "category": getattr(permission, "category", None) + "category": getattr(permission, "category", None), } for permission in Permission - ] \ No newline at end of file + ] diff --git a/backend/core/rbac.py b/backend/core/rbac.py index d252618..a105d0a 100644 --- a/backend/core/rbac.py +++ b/backend/core/rbac.py @@ -1,22 +1,25 @@ import logging from functools import wraps -from typing import List, Dict -from sqlalchemy import func, select -from models.roles import Roles -from core.config import settings + from fastapi import HTTPException, status -from models.role_mapper import RoleMapper +from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession + +from core.config import settings from models.role_attributes import RoleAttributes -from utils.custom_exception import ServerException from models.role_attributes_mapper import RoleAttributesMapper +from models.role_mapper import RoleMapper +from models.roles import Roles +from utils.custom_exception import ServerException logger = logging.getLogger("rbac") + def is_super_admin_role_name(role_name: str | None) -> bool: """True when role_name is the ENV-configured system super-admin role.""" return bool(role_name) and role_name == settings.DEFAULT_SUPER_ADMIN_ROLE + async def get_user_role_level(user_id: str, db: AsyncSession) -> int: """Return the user's highest role level, or 0 when the user has no role.""" try: @@ -31,6 +34,7 @@ async def get_user_role_level(user_id: str, db: AsyncSession) -> int: logger.error(f"Failed to get user role level: {e}") return 0 + async def get_user_role_id(user_id: str, db: AsyncSession) -> str | None: """Return the user's primary role id (highest level), or None when unassigned.""" try: @@ -46,6 +50,7 @@ async def get_user_role_id(user_id: str, db: AsyncSession) -> str | None: logger.error(f"Failed to get user role id: {e}") return None + async def user_has_role(user_id: str, role_id: str, db: AsyncSession) -> bool: """True when the user is currently assigned the given role.""" try: @@ -60,7 +65,8 @@ async def user_has_role(user_id: str, role_id: str, db: AsyncSession) -> bool: logger.error(f"Failed to check user role assignment: {e}") return False -async def get_user_attributes(user_id: str, db: AsyncSession) -> Dict[str, bool]: + +async def get_user_attributes(user_id: str, db: AsyncSession) -> dict[str, bool]: """Get user attributes""" try: # Check if user has super admin role first @@ -69,32 +75,30 @@ async def get_user_attributes(user_id: str, db: AsyncSession) -> Dict[str, bool] all_attributes_result = await db.execute(select(RoleAttributes.name)) all_attributes = [row.name for row in all_attributes_result] return {attr: True for attr in all_attributes} - + result = await db.execute( - select( - RoleAttributes.name, - RoleAttributesMapper.value - ) + select(RoleAttributes.name, RoleAttributesMapper.value) .join(RoleAttributesMapper, RoleAttributes.id == RoleAttributesMapper.attributes_id) .join(RoleMapper, RoleMapper.role_id == RoleAttributesMapper.role_id) .where(RoleMapper.user_id == user_id) ) - + attributes = {} for row in result: attr_name = row.name attr_value = row.value - + if attr_name in attributes: attributes[attr_name] = attributes[attr_name] or attr_value else: attributes[attr_name] = attr_value - + return attributes except Exception as e: ServerException(f"Failed to get user attributes: {e}") return {} + async def check_user_has_super_role(user_id: str, db: AsyncSession) -> bool: """Check if user has super admin role""" try: @@ -103,49 +107,48 @@ async def check_user_has_super_role(user_id: str, db: AsyncSession) -> bool: .join(RoleMapper, Roles.id == RoleMapper.role_id) .where(RoleMapper.user_id == user_id) ) - + user_roles = [row.name for row in result] return settings.DEFAULT_SUPER_ADMIN_ROLE in user_roles except Exception as e: logger.error(f"Failed to check super role: {e}") return False -def require_permission(required_attributes: List[str]): + +def require_permission(required_attributes: list[str]): """Permission check decorator""" + def decorator(func): @wraps(func) async def wrapper(*args, **kwargs): - token = kwargs.get('token') - db = kwargs.get('db') - + token = kwargs.get("token") + db = kwargs.get("db") + if not db: - raise HTTPException( - status_code=status.HTTP_500_INTERNAL_SERVER_ERROR - ) - + raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR) + user_id = token.get("sub") - + # Check if user has super admin role first if await check_user_has_super_role(user_id, db): return await func(*args, **kwargs) - + # If not super admin, check specific permissions user_attributes = await get_user_attributes(user_id, db) # Convert Permission enum to string value if needed def get_attr_value(attr): """Convert Permission enum to string value, or return as-is if already a string""" - if hasattr(attr, 'value'): + if hasattr(attr, "value"): return attr.value return str(attr) if attr else attr # Check if the user has at least one of the required permissions attr_values = [get_attr_value(attr) for attr in required_attributes] has_permission = any( - user_attributes.get(attr_value, False) - for attr_value in attr_values + user_attributes.get(attr_value, False) for attr_value in attr_values ) - + if not has_permission: logger.warning( f"Permission denied for user {user_id}. " @@ -153,10 +156,11 @@ def get_attr_value(attr): f"User has: {[k for k, v in user_attributes.items() if v]}" ) raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail="Permission denied" + status_code=status.HTTP_403_FORBIDDEN, detail="Permission denied" ) - + return await func(*args, **kwargs) + return wrapper - return decorator \ No newline at end of file + + return decorator diff --git a/backend/core/redis.py b/backend/core/redis.py index 69448d9..2038157 100644 --- a/backend/core/redis.py +++ b/backend/core/redis.py @@ -1,14 +1,15 @@ import redis.asyncio as aioredis + from core.config import settings _redis = None + async def init_redis(): global _redis - _redis = await aioredis.from_url( - settings.REDIS_URL, encoding="utf-8", decode_responses=True - ) + _redis = await aioredis.from_url(settings.REDIS_URL, encoding="utf-8", decode_responses=True) return _redis + def get_redis(): - return _redis \ No newline at end of file + return _redis diff --git a/backend/core/security.py b/backend/core/security.py index 3e8852f..c0a2ff7 100644 --- a/backend/core/security.py +++ b/backend/core/security.py @@ -1,33 +1,38 @@ import ast -import redis import logging -from models.users import Users -from jose import jwt, JWTError -from core.redis import get_redis -from core.config import settings -from core.dependencies import get_db -from sqlalchemy import update, select -from typing import Optional, Dict, Any from datetime import datetime, timedelta +from typing import Any + import bcrypt -from models.user_sessions import UserSessions +import redis +from fastapi import Depends, HTTPException, status +from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer +from jose import JWTError, jwt +from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession + +from core.config import settings +from core.dependencies import get_db +from core.redis import get_redis +from models.user_sessions import UserSessions +from models.users import Users from utils.custom_exception import ServerException -from fastapi import HTTPException, status, Depends -from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials logger = logging.getLogger("security") + async def hash_password(password: str) -> str: return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8") + async def verify_password(plain_password: str, hashed_password: str) -> bool: return bcrypt.checkpw( plain_password.encode("utf-8"), hashed_password.encode("utf-8"), ) -async def create_access_token(data: Dict[str, Any]) -> str: + +async def create_access_token(data: dict[str, Any]) -> str: to_encode = data.copy() now = datetime.now().astimezone() expire = now + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) @@ -35,28 +40,30 @@ async def create_access_token(data: Dict[str, Any]) -> str: encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) return encoded_jwt + async def create_password_reset_token(user_id: str, email: str) -> str: """Create password reset token""" try: now = datetime.now().astimezone() expire = now + timedelta(minutes=settings.PASSWORD_RESET_TOKEN_EXPIRE_MINUTES) - + payload = { "sub": user_id, "email": email, "token_type": "password_reset", "force_change_password": True, "iat": now, - "exp": expire + "exp": expire, } - + token = jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM) - + return token - + except Exception as e: raise ServerException(f"Failed to create password reset token: {str(e)}") + async def create_csrf_token(session_id: str) -> str: """Create CSRF token bound to a session""" try: @@ -75,38 +82,43 @@ async def create_csrf_token(session_id: str) -> str: except Exception as e: raise ServerException(f"Failed to create CSRF token: {str(e)}") + async def create_email_verification_token(user_id: str, email: str, token_type: str) -> str: """Create email verification token""" try: now = datetime.now().astimezone() expire = now + timedelta(minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES) - + payload = { "sub": user_id, "email": email, "token_type": "email_verification", "verification_type": token_type, # 'registration' or 'email_change' "iat": now, - "exp": expire + "exp": expire, } - + token = jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM) - + return token - + except Exception as e: raise ServerException(f"Failed to create email verification token: {str(e)}") -async def get_token(credentials: Optional[HTTPAuthorizationCredentials] = Depends(HTTPBearer(auto_error=False))) -> str: + +async def get_token( + credentials: HTTPAuthorizationCredentials | None = Depends(HTTPBearer(auto_error=False)), +) -> str: if not credentials: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) return credentials.credentials -async def verify_session(sid: str, token: str, redis_client) -> Dict[str, Any]: + +async def verify_session(sid: str, token: str, redis_client) -> dict[str, Any]: try: redis_key = f"session:{sid}" raw = await redis_client.get(redis_key) @@ -114,161 +126,166 @@ async def verify_session(sid: str, token: str, redis_client) -> Dict[str, Any]: raise ValueError("Invalid or expired session") try: session_data = ast.literal_eval(raw) - except (ValueError, SyntaxError): + except ValueError, SyntaxError: logger.error(f"Invalid session data: {raw}") raise ValueError("Invalid session data") - + if session_data.get("access_token") and session_data.get("access_token") != token: logger.error(f"Token mismatch: {session_data.get('access_token')} != {token}") - raise JWTError("Token mismatch") + raise JWTError("Token mismatch") return session_data except JWTError as e: logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) except Exception as e: logger.error(f"Failed to verify session: {e}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired session", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) -async def verify_token(token: str = Depends(get_token), redis_client = Depends(get_redis), db: AsyncSession = Depends(get_db)) -> Dict[str, Any]: + +async def verify_token( + token: str = Depends(get_token), + redis_client=Depends(get_redis), + db: AsyncSession = Depends(get_db), +) -> dict[str, Any]: try: payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) sid = payload.get("sid") if not sid: raise ValueError("Missing session ID") - session_data = await verify_session(sid, token, redis_client) - + await verify_session(sid, token, redis_client) + user = await db.execute(select(Users.status).where(Users.id == payload.get("sub"))) user_status = user.scalar_one_or_none() if user_status is not None and not user_status: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail="Account is disabled" - ) - + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Account is disabled") + return payload - + except JWTError as e: logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) except ValueError as e: logger.warning(f"Token validation error: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) -async def verify_password_reset_token(token: str = Depends(get_token)) -> Dict[str, Any]: + +async def verify_password_reset_token(token: str = Depends(get_token)) -> dict[str, Any]: """Verify password reset token""" try: payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) - + if payload.get("token_type") != "password_reset": raise ValueError("Invalid token type") - + # Check if force change password if not payload.get("force_change_password"): raise ValueError("Token not authorized for password reset") if payload.get("exp") < datetime.now().astimezone().timestamp(): raise ValueError("Token expired") - + user_id = payload.get("sub") email = payload.get("email") - + if not user_id or not email: raise ValueError("Invalid token payload") - + return { "token": token, "sub": user_id, "email": email, "exp": payload.get("exp"), - "iat": payload.get("iat") + "iat": payload.get("iat"), } - + except JWTError as e: - logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") + logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) except ValueError as e: - logger.error(f"Failed to verify password reset token: {str(e)}") + logger.error(f"Failed to verify password reset token: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) -async def verify_email_verification_token(token: str = Depends(get_token)) -> Dict[str, Any]: + +async def verify_email_verification_token(token: str = Depends(get_token)) -> dict[str, Any]: """Verify email verification token""" try: payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) - + if payload.get("token_type") != "email_verification": raise ValueError("Invalid token type") - + verification_type = payload.get("verification_type") if verification_type not in ["registration", "email_change"]: raise ValueError("Invalid verification type") if payload.get("exp") < datetime.now().astimezone().timestamp(): raise ValueError("Token expired") - + user_id = payload.get("sub") email = payload.get("email") - + if not user_id or not email: raise ValueError("Invalid token payload") - + return { "token": token, "sub": user_id, "email": email, "verification_type": verification_type, "exp": payload.get("exp"), - "iat": payload.get("iat") + "iat": payload.get("iat"), } - + except JWTError as e: - logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") + logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) except ValueError as e: - logger.error(f"Failed to verify email verification token: {str(e)}") + logger.error(f"Failed to verify email verification token: {str(e)}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", - headers={"WWW-Authenticate": "Bearer"} + headers={"WWW-Authenticate": "Bearer"}, ) -async def extend_session_ttl(redis_client, session_id: str, session_data: Dict[str, Any]) -> None: + +async def extend_session_ttl(redis_client, session_id: str, session_data: dict[str, Any]) -> None: """Extend session TTL and update last activity time""" try: # Update last activity time using system timezone with timezone info session_data["last_activity"] = datetime.now().astimezone().isoformat() - + # Reset TTL, start from current time ttl = settings.SESSION_EXPIRE_MINUTES * 60 await redis_client.setex(f"session:{session_id}", ttl, str(session_data)) - + except Exception as e: logger.error(f"Failed to extend session TTL: {e}") @@ -277,41 +294,34 @@ async def clear_user_all_sessions( db: AsyncSession, redis_client: redis.Redis, user_id: str, - exclude_session_id: Optional[str] = None, + exclude_session_id: str | None = None, ) -> bool: """Logout user from all devices, optionally keeping one active session.""" try: - update_query = ( - update(UserSessions) - .where( - UserSessions.user_id == user_id, - UserSessions.is_active == True - ) - .values(is_active=False) - ) - if exclude_session_id: - update_query = update_query.where(UserSessions.id != exclude_session_id) - - await db.execute(update_query) - await db.commit() - sessions_query = select(UserSessions.id).where( UserSessions.user_id == user_id, - UserSessions.is_active == False + UserSessions.is_active.is_(True), ) if exclude_session_id: sessions_query = sessions_query.where(UserSessions.id != exclude_session_id) result = await db.execute(sessions_query) - session_ids = result.scalars().all() - - if session_ids: - redis_keys = [] - for sid in session_ids: - redis_keys.append(f"session:{sid}") - redis_keys.append(f"csrf:{sid}") - await redis_client.delete(*redis_keys) - + session_ids = list(result.scalars().all()) + + if not session_ids: + return True + + await db.execute( + update(UserSessions).where(UserSessions.id.in_(session_ids)).values(is_active=False) + ) + await db.commit() + + redis_keys = [] + for sid in session_ids: + redis_keys.append(f"session:{sid}") + redis_keys.append(f"csrf:{sid}") + await redis_client.delete(*redis_keys) + return True except Exception as e: - raise ServerException(f"Failed to logout all devices: {e}") \ No newline at end of file + raise ServerException(f"Failed to logout all devices: {e}") diff --git a/backend/core/telemetry.py b/backend/core/telemetry.py index 62ca178..50af2a6 100644 --- a/backend/core/telemetry.py +++ b/backend/core/telemetry.py @@ -1,5 +1,6 @@ import logging from urllib.parse import urljoin + from fastapi import FastAPI from opentelemetry import trace from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter @@ -10,6 +11,7 @@ from opentelemetry.sdk.resources import Resource from opentelemetry.sdk.trace import ReadableSpan, Span, SpanProcessor, TracerProvider from opentelemetry.sdk.trace.export import BatchSpanProcessor + from core.config import SKIP_METHODS, otel_excluded_urls, settings from core.database import async_engine, engine @@ -72,9 +74,7 @@ def _build_resource() -> Resource: { "service.name": settings.PROJECT_NAME, "service.version": settings.PROJECT_VERSION, - "deployment.environment": ( - "development" if settings.DEBUG_MODE else "production" - ), + "deployment.environment": ("development" if settings.DEBUG_MODE else "production"), } ) @@ -91,20 +91,14 @@ def setup_telemetry(app: FastAPI) -> None: resource = _build_resource() provider = TracerProvider(resource=resource) - exporter = OTLPSpanExporter( - endpoint=_traces_endpoint(settings.OTEL_EXPORTER_OTLP_ENDPOINT) - ) + exporter = OTLPSpanExporter(endpoint=_traces_endpoint(settings.OTEL_EXPORTER_OTLP_ENDPOINT)) provider.add_span_processor( - SkipSpanProcessor( - BatchSpanProcessor(exporter, schedule_delay_millis=1000) - ) + SkipSpanProcessor(BatchSpanProcessor(exporter, schedule_delay_millis=1000)) ) trace.set_tracer_provider(provider) FastAPIInstrumentor.instrument_app(app, excluded_urls=EXCLUDED_URLS) - SQLAlchemyInstrumentor().instrument( - engines=[async_engine.sync_engine, engine] - ) + SQLAlchemyInstrumentor().instrument(engines=[async_engine.sync_engine, engine]) RedisInstrumentor().instrument() logger.info("OpenTelemetry tracing enabled") diff --git a/backend/extensions/__init__.py b/backend/extensions/__init__.py index 0b83571..b850820 100644 --- a/backend/extensions/__init__.py +++ b/backend/extensions/__init__.py @@ -1,7 +1,8 @@ from .exception_handler import add_exception_handlers from .smtp import add_smtp + def register_extensions(app): # Add new extensions imports below. add_exception_handlers(app) - add_smtp(app) \ No newline at end of file + add_smtp(app) diff --git a/backend/extensions/exception_handler.py b/backend/extensions/exception_handler.py index 8abf96b..225e8e3 100644 --- a/backend/extensions/exception_handler.py +++ b/backend/extensions/exception_handler.py @@ -1,27 +1,23 @@ -from utils.response import APIResponse -from fastapi.responses import JSONResponse -from fastapi import FastAPI, Request, HTTPException +from fastapi import FastAPI, HTTPException, Request from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse + +from utils.response import APIResponse + def add_exception_handlers(app: FastAPI): @app.exception_handler(HTTPException) async def http_exception_handler(request: Request, exc: HTTPException): status_code = exc.status_code - + if isinstance(exc.detail, dict) and "code" in exc.detail and "message" in exc.detail: - return JSONResponse( - status_code=status_code, - content=exc.detail - ) - + return JSONResponse(status_code=status_code, content=exc.detail) + message = exc.detail if exc.detail else "HTTP Error" data = None - + resp = APIResponse(code=status_code, message=message, data=data) - return JSONResponse( - status_code=status_code, - content=resp.dict(exclude_none=True) - ) + return JSONResponse(status_code=status_code, content=resp.dict(exclude_none=True)) @app.exception_handler(RequestValidationError) async def validation_exception_handler(request: Request, exc: RequestValidationError): @@ -30,15 +26,9 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE field = ".".join([str(loc) for loc in err["loc"] if isinstance(loc, (str, int))]) errors[field] = err["msg"] resp = APIResponse(code=422, message="Validation Error", data=errors) - return JSONResponse( - status_code=422, - content=resp.dict(exclude_none=True) - ) + return JSONResponse(status_code=422, content=resp.dict(exclude_none=True)) @app.exception_handler(Exception) async def internal_server_error_handler(request: Request, exc: Exception): resp = APIResponse(code=500, message="Internal Server Error", data=None) - return JSONResponse( - status_code=500, - content=resp.dict(exclude_none=True) - ) \ No newline at end of file + return JSONResponse(status_code=500, content=resp.dict(exclude_none=True)) diff --git a/backend/extensions/smtp.py b/backend/extensions/smtp.py index d52f698..1fc9a96 100644 --- a/backend/extensions/smtp.py +++ b/backend/extensions/smtp.py @@ -1,13 +1,15 @@ -import ssl -import smtplib import logging -from fastapi import FastAPI -from core.config import settings +import smtplib +import ssl +from collections.abc import Iterable from dataclasses import dataclass -from email.mime.text import MIMEText -from typing import Iterable, Optional from email.message import EmailMessage from email.mime.multipart import MIMEMultipart +from email.mime.text import MIMEText + +from fastapi import FastAPI + +from core.config import settings from utils.custom_exception import SMTPNotConfiguredException logger = logging.getLogger("smtp") @@ -18,9 +20,9 @@ class SMTPSettings: enabled: bool host: str port: int - username: Optional[str] - password: Optional[str] - from_email: Optional[str] + username: str | None + password: str | None + from_email: str | None from_name: str encryption: str @@ -65,7 +67,9 @@ def _open(self, timeout: int = 30) -> smtplib.SMTP: context = ssl.create_default_context() if enc == "ssl": - client: smtplib.SMTP = smtplib.SMTP_SSL(self._cfg.host, self._cfg.port, timeout=timeout, context=context) + client: smtplib.SMTP = smtplib.SMTP_SSL( + self._cfg.host, self._cfg.port, timeout=timeout, context=context + ) else: client = smtplib.SMTP(self._cfg.host, self._cfg.port, timeout=timeout) @@ -90,9 +94,9 @@ def send_text( to_emails: Iterable[str], subject: str, body: str, - html_body: Optional[str] = None, - from_email: Optional[str] = None, - from_name: Optional[str] = None, + html_body: str | None = None, + from_email: str | None = None, + from_name: str | None = None, timeout: int = 30, ) -> None: self._validate() @@ -106,7 +110,7 @@ def send_text( msg["Subject"] = subject msg["From"] = f"{sender_name} <{sender_email}>" msg["To"] = ", ".join(list(to_emails)) - + # Add plain text and HTML parts part1 = MIMEText(body, "plain", "utf-8") part2 = MIMEText(html_body, "html", "utf-8") @@ -138,7 +142,7 @@ def build_smtp_settings() -> SMTPSettings: # Global singleton instance -_SMTP_MAILER: Optional[SMTPMailer] = None +_SMTP_MAILER: SMTPMailer | None = None def get_mailer() -> SMTPMailer: @@ -150,7 +154,7 @@ def get_mailer() -> SMTPMailer: if _SMTP_MAILER is None: cfg = build_smtp_settings() _SMTP_MAILER = SMTPMailer(cfg) - + if cfg.enabled: logger.info( "SMTP enabled: host=%s port=%s encryption=%s from=%s", @@ -161,7 +165,7 @@ def get_mailer() -> SMTPMailer: ) else: logger.info("SMTP disabled") - + return _SMTP_MAILER @@ -170,4 +174,4 @@ def add_smtp(app: FastAPI) -> None: Initialize SMTP mailer and register to app.state. """ mailer = get_mailer() - app.state.smtp = mailer \ No newline at end of file + app.state.smtp = mailer diff --git a/backend/main.py b/backend/main.py index 9afb809..207a50b 100644 --- a/backend/main.py +++ b/backend/main.py @@ -1,14 +1,18 @@ from core.config import settings, setup_logging + setup_logging("logging_config.yaml") -from api import api_router +from contextlib import asynccontextmanager + from fastapi import FastAPI -from core.redis import init_redis + +from api import api_router from core.database import init_db +from core.redis import init_redis from core.telemetry import setup_telemetry, shutdown_telemetry -from contextlib import asynccontextmanager from extensions import register_extensions from middleware import register_middlewares -from schedule import scheduler, register_schedules +from schedule import register_schedules, scheduler + # Lifespan event handler @asynccontextmanager @@ -21,6 +25,7 @@ async def lifespan(app: FastAPI): scheduler.shutdown() shutdown_telemetry() + # Control docs exposure by environment variable DEBUG_MODE docs_url = "/" if settings.DEBUG_MODE else None redoc_url = "/redoc" if settings.DEBUG_MODE else None @@ -34,7 +39,7 @@ async def lifespan(app: FastAPI): lifespan=lifespan, docs_url=docs_url, redoc_url=redoc_url, - openapi_url=openapi_url, + openapi_url=openapi_url, ) # Register all extensions @@ -46,9 +51,11 @@ async def lifespan(app: FastAPI): # Register all API routes with a global prefix '/api' app.include_router(api_router, prefix="/api") + # Health check endpoint @app.get("/healthz", include_in_schema=False) async def healthz(): return {"status": "ok"} -setup_telemetry(app) \ No newline at end of file + +setup_telemetry(app) diff --git a/backend/middleware/__init__.py b/backend/middleware/__init__.py index 85e889a..e98fb17 100644 --- a/backend/middleware/__init__.py +++ b/backend/middleware/__init__.py @@ -1,10 +1,11 @@ from fastapi import FastAPI + from .cors import add_cors_middleware -from .request_logging import add_request_logging_middleware from .rate_limiter import add_rate_limiter_middleware +from .request_logging import add_request_logging_middleware def register_middlewares(app: FastAPI): add_rate_limiter_middleware(app) add_cors_middleware(app) - add_request_logging_middleware(app) \ No newline at end of file + add_request_logging_middleware(app) diff --git a/backend/middleware/cors.py b/backend/middleware/cors.py index 122678b..e846d9a 100644 --- a/backend/middleware/cors.py +++ b/backend/middleware/cors.py @@ -1,13 +1,16 @@ -from core.config import settings +from urllib.parse import urlparse + from fastapi import FastAPI, Request from fastapi.responses import JSONResponse from starlette.middleware.base import BaseHTTPMiddleware -from urllib.parse import urlparse -class CORSMiddleware(BaseHTTPMiddleware): +from core.config import settings + + +class CORSMiddleware(BaseHTTPMiddleware): def __init__(self, app): super().__init__(app) - + self.allowed_hosts = [ f"{settings.HOSTNAME}:{settings.BACKEND_PORT}", f"localhost:{settings.BACKEND_PORT}", @@ -18,13 +21,13 @@ def __init__(self, app): ] self.allowed_methods = "GET, POST, PUT, DELETE, PATCH, OPTIONS, HEAD" - + self.generate_cors_origins() - + self.whitelist_paths = [ # Add whitelist paths here (e.g. "/api/example/") ] - + def generate_cors_origins(self): """Generate HTTP and HTTPS versions of the sources""" self.cors_origins = [] @@ -32,18 +35,17 @@ def generate_cors_origins(self): if host == "*": self.cors_origins.append("*") continue - self.cors_origins.extend([ - f"http://{host}", - f"https://{host}" - ]) - + self.cors_origins.extend([f"http://{host}", f"https://{host}"]) + def is_whitelist_path(self, path: str) -> bool: """Check if path is in whitelist""" for whitelist_path in self.whitelist_paths: - if path == whitelist_path or (whitelist_path.endswith("/") and path.startswith(whitelist_path)): + if path == whitelist_path or ( + whitelist_path.endswith("/") and path.startswith(whitelist_path) + ): return True return False - + def _normalize_origin(self, origin: str) -> str: """Normalize origin for comparison""" try: @@ -54,27 +56,27 @@ def _normalize_origin(self, origin: str) -> str: return f"{parsed.scheme}://{parsed.hostname}:{port}" except Exception: return origin.lower() - + def is_allowed_origin(self, origin: str) -> bool: """Check if origin is allowed, with normalized comparison""" if not origin: return False - + normalized_origin = self._normalize_origin(origin) - + for allowed in self.cors_origins: if allowed == "*": continue - + try: normalized_allowed = self._normalize_origin(allowed) if normalized_origin == normalized_allowed: return True except Exception: pass - + return origin.lower() in [a.lower() for a in self.cors_origins if a != "*"] - + def get_allowed_headers(self, request: Request) -> str: """Get allowed headers string for CORS response""" allowed_headers = [ @@ -86,15 +88,15 @@ def get_allowed_headers(self, request: Request) -> str: "pragma", "x-requested-with", ] - + requested_headers = request.headers.get("access-control-request-headers", "") if requested_headers: requested_list = [h.strip().lower() for h in requested_headers.split(",")] allowed_headers.extend(requested_list) - + allowed_headers = sorted(list(set(allowed_headers))) return ", ".join(allowed_headers) - + def _request_hostname(self, request: Request) -> str | None: """Hostname the browser actually called (nginx Host / X-Forwarded-Host).""" raw = request.headers.get("x-forwarded-host") or request.headers.get("host") @@ -116,68 +118,65 @@ def _is_same_host_origin(self, origin: str, request: Request) -> bool: return False return origin_host.lower() == request_host - def _should_allow_origin( - self, origin: str, is_whitelisted: bool, request: Request - ) -> bool: + def _should_allow_origin(self, origin: str, is_whitelisted: bool, request: Request) -> bool: """Check if origin should be allowed""" if is_whitelisted: return bool(origin) return bool( origin - and ( - self.is_allowed_origin(origin) - or self._is_same_host_origin(origin, request) - ) + and (self.is_allowed_origin(origin) or self._is_same_host_origin(origin, request)) ) - + def _set_cors_origin_headers(self, headers: dict, origin: str): """Set CORS origin and credentials headers""" if origin: headers["Access-Control-Allow-Origin"] = origin headers["Access-Control-Allow-Credentials"] = "true" - + def handle_preflight(self, request: Request, origin: str, is_whitelisted: bool) -> JSONResponse: """Handle OPTIONS preflight requests""" if not self._should_allow_origin(origin, is_whitelisted, request): return JSONResponse(status_code=200, content={}) - + allowed_headers_str = self.get_allowed_headers(request) - + headers = { "Access-Control-Allow-Methods": self.allowed_methods, "Access-Control-Allow-Headers": allowed_headers_str, "Access-Control-Max-Age": "3600", "Vary": "Origin", } - + self._set_cors_origin_headers(headers, origin) return JSONResponse(status_code=200, content={}, headers=headers) - - def add_cors_headers( - self, response, origin: str, is_whitelisted: bool, request: Request - ): + + def add_cors_headers(self, response, origin: str, is_whitelisted: bool, request: Request): """Add CORS headers to response""" response.headers["Vary"] = "Origin" - + if not self._should_allow_origin(origin, is_whitelisted, request): return - + self._set_cors_origin_headers(response.headers, origin) response.headers["Access-Control-Allow-Methods"] = self.allowed_methods - response.headers["Access-Control-Allow-Headers"] = "content-type, authorization, accept, accept-language, cache-control, pragma, x-requested-with" - + response.headers["Access-Control-Allow-Headers"] = ( + "content-type, authorization, accept, accept-language, cache-control, " + "pragma, x-requested-with" + ) + async def dispatch(self, request: Request, call_next): """Handle CORS requests""" path = request.url.path origin = request.headers.get("origin") is_whitelisted = self.is_whitelist_path(path) - + if request.method == "OPTIONS": return self.handle_preflight(request, origin, is_whitelisted) - + response = await call_next(request) self.add_cors_headers(response, origin, is_whitelisted, request) return response + def add_cors_middleware(app: FastAPI): - app.add_middleware(CORSMiddleware) \ No newline at end of file + app.add_middleware(CORSMiddleware) diff --git a/backend/middleware/rate_limiter.py b/backend/middleware/rate_limiter.py index cac99ba..22813bf 100644 --- a/backend/middleware/rate_limiter.py +++ b/backend/middleware/rate_limiter.py @@ -1,12 +1,13 @@ import logging -from utils import get_real_ip -from core.redis import get_redis -from core.config import settings -from utils.response import APIResponse + from fastapi import Request, status from fastapi.responses import JSONResponse from starlette.middleware.base import BaseHTTPMiddleware -from core.config import SKIP_METHODS, SKIP_PATHS + +from core.config import SKIP_METHODS, SKIP_PATHS, settings +from core.redis import get_redis +from utils import get_real_ip +from utils.response import APIResponse logger = logging.getLogger("rate_limiter") @@ -14,6 +15,7 @@ RATE_LIMIT_WINDOW_SECONDS = settings.RATE_LIMIT_WINDOW_SECONDS BLOCK_TIME_SECONDS = settings.BLOCK_TIME_SECONDS + class RateLimiterMiddleware(BaseHTTPMiddleware): def __init__(self, app): super().__init__(app) @@ -25,11 +27,12 @@ def _configure_endpoint_limits(self): """ Configure rate limits for specific endpoints. Only endpoints listed here will override the default rate limit settings. - + Format: { "path": { - "limit": (count, window_seconds) or None, # (allowed_requests, time_window) or None to disable rate limiting + "limit": (count, window_seconds) or None, + # (allowed_requests, time_window) or None to disable rate limiting "status_codes": [int, ...], # Optional: count only these status codes "clear_on_success": bool # Optional: clear counter on 2xx responses } @@ -37,28 +40,24 @@ def _configure_endpoint_limits(self): """ self.endpoint_rate_limits = { # Add endpoint rate limits here to override default settings - "/api/auth/token": { - "limit": (60, 60), - "status_codes": None, - "clear_on_success": False - }, + "/api/auth/token": {"limit": (60, 60), "status_codes": None, "clear_on_success": False}, "/api/auth/login": { "limit": (5, 30), "status_codes": [401, 403], - "clear_on_success": True + "clear_on_success": True, }, "/api/roles/permissions": { "limit": (60, 60), "status_codes": None, - "clear_on_success": False + "clear_on_success": False, }, "/api/debug/test-ip": { "limit": (10, 60), "status_codes": None, - "clear_on_success": False + "clear_on_success": False, }, } - + def _get_rate_limit_config(self, path: str) -> dict: """ Get rate limit configuration for a path. @@ -67,12 +66,12 @@ def _get_rate_limit_config(self, path: str) -> dict: custom_config = self.endpoint_rate_limits.get(path) if custom_config: return custom_config - + # Default configuration for all endpoints return { - "limit": (RATE_LIMIT, RATE_LIMIT_WINDOW_SECONDS), + "limit": (RATE_LIMIT, RATE_LIMIT_WINDOW_SECONDS), "status_codes": None, - "clear_on_success": False + "clear_on_success": False, } async def dispatch(self, request: Request, call_next): @@ -94,26 +93,28 @@ async def dispatch(self, request: Request, call_next): # Get rate limit config rate_limit_config = self._get_rate_limit_config(path) - + # If limit is None, skip rate limiting for this endpoint if rate_limit_config.get("limit") is None: return await call_next(request) - + api_block_key = f"block:api:{ip}:{path}" - + # Check if IP is blocked for this endpoint is_blocked = await redis.get(api_block_key) if is_blocked: resp = APIResponse[None](code=429, message="Too many requests. Try again later.") - return JSONResponse(status_code=status.HTTP_429_TOO_MANY_REQUESTS, - content=resp.model_dump(exclude_none=True)) + return JSONResponse( + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + content=resp.model_dump(exclude_none=True), + ) response = await call_next(request) status_codes = rate_limit_config.get("status_codes") limit_count, window_seconds = rate_limit_config["limit"] clear_on_success = rate_limit_config.get("clear_on_success", False) is_success = 200 <= response.status_code < 300 - + should_count = False if not status_codes: should_count = True @@ -131,13 +132,16 @@ async def dispatch(self, request: Request, call_next): if api_fails >= limit_count: await redis.set(api_block_key, 1, ex=BLOCK_TIME_SECONDS) await redis.delete(api_fail_key) - resp = APIResponse[None](code=429, message="Too many requests. Try again later.") - return JSONResponse(status_code=status.HTTP_429_TOO_MANY_REQUESTS, - content=resp.model_dump(exclude_none=True)) + resp = APIResponse[None]( + code=429, message="Too many requests. Try again later." + ) + return JSONResponse( + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + content=resp.model_dump(exclude_none=True), + ) except Exception as e: logger.error( - f"Rate limiter error: method={method} path={path} " - f"error={e} ipAddress={ip}" + f"Rate limiter error: method={method} path={path} error={e} ipAddress={ip}" ) elif clear_on_success and is_success: try: @@ -157,5 +161,6 @@ async def dispatch(self, request: Request, call_next): return response return await call_next(request) + def add_rate_limiter_middleware(app): - app.add_middleware(RateLimiterMiddleware) \ No newline at end of file + app.add_middleware(RateLimiterMiddleware) diff --git a/backend/middleware/request_logging.py b/backend/middleware/request_logging.py index 8851f54..b6b6f1e 100644 --- a/backend/middleware/request_logging.py +++ b/backend/middleware/request_logging.py @@ -1,10 +1,11 @@ import logging import time -from utils import get_real_ip + from fastapi import FastAPI, Request from starlette.middleware.base import BaseHTTPMiddleware -from core.config import SKIP_METHODS, SKIP_PATHS -from core.config import settings + +from core.config import SKIP_METHODS, SKIP_PATHS, settings +from utils import get_real_ip from utils.log_sanitize import ( format_log_value, sanitize_body, @@ -52,7 +53,7 @@ async def dispatch(self, request: Request, call_next): logger.info( f"API Request: method={method} path={path} ipAddress={client_ip} " - f"user-agent=\"{user_agent}\"{request_extra}" + f'user-agent="{user_agent}"{request_extra}' ) started = time.perf_counter() @@ -75,6 +76,7 @@ async def dispatch(self, request: Request, call_next): settings.LOG_HTTP_BODY_MAX_BYTES, ), ) + # Replay on the same response so multiple Set-Cookie headers stay intact. # dict(response.headers) collapses them into one and drops csrf_token. async def _replay(content: bytes = body): @@ -85,7 +87,7 @@ async def _replay(content: bytes = body): logger.info( f"API Response: method={method} path={path} ipAddress={client_ip} " f"status_code={response.status_code} duration={duration_ms:.1f}ms " - f"user-agent=\"{user_agent}\"{response_extra}" + f'user-agent="{user_agent}"{response_extra}' ) return response diff --git a/backend/migrations/env.py b/backend/migrations/env.py index 6472674..ea7096b 100755 --- a/backend/migrations/env.py +++ b/backend/migrations/env.py @@ -1,15 +1,12 @@ -import sys import os -sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) - -import models -from logging.config import fileConfig +import sys -from sqlalchemy import engine_from_config -from sqlalchemy import pool +sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) +from logging.config import fileConfig from alembic import context +from sqlalchemy import pool # this is the Alembic Config object, which provides # access to the values within the .ini file in use. @@ -25,6 +22,7 @@ # from myapp import mymodel # target_metadata = mymodel.Base.metadata from core.database import Base + target_metadata = Base.metadata # other values from the config, defined by the needs of env.py, @@ -34,14 +32,17 @@ from core.config import settings + def include_object(object, name, type_, reflected, compare_to): """ - Filter out APScheduler related tables to avoid Alembic automatically generating migrations to delete these tables + Filter out APScheduler related tables to avoid Alembic automatically generating + migrations to delete these tables """ if type_ == "table" and name == "apscheduler_jobs": return False return True + def run_migrations_offline() -> None: """Run migrations in 'offline' mode. @@ -75,10 +76,11 @@ def run_migrations_online() -> None: """ from sqlalchemy import create_engine + connectable = create_engine(settings.DATABASE_URL, poolclass=pool.NullPool) with connectable.connect() as connection: context.configure( - connection=connection, + connection=connection, target_metadata=target_metadata, include_object=include_object, ) @@ -90,4 +92,4 @@ def run_migrations_online() -> None: if context.is_offline_mode(): run_migrations_offline() else: - run_migrations_online() \ No newline at end of file + run_migrations_online() diff --git a/backend/models/__init__.py b/backend/models/__init__.py index 1dae284..046b3ef 100644 --- a/backend/models/__init__.py +++ b/backend/models/__init__.py @@ -1,10 +1,22 @@ # Add new SQLAlchemy model imports below. -from .users import Users +from .email_verification_tokens import EmailVerificationTokens from .login_logs import LoginLogs -from .user_sessions import UserSessions from .password_reset_tokens import PasswordResetTokens -from .email_verification_tokens import EmailVerificationTokens -from .roles import Roles -from .role_mapper import RoleMapper from .role_attributes import RoleAttributes -from .role_attributes_mapper import RoleAttributesMapper \ No newline at end of file +from .role_attributes_mapper import RoleAttributesMapper +from .role_mapper import RoleMapper +from .roles import Roles +from .user_sessions import UserSessions +from .users import Users + +__all__ = [ + "EmailVerificationTokens", + "LoginLogs", + "PasswordResetTokens", + "RoleAttributes", + "RoleAttributesMapper", + "RoleMapper", + "Roles", + "UserSessions", + "Users", +] diff --git a/backend/models/email_verification_tokens.py b/backend/models/email_verification_tokens.py index 6b61963..9aec015 100644 --- a/backend/models/email_verification_tokens.py +++ b/backend/models/email_verification_tokens.py @@ -1,20 +1,26 @@ +from sqlalchemy import TIMESTAMP, Boolean, Column, ForeignKey, String, Text, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, Boolean, TIMESTAMP, ForeignKey, Text, text + class EmailVerificationTokens(Base): __tablename__ = "email_verification_tokens" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) user_id = Column(String(36), ForeignKey("users.id"), nullable=False, index=True) email = Column(String(50), nullable=False, index=True) token = Column(Text, nullable=False, unique=True, index=True) token_type = Column(String(20), nullable=False, index=True) is_used = Column(Boolean, nullable=False, default=False) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) expires_at = Column(TIMESTAMP, nullable=False) - + # Relationships - user = relationship("Users", back_populates="email_verification_tokens") \ No newline at end of file + user = relationship("Users", back_populates="email_verification_tokens") diff --git a/backend/models/login_logs.py b/backend/models/login_logs.py index 18e27aa..f547853 100644 --- a/backend/models/login_logs.py +++ b/backend/models/login_logs.py @@ -1,11 +1,13 @@ +from sqlalchemy import TIMESTAMP, Boolean, Column, ForeignKey, String, Text, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, Boolean, TIMESTAMP, Text, ForeignKey, text + class LoginLogs(Base): __tablename__ = "login_logs" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) user_id = Column(String(36), ForeignKey("users.id"), nullable=True, index=True) email = Column(String(50), nullable=False, index=True) @@ -13,8 +15,12 @@ class LoginLogs(Base): user_agent = Column(Text, nullable=False) is_success = Column(Boolean, nullable=False) failure_reason = Column(String(255), nullable=True) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) - + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) + # Relationships - user = relationship("Users", back_populates="login_logs") \ No newline at end of file + user = relationship("Users", back_populates="login_logs") diff --git a/backend/models/password_reset_tokens.py b/backend/models/password_reset_tokens.py index c1a9b2e..4e2b9d1 100644 --- a/backend/models/password_reset_tokens.py +++ b/backend/models/password_reset_tokens.py @@ -1,18 +1,24 @@ +from sqlalchemy import TIMESTAMP, Boolean, Column, ForeignKey, String, Text, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, Boolean, TIMESTAMP, ForeignKey, Text, text + class PasswordResetTokens(Base): __tablename__ = "password_reset_tokens" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) user_id = Column(String(36), ForeignKey("users.id"), nullable=False, index=True) token = Column(Text, nullable=False, unique=True, index=True) is_used = Column(Boolean, nullable=False, default=False) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) expires_at = Column(TIMESTAMP, nullable=False) - + # Relationships - user = relationship("Users", back_populates="password_reset_tokens") \ No newline at end of file + user = relationship("Users", back_populates="password_reset_tokens") diff --git a/backend/models/role_attributes.py b/backend/models/role_attributes.py index fabd54b..e2c3b87 100644 --- a/backend/models/role_attributes.py +++ b/backend/models/role_attributes.py @@ -1,17 +1,23 @@ +from sqlalchemy import TIMESTAMP, Column, String, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, TIMESTAMP, text + class RoleAttributes(Base): __tablename__ = "role_attributes" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) name = Column(String(100), nullable=False, unique=True, index=True) group = Column(String(100), nullable=True, index=True) category = Column(String(100), nullable=True, index=True) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) - + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) + # Relationships - role_mappings = relationship("RoleAttributesMapper", back_populates="attribute") \ No newline at end of file + role_mappings = relationship("RoleAttributesMapper", back_populates="attribute") diff --git a/backend/models/role_attributes_mapper.py b/backend/models/role_attributes_mapper.py index aecbe5d..3b67907 100644 --- a/backend/models/role_attributes_mapper.py +++ b/backend/models/role_attributes_mapper.py @@ -1,16 +1,22 @@ -from core.database import Base +from sqlalchemy import TIMESTAMP, Boolean, Column, ForeignKey, String, text from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, Boolean, ForeignKey, TIMESTAMP, text + +from core.database import Base + class RoleAttributesMapper(Base): __tablename__ = "role_attributes_mapper" - + role_id = Column(String(36), ForeignKey("roles.id"), primary_key=True) attributes_id = Column(String(36), ForeignKey("role_attributes.id"), primary_key=True) value = Column(Boolean, nullable=False) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) - + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) + # Relationships role = relationship("Roles", back_populates="attribute_mappings") - attribute = relationship("RoleAttributes", back_populates="role_mappings") \ No newline at end of file + attribute = relationship("RoleAttributes", back_populates="role_mappings") diff --git a/backend/models/role_mapper.py b/backend/models/role_mapper.py index 7e1506a..b2a428a 100644 --- a/backend/models/role_mapper.py +++ b/backend/models/role_mapper.py @@ -1,15 +1,21 @@ -from core.database import Base +from sqlalchemy import TIMESTAMP, Column, ForeignKey, String, text from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, ForeignKey, TIMESTAMP, text + +from core.database import Base + class RoleMapper(Base): __tablename__ = "role_mapper" - + user_id = Column(String(36), ForeignKey("users.id"), primary_key=True) role_id = Column(String(36), ForeignKey("roles.id"), primary_key=True) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) - + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) + # Relationships user = relationship("Users", back_populates="role_mappings") - role = relationship("Roles", back_populates="user_mappings") \ No newline at end of file + role = relationship("Roles", back_populates="user_mappings") diff --git a/backend/models/roles.py b/backend/models/roles.py index 2cb1aab..04d2897 100644 --- a/backend/models/roles.py +++ b/backend/models/roles.py @@ -1,19 +1,25 @@ +from sqlalchemy import TIMESTAMP, Column, Integer, String, Text, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, Integer, String, Text, TIMESTAMP, text + class Roles(Base): __tablename__ = "roles" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) name = Column(String(100), nullable=False, unique=True, index=True) description = Column(Text, nullable=True) # Higher number = higher privilege. System super-admin uses 100. level = Column(Integer, nullable=False, default=1, server_default=text("1"), index=True) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) - + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) + # Relationships user_mappings = relationship("RoleMapper", back_populates="role") - attribute_mappings = relationship("RoleAttributesMapper", back_populates="role") \ No newline at end of file + attribute_mappings = relationship("RoleAttributesMapper", back_populates="role") diff --git a/backend/models/user_sessions.py b/backend/models/user_sessions.py index 7f995f0..fc94107 100644 --- a/backend/models/user_sessions.py +++ b/backend/models/user_sessions.py @@ -1,20 +1,26 @@ +from sqlalchemy import TIMESTAMP, Boolean, Column, ForeignKey, String, Text, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, Boolean, TIMESTAMP, Text, ForeignKey, text + class UserSessions(Base): __tablename__ = "user_sessions" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) user_id = Column(String(36), ForeignKey("users.id"), nullable=False, index=True) jwt_access_token = Column(Text, nullable=False) ip_address = Column(String(45), nullable=False) user_agent = Column(Text, nullable=False) is_active = Column(Boolean, nullable=False, default=True) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) expires_at = Column(TIMESTAMP, nullable=False) - + # Relationships - user = relationship("Users", back_populates="user_sessions") \ No newline at end of file + user = relationship("Users", back_populates="user_sessions") diff --git a/backend/models/users.py b/backend/models/users.py index 850c376..09b2c2b 100644 --- a/backend/models/users.py +++ b/backend/models/users.py @@ -1,11 +1,13 @@ +from sqlalchemy import TIMESTAMP, Boolean, Column, String, text +from sqlalchemy.orm import relationship from uuid_utils import uuid7 + from core.database import Base -from sqlalchemy.orm import relationship -from sqlalchemy import Column, String, Boolean, TIMESTAMP, text + class Users(Base): __tablename__ = "users" - + id = Column(String(36), primary_key=True, default=lambda: str(uuid7()), unique=True, index=True) email = Column(String(50), unique=True, nullable=False, index=True) first_name = Column(String(100), nullable=False) @@ -13,15 +15,19 @@ class Users(Base): phone = Column(String(20), nullable=False) hash_password = Column(String(255), nullable=True) status = Column(Boolean, nullable=False, default=True) - password_reset_required = Column(Boolean, nullable=False, default=False, server_default='0') - email_verified = Column(Boolean, nullable=False, default=False, server_default='0') + password_reset_required = Column(Boolean, nullable=False, default=False, server_default="0") + email_verified = Column(Boolean, nullable=False, default=False, server_default="0") pending_email = Column(String(50), nullable=True, index=True) - created_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP')) - updated_at = Column(TIMESTAMP, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')) - + created_at = Column(TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP")) + updated_at = Column( + TIMESTAMP, + nullable=False, + server_default=text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP"), + ) + # Relationships login_logs = relationship("LoginLogs", back_populates="user") user_sessions = relationship("UserSessions", back_populates="user") role_mappings = relationship("RoleMapper", back_populates="user") password_reset_tokens = relationship("PasswordResetTokens", back_populates="user") - email_verification_tokens = relationship("EmailVerificationTokens", back_populates="user") \ No newline at end of file + email_verification_tokens = relationship("EmailVerificationTokens", back_populates="user") diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 2af308a..fbe5565 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -78,6 +78,7 @@ precision = 2 [tool.ruff] target-version = "py314" line-length = 100 +indent-width = 4 src = ["."] exclude = [ ".venv", @@ -86,10 +87,19 @@ exclude = [ [tool.ruff.lint] select = ["E", "F", "I", "UP", "B"] -ignore = ["B008"] +ignore = [ + "B008", # FastAPI Depends() in default arguments + "B904", # raise ... from err in except (FastAPI handlers) + "E712", # SQLAlchemy columns need == True/False or .is_(), not `not column` +] [tool.ruff.lint.per-file-ignores] "tests/**/*.py" = ["B"] +"main.py" = ["E402"] +"core/config.py" = ["E402"] +"migrations/env.py" = ["E402"] +"utils/email_templates.py" = ["E501"] [tool.ruff.format] quote-style = "double" +indent-style = "space" diff --git a/backend/schedule/__init__.py b/backend/schedule/__init__.py index f9b2e15..e3d5aa1 100644 --- a/backend/schedule/__init__.py +++ b/backend/schedule/__init__.py @@ -1,18 +1,19 @@ -from core.database import engine -from .cleanup_tasks import CleanupTasks from apscheduler.executors.asyncio import AsyncIOExecutor -from apscheduler.schedulers.asyncio import AsyncIOScheduler from apscheduler.jobstores.sqlalchemy import SQLAlchemyJobStore +from apscheduler.schedulers.asyncio import AsyncIOScheduler + +from core.database import engine + +from .cleanup_tasks import CleanupTasks executors = { - 'default': AsyncIOExecutor(), -} -jobstores = { - 'default': SQLAlchemyJobStore(engine=engine) + "default": AsyncIOExecutor(), } +jobstores = {"default": SQLAlchemyJobStore(engine=engine)} scheduler = AsyncIOScheduler(jobstores=jobstores, executors=executors) cleanup_tasks = CleanupTasks() + def register_schedules(): # Add new schedules imports below. scheduler.add_job( @@ -22,7 +23,7 @@ def register_schedules(): minute=0, id="cleanup_expired_sessions", name="Cleanup Expired Sessions", - replace_existing=True + replace_existing=True, ) scheduler.add_job( cleanup_tasks.cleanup_expired_email_verifications, @@ -31,5 +32,5 @@ def register_schedules(): minute=0, id="cleanup_expired_email_verifications", name="Cleanup Expired Email Verifications", - replace_existing=True - ) \ No newline at end of file + replace_existing=True, + ) diff --git a/backend/schedule/cleanup_tasks.py b/backend/schedule/cleanup_tasks.py index 6241cb3..ac737c9 100644 --- a/backend/schedule/cleanup_tasks.py +++ b/backend/schedule/cleanup_tasks.py @@ -1,11 +1,14 @@ import logging from datetime import datetime -from sqlalchemy import delete, or_, select, update, func + +from sqlalchemy import delete, func, or_, select, update + from core.database import AsyncSessionLocal -from models.user_sessions import UserSessions from models.email_verification_tokens import EmailVerificationTokens +from models.user_sessions import UserSessions from models.users import Users + class CleanupTasks: def __init__(self): self.logger = logging.getLogger("schedule") @@ -18,12 +21,14 @@ async def cleanup_expired_sessions(self): delete(UserSessions).where( or_( UserSessions.expires_at < datetime.now().astimezone(), - UserSessions.is_active == False + UserSessions.is_active.is_(False), ) ) ) await db.commit() - self.logger.info(f"Cleaned up {expired_sessions.rowcount} expired or inactive sessions") + self.logger.info( + f"Cleaned up {expired_sessions.rowcount} expired or inactive sessions" + ) except Exception as e: self.logger.error(f"Failed to cleanup expired sessions: {e}") await db.rollback() @@ -46,8 +51,9 @@ async def cleanup_expired_email_verifications(self): ) expired_tokens = await db.execute( - delete(EmailVerificationTokens) - .where(EmailVerificationTokens.email.in_(expired_emails_subquery)) + delete(EmailVerificationTokens).where( + EmailVerificationTokens.email.in_(expired_emails_subquery) + ) ) await db.commit() @@ -56,4 +62,4 @@ async def cleanup_expired_email_verifications(self): ) except Exception as e: self.logger.error(f"Failed to cleanup expired email verifications: {e}") - await db.rollback() \ No newline at end of file + await db.rollback() diff --git a/backend/tests/api/account/test_controller.py b/backend/tests/api/account/test_controller.py index dd7ee30..ee01da0 100644 --- a/backend/tests/api/account/test_controller.py +++ b/backend/tests/api/account/test_controller.py @@ -1,7 +1,9 @@ +from unittest.mock import patch + import pytest from httpx import AsyncClient + from models.users import Users -from unittest.mock import patch class TestGetUserProfileAPI: @@ -87,9 +89,7 @@ async def test_get_user_profile_malformed_auth_header(self, client: AsyncClient) malformed_headers = ["InvalidFormat", "Bearer", "Basic token123", "Bearer ", ""] for header in malformed_headers: - response = await client.get( - "/api/account/profile", headers={"Authorization": header} - ) + response = await client.get("/api/account/profile", headers={"Authorization": header}) assert response.status_code == 401 @@ -331,7 +331,9 @@ async def test_change_password_success( # data field is excluded when None due to response_model_exclude_unset=True @pytest.mark.asyncio - async def test_change_password_success_response(self, client: AsyncClient, account_auth_headers: dict): + async def test_change_password_success_response( + self, client: AsyncClient, account_auth_headers: dict + ): """Test change password returns success response when service succeeds""" password_data = { "current_password": "AccountTestPassword123!", @@ -408,9 +410,7 @@ async def test_change_password_missing_fields( self, client: AsyncClient, account_test_user: Users, account_auth_headers: dict ): """Test password change with missing required fields""" - password_data = { - "current_password": "AccountTestPassword123!" - } # Missing new_password + password_data = {"current_password": "AccountTestPassword123!"} # Missing new_password response = await client.put( "/api/account/password", @@ -626,9 +626,7 @@ class TestControllerEdgeCases: """Test controller layer edge cases and error scenarios""" @pytest.mark.asyncio - async def test_large_payload_handling( - self, client: AsyncClient, account_auth_headers: dict - ): + async def test_large_payload_handling(self, client: AsyncClient, account_auth_headers: dict): """Test handling of large payloads""" # Test with very long strings large_data = { @@ -647,9 +645,7 @@ async def test_large_payload_handling( assert response.status_code == 200 @pytest.mark.asyncio - async def test_unicode_handling( - self, client: AsyncClient, account_auth_headers: dict - ): + async def test_unicode_handling(self, client: AsyncClient, account_auth_headers: dict): """Test handling of unicode characters""" unicode_data = { "first_name": "測試", @@ -686,9 +682,7 @@ async def test_special_characters_in_phone( assert data["data"]["phone"] == "+1 (555) 123-4567" @pytest.mark.asyncio - async def test_malformed_json_handling( - self, client: AsyncClient, account_auth_headers: dict - ): + async def test_malformed_json_handling(self, client: AsyncClient, account_auth_headers: dict): """Test handling of various malformed JSON scenarios""" malformed_cases = [ '{"first_name": "Test", "last_name":}', # Missing value diff --git a/backend/tests/api/account/test_schema.py b/backend/tests/api/account/test_schema.py index 073214d..4a8c244 100644 --- a/backend/tests/api/account/test_schema.py +++ b/backend/tests/api/account/test_schema.py @@ -1,8 +1,10 @@ -import pytest from datetime import datetime -from core.config import settings + +import pytest from pydantic import ValidationError -from api.account.schema import UserProfile, UserUpdate, PasswordChange + +from api.account.schema import PasswordChange, UserProfile, UserUpdate +from core.config import settings class TestUserProfile: @@ -232,9 +234,7 @@ def test_password_change_length_constraints(self): long_password = "a" * 51 # Exceeds max_length=50 with pytest.raises(ValidationError) as exc_info: - PasswordChange( - current_password="CurrentPass123!", new_password=long_password - ) + PasswordChange(current_password="CurrentPass123!", new_password=long_password) errors = exc_info.value.errors() assert any("max_length" in str(error) for error in errors) diff --git a/backend/tests/api/account/test_service.py b/backend/tests/api/account/test_service.py index 98b5fbb..5d36543 100644 --- a/backend/tests/api/account/test_service.py +++ b/backend/tests/api/account/test_service.py @@ -1,23 +1,27 @@ +from unittest.mock import AsyncMock, MagicMock, patch + import pytest -from models.users import Users -from core.security import verify_password -from unittest.mock import AsyncMock, patch, MagicMock from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession -from api.account.schema import UserUpdate, PasswordChange -from utils.custom_exception import AuthenticationException, ServerException, SMTPNotConfiguredException -from api.account.services import get_user_by_id, update_user_profile, change_password -from extensions.smtp import SMTPMailer + +from api.account.schema import PasswordChange, UserUpdate +from api.account.services import change_password, get_user_by_id, update_user_profile from core.config import settings +from core.security import verify_password +from extensions.smtp import SMTPMailer +from models.users import Users +from utils.custom_exception import ( + AuthenticationException, + ServerException, + SMTPNotConfiguredException, +) class TestGetUserById: """Test get_user_by_id service function""" @pytest.mark.asyncio - async def test_get_user_by_id_success( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_get_user_by_id_success(self, test_db_session: AsyncSession, test_user: Users): """Test successful user retrieval by ID""" user = await get_user_by_id(test_db_session, test_user.id) @@ -92,16 +96,12 @@ async def test_update_user_profile_success_partial_fields( assert email_change_requested is False @pytest.mark.asyncio - async def test_update_user_profile_user_not_found( - self, test_db_session: AsyncSession - ): + async def test_update_user_profile_user_not_found(self, test_db_session: AsyncSession): """Test profile update with non-existent user ID""" update_data = UserUpdate(first_name="Test") non_existent_id = "non-existent-id-123" - updated_user = await update_user_profile( - test_db_session, non_existent_id, update_data - ) + updated_user = await update_user_profile(test_db_session, non_existent_id, update_data) assert updated_user is None @@ -153,9 +153,7 @@ async def test_update_user_profile_database_commit_error( update_data = UserUpdate(first_name="Test") # Mock commit to raise an exception - with patch.object( - test_db_session, "commit", side_effect=SQLAlchemyError("Commit error") - ): + with patch.object(test_db_session, "commit", side_effect=SQLAlchemyError("Commit error")): with pytest.raises(SQLAlchemyError): await update_user_profile(test_db_session, test_user.id, update_data) @@ -243,9 +241,7 @@ class TestChangePassword: """Test change_password service function""" @pytest.mark.asyncio - async def test_change_password_success( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_change_password_success(self, test_db_session: AsyncSession, test_user: Users): """Test successful password change""" password_data = PasswordChange( current_password="TestPassword123!", @@ -255,9 +251,7 @@ async def test_change_password_success( mock_redis = AsyncMock() - success = await change_password( - test_db_session, test_user.id, password_data, mock_redis - ) + success = await change_password(test_db_session, test_user.id, password_data, mock_redis) assert success is True @@ -312,12 +306,8 @@ async def test_change_password_wrong_current_password( mock_redis = AsyncMock() - with pytest.raises( - AuthenticationException, match="Current password is incorrect" - ): - await change_password( - test_db_session, test_user.id, password_data, mock_redis - ) + with pytest.raises(AuthenticationException, match="Current password is incorrect"): + await change_password(test_db_session, test_user.id, password_data, mock_redis) @pytest.mark.asyncio async def test_change_password_user_not_found(self, test_db_session: AsyncSession): @@ -331,9 +321,7 @@ async def test_change_password_user_not_found(self, test_db_session: AsyncSessio mock_redis = AsyncMock() non_existent_id = "non-existent-id-123" - success = await change_password( - test_db_session, non_existent_id, password_data, mock_redis - ) + success = await change_password(test_db_session, non_existent_id, password_data, mock_redis) assert success is False @@ -351,13 +339,9 @@ async def test_change_password_database_error( mock_redis = AsyncMock() # Mock commit to raise an exception - with patch.object( - test_db_session, "commit", side_effect=SQLAlchemyError("Commit error") - ): + with patch.object(test_db_session, "commit", side_effect=SQLAlchemyError("Commit error")): with pytest.raises(ServerException, match="Failed to change password"): - await change_password( - test_db_session, test_user.id, password_data, mock_redis - ) + await change_password(test_db_session, test_user.id, password_data, mock_redis) @pytest.mark.asyncio async def test_change_password_authentication_exception_propagation( @@ -374,9 +358,7 @@ async def test_change_password_authentication_exception_propagation( # Should raise AuthenticationException, not ServerException with pytest.raises(AuthenticationException): - await change_password( - test_db_session, test_user.id, password_data, mock_redis - ) + await change_password(test_db_session, test_user.id, password_data, mock_redis) class TestServiceIntegration: @@ -406,9 +388,7 @@ async def test_update_profile_then_change_password( ) mock_redis = AsyncMock() - success = await change_password( - test_db_session, test_user.id, password_data, mock_redis - ) + success = await change_password(test_db_session, test_user.id, password_data, mock_redis) assert success is True @@ -492,9 +472,7 @@ async def test_change_password_with_none_redis_client( logout_other_devices=True, # This should be ignored if redis_client is None ) - success = await change_password( - test_db_session, test_user.id, password_data, None - ) + success = await change_password(test_db_session, test_user.id, password_data, None) assert success is True await test_db_session.refresh(test_user) @@ -517,9 +495,7 @@ async def test_service_functions_with_invalid_user_id_format( # Test update_user_profile update_data = UserUpdate(first_name="Test") - updated_user = await update_user_profile( - test_db_session, invalid_id, update_data - ) + updated_user = await update_user_profile(test_db_session, invalid_id, update_data) assert updated_user is None # Test change_password @@ -529,7 +505,5 @@ async def test_service_functions_with_invalid_user_id_format( logout_other_devices=False, ) mock_redis = AsyncMock() - success = await change_password( - test_db_session, invalid_id, password_data, mock_redis - ) + success = await change_password(test_db_session, invalid_id, password_data, mock_redis) assert success is False diff --git a/backend/tests/api/auth/test_controller.py b/backend/tests/api/auth/test_controller.py index 9982376..3359d2c 100644 --- a/backend/tests/api/auth/test_controller.py +++ b/backend/tests/api/auth/test_controller.py @@ -1,26 +1,28 @@ -import pytest from unittest.mock import AsyncMock, patch + +import pytest from httpx import AsyncClient + from api.auth.schema import ( ActionRequiredResponse, ) +from core.redis import get_redis +from core.security import ( + verify_email_verification_token, + verify_password_reset_token, + verify_token, +) +from main import app from utils.custom_exception import ( - ConflictException, AuthenticationException, - PasswordResetRequiredException, + ConflictException, + EmailVerificationRequiredException, NotFoundException, + PasswordResetRequiredException, + RegistrationDisabledException, SMTPNotConfiguredException, ValidationException, - EmailVerificationRequiredException, - RegistrationDisabledException, -) -from core.security import ( - verify_password_reset_token, - verify_token, - verify_email_verification_token, ) -from core.redis import get_redis -from main import app class TestAuthController: @@ -185,9 +187,7 @@ async def test_login_invalid_credentials(self, client: AsyncClient): login_data = {"email": "john.doe@example.com", "password": "WrongPassword"} with patch("api.auth.controller.login") as mock_login: - mock_login.side_effect = AuthenticationException( - "Invalid email or password" - ) + mock_login.side_effect = AuthenticationException("Invalid email or password") response = await client.post("/api/auth/login", json=login_data) @@ -259,9 +259,7 @@ async def mock_verify_token(): return {"sub": "test-user-id", "sid": None} with patch("api.auth.controller.logout") as mock_logout: - mock_logout.side_effect = AuthenticationException( - "Invalid or expired session" - ) + mock_logout.side_effect = AuthenticationException("Invalid or expired session") app.dependency_overrides[verify_token] = mock_verify_token try: @@ -294,9 +292,10 @@ async def mock_verify_token(): @pytest.mark.asyncio async def test_token_refresh_success(self, client: AsyncClient): """Test successful token refresh""" - with patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, patch( - "api.auth.controller.token" - ) as mock_token: + with ( + patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, + patch("api.auth.controller.token") as mock_token, + ): mock_verify_csrf.return_value = "test-session-id" mock_token.return_value = "new-access-token" @@ -329,9 +328,7 @@ async def test_token_refresh_no_csrf_header(self, client: AsyncClient): async def test_token_refresh_invalid_csrf(self, client: AsyncClient): """Test token refresh with invalid CSRF token""" with patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf: - mock_verify_csrf.side_effect = AuthenticationException( - "Invalid or expired CSRF token" - ) + mock_verify_csrf.side_effect = AuthenticationException("Invalid or expired CSRF token") response = await client.post( "/api/auth/token", @@ -347,9 +344,10 @@ async def test_token_refresh_invalid_csrf(self, client: AsyncClient): @pytest.mark.asyncio async def test_token_refresh_invalid_session(self, client: AsyncClient): """Test token refresh with invalid session""" - with patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, patch( - "api.auth.controller.token" - ) as mock_token: + with ( + patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, + patch("api.auth.controller.token") as mock_token, + ): mock_verify_csrf.return_value = "other-session-id" response = await client.post( @@ -367,9 +365,10 @@ async def test_token_refresh_invalid_session(self, client: AsyncClient): @pytest.mark.asyncio async def test_token_refresh_user_not_found(self, client: AsyncClient): """Test token refresh with user not found""" - with patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, patch( - "api.auth.controller.token" - ) as mock_token: + with ( + patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, + patch("api.auth.controller.token") as mock_token, + ): mock_verify_csrf.return_value = "test-session-id" mock_token.side_effect = NotFoundException("User not found") @@ -387,9 +386,10 @@ async def test_token_refresh_user_not_found(self, client: AsyncClient): @pytest.mark.asyncio async def test_token_refresh_server_error(self, client: AsyncClient): """Test token refresh with server error""" - with patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, patch( - "api.auth.controller.token" - ) as mock_token: + with ( + patch("api.auth.controller.verify_csrf_token") as mock_verify_csrf, + patch("api.auth.controller.token") as mock_token, + ): mock_verify_csrf.return_value = "test-session-id" mock_token.side_effect = Exception("Database error") @@ -503,9 +503,7 @@ async def mock_verify_password_reset_token(): "csrf_token": "test-csrf-token", } - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.post( @@ -546,9 +544,7 @@ async def mock_verify_password_reset_token(): with patch("api.auth.controller.reset_password") as mock_reset: mock_reset.side_effect = AuthenticationException("Invalid or expired token") - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.post( @@ -577,9 +573,7 @@ async def mock_verify_password_reset_token(): with patch("api.auth.controller.reset_password") as mock_reset: mock_reset.side_effect = Exception("Database error") - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.post( @@ -608,9 +602,7 @@ async def mock_verify_password_reset_token(): with patch("api.auth.controller.reset_password") as mock_reset: mock_reset.side_effect = NotFoundException("User not found") - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.post( @@ -636,15 +628,9 @@ async def mock_verify_password_reset_token(): "email": "test@example.com", } - with patch( - "api.auth.controller.validate_password_reset_token" - ) as mock_validate: - mock_validate.return_value = type( - "ValidationResult", (), {"is_valid": True} - )() - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + with patch("api.auth.controller.validate_password_reset_token") as mock_validate: + mock_validate.return_value = type("ValidationResult", (), {"is_valid": True})() + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.get( @@ -680,13 +666,9 @@ async def mock_verify_password_reset_token(): "email": "test@example.com", } - with patch( - "api.auth.controller.validate_password_reset_token" - ) as mock_validate: + with patch("api.auth.controller.validate_password_reset_token") as mock_validate: mock_validate.side_effect = Exception("Database error") - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.get( @@ -748,9 +730,7 @@ async def mock_verify_password_reset_token(): "csrf_token": "test-csrf-token", } - app.dependency_overrides[verify_password_reset_token] = ( - mock_verify_password_reset_token - ) + app.dependency_overrides[verify_password_reset_token] = mock_verify_password_reset_token try: response = await client.post( @@ -769,7 +749,9 @@ async def test_forgot_password_send_email_success(self, client: AsyncClient): """Test forgot password sends reset email""" req_data = {"email": "john.doe@example.com"} with patch("api.auth.controller.forgot_password", new_callable=AsyncMock) as mock_send: - mock_send.return_value = {"reset_url": "http://localhost:3000/reset-password?token=test"} + mock_send.return_value = { + "reset_url": "http://localhost:3000/reset-password?token=test" + } response = await client.post("/api/auth/forgot-password", json=req_data) assert response.status_code == 200 @@ -798,7 +780,7 @@ async def test_forgot_password_cooldown_active(self, client: AsyncClient): with patch("api.auth.controller.forgot_password", new_callable=AsyncMock) as mock_send: mock_send.side_effect = ValidationException( "Please wait 60 seconds before requesting another password reset email", - details={"cooldown_seconds": 60} + details={"cooldown_seconds": 60}, ) response = await client.post("/api/auth/forgot-password", json=req_data) @@ -836,7 +818,9 @@ async def test_forgot_password_smtp_disabled(self, client: AsyncClient): @pytest.mark.asyncio async def test_get_password_reset_cooldown_success(self, client: AsyncClient): """Test get password reset cooldown returns remaining time""" - with patch("api.auth.controller.get_password_reset_cooldown", new_callable=AsyncMock) as mock_cooldown: + with patch( + "api.auth.controller.get_password_reset_cooldown", new_callable=AsyncMock + ) as mock_cooldown: mock_cooldown.return_value = {"cooldown_seconds": 120} response = await client.get("/api/auth/forgot-password/cooldown?email=test@example.com") @@ -850,7 +834,9 @@ async def test_get_password_reset_cooldown_success(self, client: AsyncClient): @pytest.mark.asyncio async def test_get_password_reset_cooldown_no_cooldown(self, client: AsyncClient): """Test get password reset cooldown returns 0 when no cooldown""" - with patch("api.auth.controller.get_password_reset_cooldown", new_callable=AsyncMock) as mock_cooldown: + with patch( + "api.auth.controller.get_password_reset_cooldown", new_callable=AsyncMock + ) as mock_cooldown: mock_cooldown.return_value = {"cooldown_seconds": 0} response = await client.get("/api/auth/forgot-password/cooldown?email=test@example.com") @@ -862,7 +848,9 @@ async def test_get_password_reset_cooldown_no_cooldown(self, client: AsyncClient @pytest.mark.asyncio async def test_get_password_reset_cooldown_server_error(self, client: AsyncClient): """Test get password reset cooldown with server error""" - with patch("api.auth.controller.get_password_reset_cooldown", new_callable=AsyncMock) as mock_cooldown: + with patch( + "api.auth.controller.get_password_reset_cooldown", new_callable=AsyncMock + ) as mock_cooldown: mock_cooldown.side_effect = Exception("Redis error") response = await client.get("/api/auth/forgot-password/cooldown?email=test@example.com") assert response.status_code == 500 @@ -870,6 +858,7 @@ async def test_get_password_reset_cooldown_server_error(self, client: AsyncClien @pytest.mark.asyncio async def test_verify_email_success(self, client: AsyncClient): """Test verify email success""" + async def mock_verify_email_token(): return { "sub": "test-user-id", @@ -895,9 +884,7 @@ async def mock_verify_email_token(): "access_token": "test-access-token", "csrf_token": "test-csrf-token", } - app.dependency_overrides[verify_email_verification_token] = ( - mock_verify_email_token - ) + app.dependency_overrides[verify_email_verification_token] = mock_verify_email_token try: response = await client.get( @@ -925,6 +912,7 @@ async def test_verify_email_invalid_token(self, client: AsyncClient): @pytest.mark.asyncio async def test_verify_email_user_not_found(self, client: AsyncClient): """Test verify email when user not found""" + async def mock_verify_email_token(): return { "sub": "test-user-id", @@ -935,9 +923,7 @@ async def mock_verify_email_token(): with patch("api.auth.controller.verify_email") as mock_verify: mock_verify.side_effect = NotFoundException("User not found") - app.dependency_overrides[verify_email_verification_token] = ( - mock_verify_email_token - ) + app.dependency_overrides[verify_email_verification_token] = mock_verify_email_token try: response = await client.get( "/api/auth/verify-email", @@ -953,6 +939,7 @@ async def mock_verify_email_token(): @pytest.mark.asyncio async def test_verify_email_conflict(self, client: AsyncClient): """Test verify email when email already exists""" + async def mock_verify_email_token(): return { "sub": "test-user-id", @@ -963,9 +950,7 @@ async def mock_verify_email_token(): with patch("api.auth.controller.verify_email") as mock_verify: mock_verify.side_effect = ConflictException("Email already exists") - app.dependency_overrides[verify_email_verification_token] = ( - mock_verify_email_token - ) + app.dependency_overrides[verify_email_verification_token] = mock_verify_email_token try: response = await client.get( "/api/auth/verify-email", @@ -981,6 +966,7 @@ async def mock_verify_email_token(): @pytest.mark.asyncio async def test_verify_email_server_error(self, client: AsyncClient): """Test verify email with server error""" + async def mock_verify_email_token(): return { "sub": "test-user-id", @@ -991,9 +977,7 @@ async def mock_verify_email_token(): with patch("api.auth.controller.verify_email") as mock_verify: mock_verify.side_effect = Exception("Database error") - app.dependency_overrides[verify_email_verification_token] = ( - mock_verify_email_token - ) + app.dependency_overrides[verify_email_verification_token] = mock_verify_email_token try: response = await client.get( "/api/auth/verify-email", @@ -1010,7 +994,9 @@ async def mock_verify_email_token(): async def test_resend_verification_email_success(self, client: AsyncClient): """Test resend verification email success""" req_data = {"email": "john.doe@example.com"} - with patch("api.auth.controller.resend_verification_email", new_callable=AsyncMock) as mock_send: + with patch( + "api.auth.controller.resend_verification_email", new_callable=AsyncMock + ) as mock_send: mock_send.return_value = {"message": "Verification email sent"} response = await client.post("/api/auth/resend-verification", json=req_data) assert response.status_code == 200 @@ -1023,7 +1009,9 @@ async def test_resend_verification_email_success(self, client: AsyncClient): async def test_resend_verification_email_cooldown(self, client: AsyncClient): """Test resend verification email with cooldown""" req_data = {"email": "john.doe@example.com"} - with patch("api.auth.controller.resend_verification_email", new_callable=AsyncMock) as mock_send: + with patch( + "api.auth.controller.resend_verification_email", new_callable=AsyncMock + ) as mock_send: mock_send.side_effect = ValidationException("Please wait") response = await client.post("/api/auth/resend-verification", json=req_data) assert response.status_code == 400 @@ -1035,7 +1023,9 @@ async def test_resend_verification_email_cooldown(self, client: AsyncClient): async def test_resend_verification_email_disabled_account(self, client: AsyncClient): """Test resend verification email with disabled account""" req_data = {"email": "disabled@example.com"} - with patch("api.auth.controller.resend_verification_email", new_callable=AsyncMock) as mock_send: + with patch( + "api.auth.controller.resend_verification_email", new_callable=AsyncMock + ) as mock_send: mock_send.side_effect = AuthenticationException("Account is disabled") response = await client.post("/api/auth/resend-verification", json=req_data) assert response.status_code == 403 @@ -1047,7 +1037,9 @@ async def test_resend_verification_email_disabled_account(self, client: AsyncCli async def test_resend_verification_email_not_found(self, client: AsyncClient): """Test resend verification email when user not found""" req_data = {"email": "missing@example.com"} - with patch("api.auth.controller.resend_verification_email", new_callable=AsyncMock) as mock_send: + with patch( + "api.auth.controller.resend_verification_email", new_callable=AsyncMock + ) as mock_send: mock_send.side_effect = NotFoundException("User not registered") response = await client.post("/api/auth/resend-verification", json=req_data) assert response.status_code == 404 @@ -1059,7 +1051,9 @@ async def test_resend_verification_email_not_found(self, client: AsyncClient): async def test_resend_verification_email_smtp_disabled(self, client: AsyncClient): """Test resend verification email when SMTP disabled""" req_data = {"email": "john.doe@example.com"} - with patch("api.auth.controller.resend_verification_email", new_callable=AsyncMock) as mock_send: + with patch( + "api.auth.controller.resend_verification_email", new_callable=AsyncMock + ) as mock_send: mock_send.side_effect = SMTPNotConfiguredException("SMTP is disabled") response = await client.post("/api/auth/resend-verification", json=req_data) assert response.status_code == 503 @@ -1071,7 +1065,9 @@ async def test_resend_verification_email_smtp_disabled(self, client: AsyncClient async def test_resend_verification_email_server_error(self, client: AsyncClient): """Test resend verification email with server error""" req_data = {"email": "john.doe@example.com"} - with patch("api.auth.controller.resend_verification_email", new_callable=AsyncMock) as mock_send: + with patch( + "api.auth.controller.resend_verification_email", new_callable=AsyncMock + ) as mock_send: mock_send.side_effect = Exception("Database error") response = await client.post("/api/auth/resend-verification", json=req_data) assert response.status_code == 500 @@ -1081,8 +1077,10 @@ async def test_get_email_verification_cooldown(self, client: AsyncClient): """Test get email verification cooldown returns remaining time""" mock_redis = AsyncMock() mock_redis.ttl.return_value = 120 + async def override_get_redis(): return mock_redis + app.dependency_overrides[get_redis] = override_get_redis try: @@ -1102,8 +1100,10 @@ async def test_get_email_verification_cooldown_server_error(self, client: AsyncC """Test get email verification cooldown with server error""" mock_redis = AsyncMock() mock_redis.ttl.side_effect = Exception("Redis error") + async def override_get_redis(): return mock_redis + app.dependency_overrides[get_redis] = override_get_redis try: @@ -1112,4 +1112,4 @@ async def override_get_redis(): ) assert response.status_code == 500 finally: - app.dependency_overrides.pop(get_redis, None) \ No newline at end of file + app.dependency_overrides.pop(get_redis, None) diff --git a/backend/tests/api/auth/test_schema.py b/backend/tests/api/auth/test_schema.py index 1f656c8..023da12 100644 --- a/backend/tests/api/auth/test_schema.py +++ b/backend/tests/api/auth/test_schema.py @@ -1,20 +1,22 @@ -import pytest from datetime import datetime + +import pytest from pydantic import ValidationError + from api.auth.schema import ( - UserRegister, - UserLogin, - UserResponse, - UserLoginResponse, - TokenResponse, - PasswordResetRequiredResponse, - LogoutRequest, - ResetPasswordRequest, - TokenValidationResponse, ForgotPasswordRequest, - PasswordResetCooldownResponse, LoginResult, + LogoutRequest, + PasswordResetCooldownResponse, + PasswordResetRequiredResponse, + ResetPasswordRequest, SessionResult, + TokenResponse, + TokenValidationResponse, + UserLogin, + UserLoginResponse, + UserRegister, + UserResponse, ) @@ -467,4 +469,4 @@ def test_datetime_validation(self): } token_response = TokenResponse(**data) - assert token_response.expires_at is not None \ No newline at end of file + assert token_response.expires_at is not None diff --git a/backend/tests/api/auth/test_service.py b/backend/tests/api/auth/test_service.py index 7a68f0e..f998ee1 100644 --- a/backend/tests/api/auth/test_service.py +++ b/backend/tests/api/auth/test_service.py @@ -1,48 +1,50 @@ -import pytest -from unittest.mock import AsyncMock, MagicMock, patch from datetime import datetime, timedelta +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession -from models.users import Users -from models.user_sessions import UserSessions -from models.password_reset_tokens import PasswordResetTokens -from models.email_verification_tokens import EmailVerificationTokens -from core.security import hash_password, create_access_token, create_csrf_token -from extensions.smtp import SMTPMailer + +from api.auth.schema import UserLogin, UserRegister from api.auth.services import ( - register, + _create_user, + _create_user_session, + _request_email_change_verification_email, + _request_registration_verification_email, + _send_registration_verification_email, + _update_session_expiry, + forgot_password, + get_or_create_csrf_token, + get_password_reset_cooldown, login, logout, logout_all_devices, - token, + register, + resend_verification_email, reset_password, + token, validate_password_reset_token, - forgot_password, - get_password_reset_cooldown, - verify_email, - resend_verification_email, - _send_registration_verification_email, - _request_registration_verification_email, - _request_email_change_verification_email, - _create_user, - _create_user_session, - _update_session_expiry, - get_or_create_csrf_token, verify_csrf_token, + verify_email, ) -from api.auth.schema import UserRegister, UserLogin +from core.config import settings +from core.security import create_access_token, create_csrf_token, hash_password +from extensions.smtp import SMTPMailer +from models.email_verification_tokens import EmailVerificationTokens +from models.password_reset_tokens import PasswordResetTokens +from models.user_sessions import UserSessions +from models.users import Users from utils.custom_exception import ( - ConflictException, AuthenticationException, + ConflictException, + EmailVerificationRequiredException, + NotFoundException, PasswordResetRequiredException, + RegistrationDisabledException, ServerException, - NotFoundException, SMTPNotConfiguredException, ValidationException, - EmailVerificationRequiredException, - RegistrationDisabledException, ) -from core.config import settings class TestAuthService: @@ -89,9 +91,7 @@ async def test_register_email_already_exists( with patch.object(settings, "REGISTRATION_ENABLE", True): with pytest.raises(ConflictException) as exc_info: - await register( - test_db_session, mock_redis, user_data, "127.0.0.1", "TestAgent/1.0" - ) + await register(test_db_session, mock_redis, user_data, "127.0.0.1", "TestAgent/1.0") assert "Email already exists" in str(exc_info.value) @@ -109,9 +109,7 @@ async def test_register_disabled(self, test_db_session: AsyncSession): with patch.object(settings, "REGISTRATION_ENABLE", False): with pytest.raises(RegistrationDisabledException) as exc_info: - await register( - test_db_session, mock_redis, user_data, "127.0.0.1", "TestAgent/1.0" - ) + await register(test_db_session, mock_redis, user_data, "127.0.0.1", "TestAgent/1.0") assert "Registration is disabled" in str(exc_info.value) @@ -121,9 +119,7 @@ async def test_login_success(self, test_db_session: AsyncSession, test_user: Use mock_redis = AsyncMock() login_data = UserLogin(email=test_user.email, password="TestPassword123!") - result = await login( - test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0" - ) + result = await login(test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0") assert "user" in result assert "session_id" in result @@ -134,14 +130,10 @@ async def test_login_success(self, test_db_session: AsyncSession, test_user: Use async def test_login_user_not_found(self, test_db_session: AsyncSession): """Test login with non-existent user""" mock_redis = AsyncMock() - login_data = UserLogin( - email="nonexistent@example.com", password="TestPassword123!" - ) + login_data = UserLogin(email="nonexistent@example.com", password="TestPassword123!") with pytest.raises(AuthenticationException) as exc_info: - await login( - test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0" - ) + await login(test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0") assert "Invalid email or password" in str(exc_info.value) @@ -169,24 +161,18 @@ async def test_login_disabled_account(self, test_db_session: AsyncSession): login_data = UserLogin(email=disabled_user.email, password="TestPassword123!") with pytest.raises(AuthenticationException) as exc_info: - await login( - test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0" - ) + await login(test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0") assert "Account is disabled" in str(exc_info.value) @pytest.mark.asyncio - async def test_login_invalid_password( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_login_invalid_password(self, test_db_session: AsyncSession, test_user: Users): """Test login with invalid password""" mock_redis = AsyncMock() login_data = UserLogin(email=test_user.email, password="WrongPassword123!") with pytest.raises(AuthenticationException) as exc_info: - await login( - test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0" - ) + await login(test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0") assert "Invalid email or password" in str(exc_info.value) @@ -214,9 +200,7 @@ async def test_login_password_reset_required(self, test_db_session: AsyncSession login_data = UserLogin(email=reset_user.email, password="TestPassword123!") with pytest.raises(PasswordResetRequiredException) as exc_info: - await login( - test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0" - ) + await login(test_db_session, mock_redis, login_data, "127.0.0.1", "TestAgent/1.0") assert "Password reset required" in str(exc_info.value) assert exc_info.value.details is not None @@ -235,9 +219,7 @@ async def test_logout_success( mock_redis = AsyncMock() mock_redis.delete.return_value = 1 - result = await logout( - test_db_session, mock_redis, test_user.id, test_user_session.id - ) + result = await logout(test_db_session, mock_redis, test_user.id, test_user_session.id) assert result is True mock_redis.delete.assert_called_once_with( @@ -252,15 +234,11 @@ async def test_logout_all_devices_success( """Test successful logout all devices""" mock_redis = AsyncMock() - with patch( - "api.auth.services.clear_user_all_sessions", return_value=True - ) as mock_clear: + with patch("api.auth.services.clear_user_all_sessions", return_value=True) as mock_clear: result = await logout_all_devices(test_db_session, mock_redis, test_user.id) assert result is True - mock_clear.assert_called_once_with( - test_db_session, mock_redis, test_user.id - ) + mock_clear.assert_called_once_with(test_db_session, mock_redis, test_user.id) @pytest.mark.asyncio async def test_logout_all_devices_failure( @@ -299,9 +277,7 @@ async def test_token_refresh_success( mock_redis.get.return_value = str(session_data) mock_redis.setex.return_value = True - with patch( - "api.auth.services.extend_session_ttl", return_value=True - ) as mock_extend: + with patch("api.auth.services.extend_session_ttl", return_value=True) as mock_extend: result = await token(test_db_session, mock_redis, test_user_session.id) assert result is not None @@ -406,9 +382,7 @@ async def test_reset_password_success(self, test_db_session: AsyncSession): token_data = {"sub": user_id, "token": "test_reset_token_123"} - with patch( - "api.auth.services.clear_user_all_sessions", return_value=True - ) as mock_clear: + with patch("api.auth.services.clear_user_all_sessions", return_value=True) as mock_clear: result = await reset_password( test_db_session, mock_redis, @@ -466,9 +440,7 @@ async def test_reset_password_user_not_found( assert "User not found" in str(exc_info.value) @pytest.mark.asyncio - async def test_validate_password_reset_token_success( - self, test_db_session: AsyncSession - ): + async def test_validate_password_reset_token_success(self, test_db_session: AsyncSession): """Test successful password reset token validation""" user_id = "test-validate-token-service-user" hashed_pwd = await hash_password("TestPassword123!") @@ -501,9 +473,7 @@ async def test_validate_password_reset_token_success( assert result.is_valid is True @pytest.mark.asyncio - async def test_validate_password_reset_token_invalid( - self, test_db_session: AsyncSession - ): + async def test_validate_password_reset_token_invalid(self, test_db_session: AsyncSession): """Test validation of invalid password reset token""" token_data = {"sub": "nonexistent_user", "token": "invalid_token"} @@ -527,13 +497,11 @@ async def test_validate_password_reset_token_user_not_found( assert "User not found or account disabled" in str(exc_info.value) @pytest.mark.asyncio - async def test_validate_password_reset_token_not_required( - self, test_db_session: AsyncSession - ): + async def test_validate_password_reset_token_not_required(self, test_db_session: AsyncSession): """Test password reset token validation when password reset not required""" user_id = "test-no-reset-validate-user" hashed_pwd = await hash_password("TestPassword123!") - + no_reset_user = Users( id=user_id, email="noresetvalidate@example.com", @@ -563,14 +531,12 @@ async def test_validate_password_reset_token_not_required( assert result.is_valid is True @pytest.mark.asyncio - async def test_forgot_password_success( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_forgot_password_success(self, test_db_session: AsyncSession, test_user: Users): """Test successful forgot password flow""" mock_redis = AsyncMock() mock_redis.ttl.return_value = -1 # No cooldown mock_redis.setex.return_value = True - + mock_mailer = MagicMock(spec=SMTPMailer) mock_mailer.enabled = True mock_mailer.send_text = MagicMock() @@ -604,9 +570,7 @@ async def test_forgot_password_user_not_found(self, test_db_session: AsyncSessio ) @pytest.mark.asyncio - async def test_forgot_password_account_disabled( - self, test_db_session: AsyncSession - ): + async def test_forgot_password_account_disabled(self, test_db_session: AsyncSession): """Test forgot password with disabled account""" hashed_pwd = await hash_password("TestPassword123!") disabled_user = Users( @@ -624,7 +588,7 @@ async def test_forgot_password_account_disabled( mock_redis = AsyncMock() mock_redis.ttl.return_value = -1 - + mock_mailer = MagicMock(spec=SMTPMailer) mock_mailer.enabled = True @@ -644,7 +608,7 @@ async def test_forgot_password_cooldown_active( """Test forgot password when cooldown is active""" mock_redis = AsyncMock() mock_redis.ttl.return_value = 60 # 60 seconds remaining - + mock_mailer = MagicMock(spec=SMTPMailer) mock_mailer.enabled = True @@ -665,7 +629,7 @@ async def test_forgot_password_smtp_disabled( """Test forgot password when SMTP is disabled""" mock_redis = AsyncMock() mock_redis.ttl.return_value = -1 - + mock_mailer = MagicMock(spec=SMTPMailer) mock_mailer.enabled = False @@ -725,9 +689,11 @@ async def test_register_email_verification_cooldown(self, test_db_session: Async password="TestPassword123!", ) - with patch.object(settings, "REGISTRATION_ENABLE", True), patch.object( - settings, "EMAIL_VERIFICATION_ENABLE", True - ), patch.object(settings, "SMTP_ENABLE", True): + with ( + patch.object(settings, "REGISTRATION_ENABLE", True), + patch.object(settings, "EMAIL_VERIFICATION_ENABLE", True), + patch.object(settings, "SMTP_ENABLE", True), + ): with pytest.raises(EmailVerificationRequiredException): await register( test_db_session, @@ -756,9 +722,11 @@ async def test_register_email_verification_send(self, test_db_session: AsyncSess password="TestPassword123!", ) - with patch.object(settings, "REGISTRATION_ENABLE", True), patch.object( - settings, "EMAIL_VERIFICATION_ENABLE", True - ), patch.object(settings, "SMTP_ENABLE", True): + with ( + patch.object(settings, "REGISTRATION_ENABLE", True), + patch.object(settings, "EMAIL_VERIFICATION_ENABLE", True), + patch.object(settings, "SMTP_ENABLE", True), + ): with pytest.raises(EmailVerificationRequiredException): await register( test_db_session, @@ -796,8 +764,9 @@ async def test_login_email_verification_cooldown(self, test_db_session: AsyncSes login_data = UserLogin(email=user.email, password="TestPassword123!") - with patch.object(settings, "EMAIL_VERIFICATION_ENABLE", True), patch.object( - settings, "SMTP_ENABLE", True + with ( + patch.object(settings, "EMAIL_VERIFICATION_ENABLE", True), + patch.object(settings, "SMTP_ENABLE", True), ): with pytest.raises(EmailVerificationRequiredException) as exc_info: await login( @@ -837,8 +806,9 @@ async def test_login_email_verification_send(self, test_db_session: AsyncSession login_data = UserLogin(email=user.email, password="TestPassword123!") - with patch.object(settings, "EMAIL_VERIFICATION_ENABLE", True), patch.object( - settings, "SMTP_ENABLE", True + with ( + patch.object(settings, "EMAIL_VERIFICATION_ENABLE", True), + patch.object(settings, "SMTP_ENABLE", True), ): with pytest.raises(EmailVerificationRequiredException): await login( @@ -853,7 +823,9 @@ async def test_login_email_verification_send(self, test_db_session: AsyncSession mock_mailer.send_text.assert_called() @pytest.mark.asyncio - async def test_logout_server_error(self, test_db_session: AsyncSession, test_user: Users, test_user_session: UserSessions): + async def test_logout_server_error( + self, test_db_session: AsyncSession, test_user: Users, test_user_session: UserSessions + ): """Test logout server error""" mock_redis = AsyncMock() mock_redis.delete.side_effect = Exception("Redis error") @@ -947,7 +919,9 @@ async def test_reset_password_server_exception( token_data = {"sub": test_user.id, "token": "test_reset_token"} - with patch("api.auth.services.clear_user_all_sessions", side_effect=Exception("Redis error")): + with patch( + "api.auth.services.clear_user_all_sessions", side_effect=Exception("Redis error") + ): with pytest.raises(ServerException): await reset_password( test_db_session, @@ -959,14 +933,18 @@ async def test_reset_password_server_exception( ) @pytest.mark.asyncio - async def test_validate_password_reset_token_server_exception(self, test_db_session: AsyncSession): + async def test_validate_password_reset_token_server_exception( + self, test_db_session: AsyncSession + ): """Test validate password reset token server exception""" with patch.object(test_db_session, "execute", side_effect=Exception("DB error")): with pytest.raises(ServerException): await validate_password_reset_token(test_db_session, {"sub": "x", "token": "y"}) @pytest.mark.asyncio - async def test_forgot_password_server_exception(self, test_db_session: AsyncSession, test_user: Users): + async def test_forgot_password_server_exception( + self, test_db_session: AsyncSession, test_user: Users + ): """Test forgot password server exception""" mock_redis = AsyncMock() mock_redis.ttl.return_value = -1 @@ -982,7 +960,9 @@ class TestAuthEmailVerificationHelpers: """Test email verification helper functions""" @pytest.mark.asyncio - async def test_send_registration_verification_email(self, test_db_session: AsyncSession, test_user: Users): + async def test_send_registration_verification_email( + self, test_db_session: AsyncSession, test_user: Users + ): """Test sending registration verification email""" mock_mailer = MagicMock(spec=SMTPMailer) mock_mailer.enabled = True @@ -994,7 +974,9 @@ async def test_send_registration_verification_email(self, test_db_session: Async mock_mailer.send_text.assert_called_once() @pytest.mark.asyncio - async def test_request_registration_verification_email(self, test_db_session: AsyncSession, test_user: Users): + async def test_request_registration_verification_email( + self, test_db_session: AsyncSession, test_user: Users + ): """Test creating registration verification token record""" token_meta = await _request_registration_verification_email(test_db_session, test_user) @@ -1005,10 +987,14 @@ async def test_request_registration_verification_email(self, test_db_session: As assert result.scalar_one_or_none() is not None @pytest.mark.asyncio - async def test_request_email_change_verification_email(self, test_db_session: AsyncSession, test_user: Users): + async def test_request_email_change_verification_email( + self, test_db_session: AsyncSession, test_user: Users + ): """Test creating email change verification token record""" new_email = "change@example.com" - token_meta = await _request_email_change_verification_email(test_db_session, test_user, new_email) + token_meta = await _request_email_change_verification_email( + test_db_session, test_user, new_email + ) assert token_meta["verification_token"] result = await test_db_session.execute( @@ -1267,7 +1253,9 @@ async def test_create_user_server_exception(self, test_db_session: AsyncSession) await _create_user(test_db_session, user_data) @pytest.mark.asyncio - async def test_create_user_session_server_exception(self, test_db_session: AsyncSession, test_user: Users): + async def test_create_user_session_server_exception( + self, test_db_session: AsyncSession, test_user: Users + ): """Test _create_user_session server exception""" mock_redis = AsyncMock() mock_redis.setex.side_effect = Exception("Redis error") @@ -1448,7 +1436,5 @@ async def test_send_registration_verification_email_smtp_disabled( mock_mailer = MagicMock(spec=SMTPMailer) mock_mailer.enabled = False with patch.object(settings, "SMTP_ENABLE", False): - await _send_registration_verification_email( - test_db_session, mock_mailer, test_user - ) - mock_mailer.send_text.assert_not_called() \ No newline at end of file + await _send_registration_verification_email(test_db_session, mock_mailer, test_user) + mock_mailer.send_text.assert_not_called() diff --git a/backend/tests/api/roles/test_controller.py b/backend/tests/api/roles/test_controller.py index cae4770..4f46929 100644 --- a/backend/tests/api/roles/test_controller.py +++ b/backend/tests/api/roles/test_controller.py @@ -1,19 +1,21 @@ -import pytest from unittest.mock import patch + +import pytest from httpx import AsyncClient + from api.roles.schema import ( - RoleResponse, - RolesListResponse, - RoleAttributesGroupedResponse, - RoleAttributeMappingBatchResponse, AttributeMappingResult, PermissionCheckResponse, + RoleAttributeMappingBatchResponse, + RoleAttributesGroupedResponse, + RoleResponse, + RolesListResponse, ) from utils.custom_exception import ( + AuthorizationException, ConflictException, NotFoundException, ServerException, - AuthorizationException, ) @@ -21,9 +23,7 @@ class TestGetRolesAPI: """Test GET /api/roles endpoint""" @pytest.mark.asyncio - async def test_get_roles_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_get_roles_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful roles retrieval with valid authentication""" with patch("api.roles.controller.get_all_roles") as mock_get_roles: mock_roles = RolesListResponse( @@ -55,9 +55,7 @@ async def test_get_roles_unauthorized(self, client: AsyncClient): assert response.status_code == 401 @pytest.mark.asyncio - async def test_get_roles_server_error( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_get_roles_server_error(self, client: AsyncClient, users_auth_headers: dict): """Test roles retrieval with server error""" with patch( "api.roles.controller.get_all_roles", @@ -74,9 +72,7 @@ class TestCreateRoleAPI: """Test POST /api/roles endpoint""" @pytest.mark.asyncio - async def test_create_role_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_role_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful role creation""" role_data = { "name": "manager", @@ -106,9 +102,7 @@ async def test_create_role_success( assert data["data"]["name"] == "manager" @pytest.mark.asyncio - async def test_create_role_conflict( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_role_conflict(self, client: AsyncClient, users_auth_headers: dict): """Test role creation with name conflict""" role_data = {"name": "admin", "description": "Admin role", "level": 10} @@ -157,15 +151,11 @@ async def test_create_role_unauthorized(self, client: AsyncClient): assert response.status_code == 401 @pytest.mark.asyncio - async def test_create_role_server_error( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_role_server_error(self, client: AsyncClient, users_auth_headers: dict): """Test role creation with server error""" role_data = {"name": "test-role", "level": 10} - with patch( - "api.roles.controller.create_role", side_effect=Exception("Database error") - ): + with patch("api.roles.controller.create_role", side_effect=Exception("Database error")): response = await client.post( "/api/roles", json=role_data, @@ -178,16 +168,16 @@ class TestUpdateRoleAPI: """Test PUT /api/roles/{role_id} endpoint""" @pytest.mark.asyncio - async def test_update_role_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_role_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful role update""" role_id = "role-123" role_data = {"name": "updated_manager", "description": "Updated manager role", "level": 10} with patch("api.roles.controller.update_role") as mock_update_role: mock_role = RoleResponse( - id=role_id, name="updated_manager", description="Updated manager role", + id=role_id, + name="updated_manager", + description="Updated manager role", level=10, ) mock_update_role.return_value = mock_role @@ -205,9 +195,7 @@ async def test_update_role_success( assert data["data"]["name"] == "updated_manager" @pytest.mark.asyncio - async def test_update_role_not_found( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_role_not_found(self, client: AsyncClient, users_auth_headers: dict): """Test role update with non-existent role""" role_id = "non-existent-role" role_data = {"name": "updated_role", "level": 10} @@ -251,9 +239,7 @@ async def test_update_role_forbidden_super_admin( assert data["message"] == "Cannot modify the system super-admin role" @pytest.mark.asyncio - async def test_update_role_conflict( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_role_conflict(self, client: AsyncClient, users_auth_headers: dict): """Test role update with name conflict""" role_id = "role-123" role_data = {"name": "existing-role", "level": 10} @@ -281,9 +267,7 @@ async def test_update_role_unauthorized(self, client: AsyncClient): assert response.status_code == 401 @pytest.mark.asyncio - async def test_update_role_server_error( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_role_server_error(self, client: AsyncClient, users_auth_headers: dict): """Test role update with server error""" role_id = "role-123" role_data = {"name": "updated_role", "level": 10} @@ -304,9 +288,7 @@ class TestDeleteRoleAPI: """Test DELETE /api/roles/{role_id} endpoint""" @pytest.mark.asyncio - async def test_delete_role_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_delete_role_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful role deletion""" role_id = "role-123" @@ -324,9 +306,7 @@ async def test_delete_role_success( assert data["message"] == "Role deleted successfully" @pytest.mark.asyncio - async def test_delete_role_not_found( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_delete_role_not_found(self, client: AsyncClient, users_auth_headers: dict): """Test role deletion with non-existent role""" role_id = "non-existent-role" @@ -366,9 +346,7 @@ async def test_delete_role_forbidden_super_admin( assert data["message"] == "Cannot modify the system super-admin role" @pytest.mark.asyncio - async def test_delete_role_conflict( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_delete_role_conflict(self, client: AsyncClient, users_auth_headers: dict): """Test role deletion when role is assigned to users""" role_id = "role-123" @@ -395,15 +373,11 @@ async def test_delete_role_unauthorized(self, client: AsyncClient): assert response.status_code == 401 @pytest.mark.asyncio - async def test_delete_role_server_error( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_delete_role_server_error(self, client: AsyncClient, users_auth_headers: dict): """Test role deletion with server error""" role_id = "role-123" - with patch( - "api.roles.controller.delete_role", side_effect=Exception("Database error") - ): + with patch("api.roles.controller.delete_role", side_effect=Exception("Database error")): response = await client.delete( f"/api/roles/{role_id}", headers={"Authorization": users_auth_headers["Authorization"]}, @@ -421,9 +395,7 @@ async def test_get_role_attribute_mapping_success( """Test successful role attributes mapping retrieval""" role_id = "role-123" - with patch( - "api.roles.controller.get_role_attribute_mapping" - ) as mock_get_mapping: + with patch("api.roles.controller.get_role_attribute_mapping") as mock_get_mapping: mock_mapping = RoleAttributesGroupedResponse( groups=[ { @@ -453,7 +425,12 @@ async def test_get_role_attribute_mapping_success( assert data["message"] == "Successfully retrieved role attributes mapping" assert "data" in data assert "groups" in data["data"] - groups = {g["group"]: {cat: {a["name"]: a for a in attrs} for cat, attrs in g["categories"].items()} for g in data["data"]["groups"]} + groups = { + g["group"]: { + cat: {a["name"]: a for a in attrs} for cat, attrs in g["categories"].items() + } + for g in data["data"]["groups"] + } assert groups["user-role-management"]["user"]["attr-1"]["value"] is True assert groups["user-role-management"]["user"]["attr-2"]["value"] is False assert groups["user-role-management"]["role"]["attr-3"]["value"] is True @@ -465,9 +442,7 @@ async def test_get_role_attribute_mapping_not_found( """Test role attributes mapping retrieval with non-existent role""" role_id = "non-existent-role" - with patch( - "api.roles.controller.get_role_attribute_mapping" - ) as mock_get_mapping: + with patch("api.roles.controller.get_role_attribute_mapping") as mock_get_mapping: mock_get_mapping.side_effect = NotFoundException("Role not found") response = await client.get( @@ -487,9 +462,7 @@ async def test_get_role_attribute_mapping_forbidden_super_admin( """Test reading attributes of system super-admin role returns 403""" role_id = "role-super" - with patch( - "api.roles.controller.get_role_attribute_mapping" - ) as mock_get_mapping: + with patch("api.roles.controller.get_role_attribute_mapping") as mock_get_mapping: mock_get_mapping.side_effect = AuthorizationException( "Cannot modify the system super-admin role" ) @@ -540,9 +513,7 @@ async def test_update_role_attribute_mapping_success( role_id = "role-123" attributes_data = {"attributes": {"attr-1": True, "attr-2": False}} - with patch( - "api.roles.controller.update_role_attribute_mapping" - ) as mock_update_mapping: + with patch("api.roles.controller.update_role_attribute_mapping") as mock_update_mapping: mock_batch_result = RoleAttributeMappingBatchResponse( results=[ AttributeMappingResult( @@ -583,9 +554,7 @@ async def test_update_role_attribute_mapping_partial_success( role_id = "role-123" attributes_data = {"attributes": {"attr-1": True, "invalid-attr": False}} - with patch( - "api.roles.controller.update_role_attribute_mapping" - ) as mock_update_mapping: + with patch("api.roles.controller.update_role_attribute_mapping") as mock_update_mapping: mock_batch_result = RoleAttributeMappingBatchResponse( results=[ AttributeMappingResult( @@ -624,13 +593,9 @@ async def test_update_role_attribute_mapping_all_failed( ): """Test role attributes mapping update with all failures""" role_id = "role-123" - attributes_data = { - "attributes": {"invalid-attr-1": True, "invalid-attr-2": False} - } + attributes_data = {"attributes": {"invalid-attr-1": True, "invalid-attr-2": False}} - with patch( - "api.roles.controller.update_role_attribute_mapping" - ) as mock_update_mapping: + with patch("api.roles.controller.update_role_attribute_mapping") as mock_update_mapping: mock_batch_result = RoleAttributeMappingBatchResponse( results=[ AttributeMappingResult( @@ -671,9 +636,7 @@ async def test_update_role_attribute_mapping_not_found( role_id = "non-existent-role" attributes_data = {"attributes": {"attr-1": True}} - with patch( - "api.roles.controller.update_role_attribute_mapping" - ) as mock_update_mapping: + with patch("api.roles.controller.update_role_attribute_mapping") as mock_update_mapping: mock_update_mapping.side_effect = NotFoundException("Role not found") response = await client.put( @@ -695,9 +658,7 @@ async def test_update_role_attribute_mapping_forbidden_super_admin( role_id = "role-super" attributes_data = {"attributes": {"view-users": True}} - with patch( - "api.roles.controller.update_role_attribute_mapping" - ) as mock_update_mapping: + with patch("api.roles.controller.update_role_attribute_mapping") as mock_update_mapping: mock_update_mapping.side_effect = AuthorizationException( "Cannot modify the system super-admin role" ) @@ -714,15 +675,11 @@ async def test_update_role_attribute_mapping_forbidden_super_admin( assert data["message"] == "Cannot modify the system super-admin role" @pytest.mark.asyncio - async def test_update_role_attribute_mapping_unauthorized( - self, client: AsyncClient - ): + async def test_update_role_attribute_mapping_unauthorized(self, client: AsyncClient): """Test role attributes mapping update without authentication""" role_id = "role-123" attributes_data = {"attributes": {"attr-1": True}} - response = await client.put( - f"/api/roles/{role_id}/attributes", json=attributes_data - ) + response = await client.put(f"/api/roles/{role_id}/attributes", json=attributes_data) assert response.status_code == 401 @pytest.mark.asyncio @@ -753,9 +710,7 @@ async def test_get_user_permissions_success( self, client: AsyncClient, users_auth_headers: dict ): """Test successful user permissions retrieval""" - with patch( - "api.roles.controller.check_user_permissions" - ) as mock_check_permissions: + with patch("api.roles.controller.check_user_permissions") as mock_check_permissions: mock_result = PermissionCheckResponse( permissions={ "view-users": True, @@ -797,4 +752,4 @@ async def test_get_user_permissions_server_error( "/api/roles/permissions", headers={"Authorization": users_auth_headers["Authorization"]}, ) - assert response.status_code == 500 \ No newline at end of file + assert response.status_code == 500 diff --git a/backend/tests/api/roles/test_schema.py b/backend/tests/api/roles/test_schema.py index 8edcdca..48ec397 100644 --- a/backend/tests/api/roles/test_schema.py +++ b/backend/tests/api/roles/test_schema.py @@ -1,16 +1,17 @@ import pytest from pydantic import ValidationError + from api.roles.schema import ( - RoleResponse, - RolesListResponse, - RoleCreate, - RoleUpdate, - RoleAttributesMapping, - RoleAttributesGroupedResponse, AttributeMappingResult, - RoleAttributeMappingBatchResponse, PermissionCheckRequest, PermissionCheckResponse, + RoleAttributeMappingBatchResponse, + RoleAttributesGroupedResponse, + RoleAttributesMapping, + RoleCreate, + RoleResponse, + RolesListResponse, + RoleUpdate, ) @@ -422,4 +423,4 @@ def test_permission_check_response_missing_permissions(self): errors = exc_info.value.errors() assert len(errors) == 1 assert errors[0]["type"] == "missing" - assert errors[0]["loc"] == ("permissions",) \ No newline at end of file + assert errors[0]["loc"] == ("permissions",) diff --git a/backend/tests/api/roles/test_service.py b/backend/tests/api/roles/test_service.py index a921602..9f43ece 100644 --- a/backend/tests/api/roles/test_service.py +++ b/backend/tests/api/roles/test_service.py @@ -1,28 +1,30 @@ -import pytest from datetime import datetime from unittest.mock import patch -from sqlalchemy.ext.asyncio import AsyncSession + +import pytest from sqlalchemy import select -from models.roles import Roles +from sqlalchemy.ext.asyncio import AsyncSession + +from api.roles.schema import RoleCreate, RoleUpdate +from api.roles.services import ( + check_user_permissions, + create_role, + delete_role, + get_all_roles, + get_role_attribute_mapping, + update_role, + update_role_attribute_mapping, +) from models.role_attributes import RoleAttributes from models.role_attributes_mapper import RoleAttributesMapper from models.role_mapper import RoleMapper +from models.roles import Roles from models.users import Users from utils.custom_exception import ( + AuthorizationException, ConflictException, NotFoundException, ServerException, - AuthorizationException, -) -from api.roles.schema import RoleCreate, RoleUpdate -from api.roles.services import ( - get_all_roles, - create_role, - update_role, - delete_role, - get_role_attribute_mapping, - update_role_attribute_mapping, - check_user_permissions, ) @@ -92,9 +94,7 @@ async def test_get_all_roles_empty(self, test_db_session: AsyncSession): @pytest.mark.asyncio async def test_get_all_roles_database_error(self, test_db_session: AsyncSession): """Test get_all_roles with database error""" - with patch.object( - test_db_session, "execute", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "execute", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: await get_all_roles(test_db_session, actor_user_id="actor-1") @@ -108,7 +108,8 @@ class TestCreateRole: async def test_create_role_success(self, test_db_session: AsyncSession): """Test successful role creation""" role_data = RoleCreate( - name="manager", description="Manager role with special permissions", + name="manager", + description="Manager role with special permissions", level=10, ) @@ -122,9 +123,7 @@ async def test_create_role_success(self, test_db_session: AsyncSession): assert result.id is not None @pytest.mark.asyncio - async def test_create_role_rejects_level_too_high( - self, test_db_session: AsyncSession - ): + async def test_create_role_rejects_level_too_high(self, test_db_session: AsyncSession): """Test role creation rejects level > actor level""" actor = await _create_actor(test_db_session, level=20, user_id="actor-low") role_data = RoleCreate(name="too-high", description="blocked", level=21) @@ -135,9 +134,7 @@ async def test_create_role_rejects_level_too_high( assert "higher than your own" in str(exc_info.value) @pytest.mark.asyncio - async def test_create_role_allows_same_level( - self, test_db_session: AsyncSession - ): + async def test_create_role_allows_same_level(self, test_db_session: AsyncSession): """Test role creation allows level equal to actor level""" actor = await _create_actor(test_db_session, level=20, user_id="actor-same") role_data = RoleCreate(name="peer-role", description="same level", level=20) @@ -150,7 +147,8 @@ async def test_create_role_allows_same_level( @pytest.mark.asyncio async def test_create_role_minimal_data(self, test_db_session: AsyncSession): """Test role creation with minimal data""" - role_data = RoleCreate(name="guest", + role_data = RoleCreate( + name="guest", level=10, ) @@ -166,13 +164,13 @@ async def test_create_role_minimal_data(self, test_db_session: AsyncSession): async def test_create_role_name_conflict(self, test_db_session: AsyncSession): """Test role creation with existing name""" # Create existing role - existing_role = Roles( - id="existing-role", name="admin", description="Existing admin role" - ) + existing_role = Roles(id="existing-role", name="admin", description="Existing admin role") test_db_session.add(existing_role) await test_db_session.commit() - role_data = RoleCreate(name="admin", description="New admin role", + role_data = RoleCreate( + name="admin", + description="New admin role", level=10, ) @@ -185,14 +183,13 @@ async def test_create_role_name_conflict(self, test_db_session: AsyncSession): @pytest.mark.asyncio async def test_create_role_database_error(self, test_db_session: AsyncSession): """Test create_role with database error""" - role_data = RoleCreate(name="test-role", + role_data = RoleCreate( + name="test-role", level=10, ) actor = await _create_actor(test_db_session) - with patch.object( - test_db_session, "commit", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "commit", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: await create_role(test_db_session, role_data, actor_user_id=actor) @@ -291,9 +288,7 @@ async def test_update_role_database_error(self, test_db_session: AsyncSession): actor = await _create_actor(test_db_session) role_data = RoleUpdate(name="updated_role") - with patch.object( - test_db_session, "execute", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "execute", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: await update_role(test_db_session, "role-1", role_data, actor_user_id=actor) @@ -317,9 +312,7 @@ async def test_delete_role_success(self, test_db_session: AsyncSession): assert result is True # Verify role is deleted - deleted_role = await test_db_session.execute( - select(Roles).where(Roles.id == "role-1") - ) + deleted_role = await test_db_session.execute(select(Roles).where(Roles.id == "role-1")) assert deleted_role.scalar_one_or_none() is None @pytest.mark.asyncio @@ -386,9 +379,7 @@ async def test_delete_role_database_error(self, test_db_session: AsyncSession): actor = await _create_actor(test_db_session) - with patch.object( - test_db_session, "execute", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "execute", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: await delete_role(test_db_session, "role-1", actor_user_id=actor) @@ -399,19 +390,17 @@ class TestGetRoleAttributeMapping: """Test get_role_attribute_mapping service function""" @pytest.mark.asyncio - async def test_get_role_attribute_mapping_success( - self, test_db_session: AsyncSession - ): + async def test_get_role_attribute_mapping_success(self, test_db_session: AsyncSession): """Test successful role attributes mapping retrieval""" # Create test role and attributes role = Roles(id="role-1", name="admin", description="Admin role") - attr1 = RoleAttributes(id="attr-1", name="view-users", group="user-role-management", category="user") + attr1 = RoleAttributes( + id="attr-1", name="view-users", group="user-role-management", category="user" + ) attr2 = RoleAttributes( id="attr-2", name="manage-roles", group="user-role-management", category="role" ) - attr3 = RoleAttributes( - id="attr-3", name="edit-content", group=None, category=None - ) + attr3 = RoleAttributes(id="attr-3", name="edit-content", group=None, category=None) test_db_session.add(role) test_db_session.add(attr1) @@ -420,12 +409,8 @@ async def test_get_role_attribute_mapping_success( await test_db_session.commit() # Create attribute mappings - mapping1 = RoleAttributesMapper( - role_id="role-1", attributes_id="attr-1", value=True - ) - mapping2 = RoleAttributesMapper( - role_id="role-1", attributes_id="attr-2", value=False - ) + mapping1 = RoleAttributesMapper(role_id="role-1", attributes_id="attr-1", value=True) + mapping2 = RoleAttributesMapper(role_id="role-1", attributes_id="attr-2", value=False) test_db_session.add(mapping1) test_db_session.add(mapping2) @@ -433,7 +418,10 @@ async def test_get_role_attribute_mapping_success( result = await get_role_attribute_mapping(test_db_session, "role-1") - groups = {g.group: {cat: {a.name: a for a in attrs} for cat, attrs in g.categories.items()} for g in result.groups} + groups = { + g.group: {cat: {a.name: a for a in attrs} for cat, attrs in g.categories.items()} + for g in result.groups + } assert set(groups.keys()) == {"default", "user-role-management"} assert groups["user-role-management"]["user"]["view-users"].value is True @@ -444,9 +432,7 @@ async def test_get_role_attribute_mapping_success( assert groups["default"]["uncategorized"]["edit-content"].value is False @pytest.mark.asyncio - async def test_get_role_attribute_mapping_role_not_found( - self, test_db_session: AsyncSession - ): + async def test_get_role_attribute_mapping_role_not_found(self, test_db_session: AsyncSession): """Test get_role_attribute_mapping with non-existent role""" with pytest.raises(NotFoundException) as exc_info: await get_role_attribute_mapping(test_db_session, "non-existent-role") @@ -454,13 +440,9 @@ async def test_get_role_attribute_mapping_role_not_found( assert "Role not found" in str(exc_info.value) @pytest.mark.asyncio - async def test_get_role_attribute_mapping_database_error( - self, test_db_session: AsyncSession - ): + async def test_get_role_attribute_mapping_database_error(self, test_db_session: AsyncSession): """Test get_role_attribute_mapping with database error""" - with patch.object( - test_db_session, "execute", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "execute", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: await get_role_attribute_mapping(test_db_session, "role-1") @@ -471,16 +453,12 @@ class TestUpdateRoleAttributeMapping: """Test update_role_attribute_mapping service function""" @pytest.mark.asyncio - async def test_update_role_attribute_mapping_success( - self, test_db_session: AsyncSession - ): + async def test_update_role_attribute_mapping_success(self, test_db_session: AsyncSession): """Test successful role attributes mapping update""" # Create test role and attributes role = Roles(id="role-1", name="admin", description="Admin role") attr1 = RoleAttributes(id="attr-1", name="view-users") - attr2 = RoleAttributes( - id="attr-2", name="manage-roles" - ) + attr2 = RoleAttributes(id="attr-2", name="manage-roles") test_db_session.add(role) test_db_session.add(attr1) @@ -560,11 +538,11 @@ async def test_update_role_attribute_mapping_role_not_found( with pytest.raises(NotFoundException) as exc_info: await update_role_attribute_mapping( - test_db_session, - "non-existent-role", - attributes_data, - actor_user_id=await _create_actor(test_db_session), - ) + test_db_session, + "non-existent-role", + attributes_data, + actor_user_id=await _create_actor(test_db_session), + ) assert "Role not found" in str(exc_info.value) @@ -576,9 +554,7 @@ async def test_update_role_attribute_mapping_database_error( attributes_data = {"view-users": True} actor = await _create_actor(test_db_session) - with patch.object( - test_db_session, "execute", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "execute", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: await update_role_attribute_mapping( test_db_session, @@ -614,20 +590,12 @@ async def test_check_user_permissions_success(self, test_db_session: AsyncSessio # Use a different role name to avoid system super-admin protection role = Roles(id="role-1", name="test-role", description="Test role") attr1 = RoleAttributes(id="attr-1", name="view-users") - attr2 = RoleAttributes( - id="attr-2", name="manage-roles" - ) - attr3 = RoleAttributes( - id="attr-3", name="edit-content" - ) + attr2 = RoleAttributes(id="attr-2", name="manage-roles") + attr3 = RoleAttributes(id="attr-3", name="edit-content") role_mapping = RoleMapper(user_id="user-1", role_id="role-1") - attr_mapping1 = RoleAttributesMapper( - role_id="role-1", attributes_id="attr-1", value=True - ) - attr_mapping2 = RoleAttributesMapper( - role_id="role-1", attributes_id="attr-2", value=False - ) + attr_mapping1 = RoleAttributesMapper(role_id="role-1", attributes_id="attr-1", value=True) + attr_mapping2 = RoleAttributesMapper(role_id="role-1", attributes_id="attr-2", value=False) test_db_session.add(user) test_db_session.add(role) @@ -641,9 +609,7 @@ async def test_check_user_permissions_success(self, test_db_session: AsyncSessio required_attributes = ["view-users", "manage-roles", "edit-content"] - result = await check_user_permissions( - test_db_session, "user-1", required_attributes - ) + result = await check_user_permissions(test_db_session, "user-1", required_attributes) assert result.permissions["view-users"] is True assert result.permissions["manage-roles"] is False @@ -662,19 +628,13 @@ async def test_check_user_permissions_no_role(self, test_db_session: AsyncSessio assert result.permissions["manage-roles"] is False @pytest.mark.asyncio - async def test_check_user_permissions_database_error( - self, test_db_session: AsyncSession - ): + async def test_check_user_permissions_database_error(self, test_db_session: AsyncSession): """Test check_user_permissions with database error""" required_attributes = ["view-users"] - with patch.object( - test_db_session, "execute", side_effect=Exception("Database error") - ): + with patch.object(test_db_session, "execute", side_effect=Exception("Database error")): with pytest.raises(ServerException) as exc_info: - await check_user_permissions( - test_db_session, "user-1", required_attributes - ) + await check_user_permissions(test_db_session, "user-1", required_attributes) assert "Failed to check user permissions" in str(exc_info.value) @@ -683,9 +643,7 @@ class TestSuperAdminRoleProtection: """Test system super-admin role cannot be listed or mutated via roles API""" @pytest.mark.asyncio - async def test_get_all_roles_excludes_super_admin( - self, test_db_session: AsyncSession - ): + async def test_get_all_roles_excludes_super_admin(self, test_db_session: AsyncSession): super_role = Roles( id="role-super", name="super-admin", description="System super-admin", level=100 ) @@ -713,26 +671,22 @@ async def test_get_all_roles_excludes_super_admin( assert result.actor_role_id == "role-super" @pytest.mark.asyncio - async def test_create_role_rejects_super_admin_name( - self, test_db_session: AsyncSession - ): + async def test_create_role_rejects_super_admin_name(self, test_db_session: AsyncSession): with pytest.raises(AuthorizationException) as exc_info: await create_role( test_db_session, - RoleCreate(name="super-admin", description="blocked", - level=10, - ), + RoleCreate( + name="super-admin", + description="blocked", + level=10, + ), actor_user_id=await _create_actor(test_db_session), ) assert "Cannot create the system super-admin role" in str(exc_info.value) @pytest.mark.asyncio - async def test_update_role_rejects_mutating_super_admin( - self, test_db_session: AsyncSession - ): - role = Roles( - id="role-super", name="super-admin", description="System super-admin" - ) + async def test_update_role_rejects_mutating_super_admin(self, test_db_session: AsyncSession): + role = Roles(id="role-super", name="super-admin", description="System super-admin") test_db_session.add(role) await test_db_session.commit() @@ -746,9 +700,7 @@ async def test_update_role_rejects_mutating_super_admin( assert "Cannot modify the system super-admin role" in str(exc_info.value) @pytest.mark.asyncio - async def test_update_role_rejects_rename_to_super_admin( - self, test_db_session: AsyncSession - ): + async def test_update_role_rejects_rename_to_super_admin(self, test_db_session: AsyncSession): role = Roles(id="role-1", name="manager", description="Manager") test_db_session.add(role) await test_db_session.commit() @@ -760,17 +712,11 @@ async def test_update_role_rejects_rename_to_super_admin( RoleUpdate(name="super-admin"), actor_user_id=await _create_actor(test_db_session), ) - assert "Cannot rename a role to the system super-admin role" in str( - exc_info.value - ) + assert "Cannot rename a role to the system super-admin role" in str(exc_info.value) @pytest.mark.asyncio - async def test_delete_role_rejects_super_admin( - self, test_db_session: AsyncSession - ): - role = Roles( - id="role-super", name="super-admin", description="System super-admin" - ) + async def test_delete_role_rejects_super_admin(self, test_db_session: AsyncSession): + role = Roles(id="role-super", name="super-admin", description="System super-admin") test_db_session.add(role) await test_db_session.commit() @@ -780,12 +726,8 @@ async def test_delete_role_rejects_super_admin( assert "Cannot modify the system super-admin role" in str(exc_info.value) @pytest.mark.asyncio - async def test_get_role_attributes_rejects_super_admin( - self, test_db_session: AsyncSession - ): - role = Roles( - id="role-super", name="super-admin", description="System super-admin" - ) + async def test_get_role_attributes_rejects_super_admin(self, test_db_session: AsyncSession): + role = Roles(id="role-super", name="super-admin", description="System super-admin") test_db_session.add(role) await test_db_session.commit() @@ -794,20 +736,16 @@ async def test_get_role_attributes_rejects_super_admin( assert "Cannot modify the system super-admin role" in str(exc_info.value) @pytest.mark.asyncio - async def test_update_role_attributes_rejects_super_admin( - self, test_db_session: AsyncSession - ): - role = Roles( - id="role-super", name="super-admin", description="System super-admin" - ) + async def test_update_role_attributes_rejects_super_admin(self, test_db_session: AsyncSession): + role = Roles(id="role-super", name="super-admin", description="System super-admin") test_db_session.add(role) await test_db_session.commit() with pytest.raises(AuthorizationException) as exc_info: await update_role_attribute_mapping( - test_db_session, - "role-super", - {"view-users": True}, - actor_user_id=await _create_actor(test_db_session), - ) - assert "Cannot modify the system super-admin role" in str(exc_info.value) \ No newline at end of file + test_db_session, + "role-super", + {"view-users": True}, + actor_user_id=await _create_actor(test_db_session), + ) + assert "Cannot modify the system super-admin role" in str(exc_info.value) diff --git a/backend/tests/api/users/test_controller.py b/backend/tests/api/users/test_controller.py index b4a2fe8..0af1cf7 100644 --- a/backend/tests/api/users/test_controller.py +++ b/backend/tests/api/users/test_controller.py @@ -1,19 +1,18 @@ -import pytest from datetime import datetime -from httpx import AsyncClient from unittest.mock import patch -from utils.custom_exception import NotFoundException, ConflictException -from api.users.schema import UserResponse, UserPagination -from api.users.schema import UserDeleteBatchResponse, UserDeleteResult + +import pytest +from httpx import AsyncClient + +from api.users.schema import UserDeleteBatchResponse, UserDeleteResult, UserPagination, UserResponse +from utils.custom_exception import ConflictException, NotFoundException class TestUsersController: """Test Users controller API endpoints""" @pytest.mark.asyncio - async def test_get_users_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_get_users_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful users retrieval with valid authentication""" with patch("api.users.controller.get_all_users") as mock_get_users: mock_pagination = UserPagination( @@ -56,14 +55,10 @@ async def test_get_users_unauthorized(self, client: AsyncClient): assert response.status_code == 401 @pytest.mark.asyncio - async def test_get_users_with_filters( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_get_users_with_filters(self, client: AsyncClient, users_auth_headers: dict): """Test users retrieval with search filters""" with patch("api.users.controller.get_all_users") as mock_get_users: - mock_pagination = UserPagination( - users=[], total=0, page=1, per_page=10, total_pages=0 - ) + mock_pagination = UserPagination(users=[], total=0, page=1, per_page=10, total_pages=0) mock_get_users.return_value = mock_pagination response = await client.get( @@ -86,9 +81,7 @@ async def test_get_users_server_error(self, client: AsyncClient, users_auth_head assert response.status_code == 500 @pytest.mark.asyncio - async def test_create_user_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_user_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful user creation""" user_data = { "first_name": "John", @@ -126,9 +119,7 @@ async def test_create_user_success( assert data["data"]["email"] == user_data["email"] @pytest.mark.asyncio - async def test_create_user_email_exists( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_user_email_exists(self, client: AsyncClient, users_auth_headers: dict): """Test user creation with existing email""" user_data = { "first_name": "John", @@ -154,9 +145,7 @@ async def test_create_user_email_exists( assert data["message"] == "Email already exists" @pytest.mark.asyncio - async def test_create_user_server_error( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_user_server_error(self, client: AsyncClient, users_auth_headers: dict): """Test user creation with server error""" user_data = { "first_name": "John", @@ -179,9 +168,7 @@ async def test_create_user_server_error( assert response.status_code == 500 @pytest.mark.asyncio - async def test_update_user_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_user_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful user update""" user_id = "test-user-id" update_data = { @@ -216,9 +203,7 @@ async def test_update_user_success( assert data["data"]["first_name"] == update_data["first_name"] @pytest.mark.asyncio - async def test_update_user_not_found( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_user_not_found(self, client: AsyncClient, users_auth_headers: dict): """Test user update with non-existent user""" user_id = "nonexistent-user-id" update_data = {"first_name": "Updated"} @@ -238,9 +223,7 @@ async def test_update_user_not_found( assert data["message"] == "User not found" @pytest.mark.asyncio - async def test_update_user_email_conflict( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_update_user_email_conflict(self, client: AsyncClient, users_auth_headers: dict): """Test user update with email conflict""" user_id = "test-user-id" update_data = {"email": "existing@example.com"} @@ -258,22 +241,26 @@ async def test_update_user_email_conflict( assert data["message"] == "Email already exists" @pytest.mark.asyncio - async def test_delete_users_success( - self, client: AsyncClient, users_auth_headers: dict - ): - """Test successful users deletion""" + async def test_delete_users_success(self, client: AsyncClient, users_auth_headers: dict): + """Test successful users deletion""" delete_data = {"user_ids": ["user1", "user2", "user3"]} with patch("api.users.controller.delete_users") as mock_delete_users: mock_batch_response = UserDeleteBatchResponse( results=[ - UserDeleteResult(user_id="user1", status="success", message="User deleted successfully"), - UserDeleteResult(user_id="user2", status="success", message="User deleted successfully"), - UserDeleteResult(user_id="user3", status="success", message="User deleted successfully") + UserDeleteResult( + user_id="user1", status="success", message="User deleted successfully" + ), + UserDeleteResult( + user_id="user2", status="success", message="User deleted successfully" + ), + UserDeleteResult( + user_id="user3", status="success", message="User deleted successfully" + ), ], total_users=3, success_count=3, - failed_count=0 + failed_count=0, ) mock_delete_users.return_value = mock_batch_response @@ -296,18 +283,22 @@ async def test_delete_users_success( async def test_delete_users_partial_success( self, client: AsyncClient, users_auth_headers: dict ): - """Test users deletion with partial success""" + """Test users deletion with partial success""" delete_data = {"user_ids": ["user1", "nonexistent2"]} with patch("api.users.controller.delete_users") as mock_delete_users: mock_batch_response = UserDeleteBatchResponse( results=[ - UserDeleteResult(user_id="user1", status="success", message="User deleted successfully"), - UserDeleteResult(user_id="nonexistent2", status="failed", message="User not found") + UserDeleteResult( + user_id="user1", status="success", message="User deleted successfully" + ), + UserDeleteResult( + user_id="nonexistent2", status="failed", message="User not found" + ), ], total_users=2, success_count=1, - failed_count=1 + failed_count=1, ) mock_delete_users.return_value = mock_batch_response @@ -327,21 +318,23 @@ async def test_delete_users_partial_success( assert data["data"]["failed_count"] == 1 @pytest.mark.asyncio - async def test_delete_users_all_failed( - self, client: AsyncClient, users_auth_headers: dict - ): - """Test users deletion with all failed""" + async def test_delete_users_all_failed(self, client: AsyncClient, users_auth_headers: dict): + """Test users deletion with all failed""" delete_data = {"user_ids": ["nonexistent1", "nonexistent2"]} with patch("api.users.controller.delete_users") as mock_delete_users: mock_batch_response = UserDeleteBatchResponse( results=[ - UserDeleteResult(user_id="nonexistent1", status="failed", message="User not found"), - UserDeleteResult(user_id="nonexistent2", status="failed", message="User not found") + UserDeleteResult( + user_id="nonexistent1", status="failed", message="User not found" + ), + UserDeleteResult( + user_id="nonexistent2", status="failed", message="User not found" + ), ], total_users=2, success_count=0, - failed_count=2 + failed_count=2, ) mock_delete_users.return_value = mock_batch_response @@ -375,9 +368,7 @@ async def test_delete_users_server_error(self, client: AsyncClient, users_auth_h assert response.status_code == 500 @pytest.mark.asyncio - async def test_reset_password_success( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_reset_password_success(self, client: AsyncClient, users_auth_headers: dict): """Test successful password reset""" user_id = "test-user-id" password_data = {"new_password": "NewPassword123!"} @@ -394,10 +385,7 @@ async def test_reset_password_success( assert response.status_code == 200 data = response.json() assert data["code"] == 200 - assert ( - data["message"] - == "Password reset successfully and all devices logged out" - ) + assert data["message"] == "Password reset successfully and all devices logged out" @pytest.mark.asyncio async def test_reset_password_user_not_found( @@ -422,9 +410,7 @@ async def test_reset_password_user_not_found( assert data["message"] == "User not found" @pytest.mark.asyncio - async def test_reset_password_server_error( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_reset_password_server_error(self, client: AsyncClient, users_auth_headers: dict): """Test password reset with server error""" user_id = "test-user-id" password_data = {"new_password": "NewPassword123!"} @@ -445,9 +431,7 @@ class TestUsersControllerValidation: """Test Users controller input validation""" @pytest.mark.asyncio - async def test_create_user_invalid_email( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_user_invalid_email(self, client: AsyncClient, users_auth_headers: dict): """Test user creation with invalid email format""" user_data = { "first_name": "John", @@ -467,9 +451,7 @@ async def test_create_user_invalid_email( assert response.status_code == 422 @pytest.mark.asyncio - async def test_create_user_short_password( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_create_user_short_password(self, client: AsyncClient, users_auth_headers: dict): """Test user creation with password too short""" user_data = { "first_name": "John", @@ -501,9 +483,7 @@ async def test_get_users_invalid_pagination( assert response.status_code == 422 @pytest.mark.asyncio - async def test_delete_users_empty_list( - self, client: AsyncClient, users_auth_headers: dict - ): + async def test_delete_users_empty_list(self, client: AsyncClient, users_auth_headers: dict): """Test users deletion with empty user list""" delete_data = {"user_ids": []} @@ -514,4 +494,4 @@ async def test_delete_users_empty_list( headers={"Authorization": users_auth_headers["Authorization"]}, ) - assert response.status_code == 422 \ No newline at end of file + assert response.status_code == 422 diff --git a/backend/tests/api/users/test_schema.py b/backend/tests/api/users/test_schema.py index 8bdb7e8..9b50548 100644 --- a/backend/tests/api/users/test_schema.py +++ b/backend/tests/api/users/test_schema.py @@ -1,16 +1,18 @@ -import pytest from datetime import datetime + +import pytest from pydantic import ValidationError + from api.users.schema import ( - UserResponse, - UserPagination, - UserSortBy, + PasswordReset, UserCreate, - UserUpdate, UserDelete, - PasswordReset, - UserDeleteResult, UserDeleteBatchResponse, + UserDeleteResult, + UserPagination, + UserResponse, + UserSortBy, + UserUpdate, ) @@ -37,7 +39,7 @@ def test_user_response_valid_data(self): assert user.first_name == "John" assert user.last_name == "Doe" assert user.phone == "+1234567890" - assert user.status == True + assert user.status assert user.role == "admin" assert user.role_level is None @@ -189,7 +191,7 @@ def test_user_create_valid_data(self): assert user.email == "john.doe@example.com" assert user.phone == "+1234567890" assert user.password == "TestPassword123!" - assert user.status == True + assert user.status assert user.role == "admin" def test_user_create_without_optional_fields(self): @@ -204,7 +206,7 @@ def test_user_create_without_optional_fields(self): user = UserCreate(**user_data) - assert user.status == True # Default value + assert user.status # Default value assert user.role is None def test_user_create_invalid_email(self): @@ -298,7 +300,7 @@ def test_user_update_valid_data(self): assert user.last_name == "Name" assert user.email == "updated@example.com" assert user.phone == "+1234567890" - assert user.status == False + assert not user.status assert user.role == "user" def test_user_update_partial_data(self): @@ -486,7 +488,7 @@ def test_schema_serialization(self): assert serialized["email"] == "john.doe@example.com" assert serialized["phone"] == "+1234567890" assert serialized["password"] == "TestPassword123!" - assert serialized["status"] == True + assert serialized["status"] assert serialized["role"] == "admin" def test_schema_exclude_unset(self): @@ -510,7 +512,7 @@ def test_user_delete_result_success(self): result_data = { "user_id": "user123", "status": "success", - "message": "User deleted successfully" + "message": "User deleted successfully", } result = UserDeleteResult(**result_data) @@ -521,11 +523,7 @@ def test_user_delete_result_success(self): def test_user_delete_result_failed(self): """Test UserDeleteResult with failed status""" - result_data = { - "user_id": "user456", - "status": "failed", - "message": "User not found" - } + result_data = {"user_id": "user456", "status": "failed", "message": "User not found"} result = UserDeleteResult(**result_data) @@ -547,11 +545,7 @@ def test_user_delete_result_missing_fields(self): def test_user_delete_result_invalid_status(self): """Test UserDeleteResult with invalid status""" with pytest.raises(ValidationError) as exc_info: - UserDeleteResult( - user_id="user123", - status="invalid_status", - message="Test message" - ) + UserDeleteResult(user_id="user123", status="invalid_status", message="Test message") errors = exc_info.value.errors() assert any("status" in str(error) for error in errors) @@ -564,23 +558,14 @@ def test_user_delete_batch_response_success(self): """Test UserDeleteBatchResponse with all successful deletions""" results = [ UserDeleteResult( - user_id="user1", - status="success", - message="User deleted successfully" + user_id="user1", status="success", message="User deleted successfully" ), UserDeleteResult( - user_id="user2", - status="success", - message="User deleted successfully" - ) + user_id="user2", status="success", message="User deleted successfully" + ), ] - batch_data = { - "results": results, - "total_users": 2, - "success_count": 2, - "failed_count": 0 - } + batch_data = {"results": results, "total_users": 2, "success_count": 2, "failed_count": 0} batch_response = UserDeleteBatchResponse(**batch_data) @@ -594,23 +579,12 @@ def test_user_delete_batch_response_partial_success(self): """Test UserDeleteBatchResponse with partial success""" results = [ UserDeleteResult( - user_id="user1", - status="success", - message="User deleted successfully" + user_id="user1", status="success", message="User deleted successfully" ), - UserDeleteResult( - user_id="user2", - status="failed", - message="User not found" - ) + UserDeleteResult(user_id="user2", status="failed", message="User not found"), ] - batch_data = { - "results": results, - "total_users": 2, - "success_count": 1, - "failed_count": 1 - } + batch_data = {"results": results, "total_users": 2, "success_count": 1, "failed_count": 1} batch_response = UserDeleteBatchResponse(**batch_data) @@ -622,24 +596,11 @@ def test_user_delete_batch_response_partial_success(self): def test_user_delete_batch_response_all_failed(self): """Test UserDeleteBatchResponse with all failed deletions""" results = [ - UserDeleteResult( - user_id="user1", - status="failed", - message="User not found" - ), - UserDeleteResult( - user_id="user2", - status="failed", - message="User not found" - ) + UserDeleteResult(user_id="user1", status="failed", message="User not found"), + UserDeleteResult(user_id="user2", status="failed", message="User not found"), ] - batch_data = { - "results": results, - "total_users": 2, - "success_count": 0, - "failed_count": 2 - } + batch_data = {"results": results, "total_users": 2, "success_count": 0, "failed_count": 2} batch_response = UserDeleteBatchResponse(**batch_data) @@ -651,12 +612,7 @@ def test_user_delete_batch_response_all_failed(self): def test_user_delete_batch_response_empty_results(self): """Test UserDeleteBatchResponse with empty results""" - batch_data = { - "results": [], - "total_users": 0, - "success_count": 0, - "failed_count": 0 - } + batch_data = {"results": [], "total_users": 0, "success_count": 0, "failed_count": 0} batch_response = UserDeleteBatchResponse(**batch_data) @@ -679,19 +635,10 @@ def test_user_delete_batch_response_missing_fields(self): def test_user_delete_batch_response_serialization(self): """Test UserDeleteBatchResponse serialization""" results = [ - UserDeleteResult( - user_id="user1", - status="success", - message="User deleted successfully" - ) + UserDeleteResult(user_id="user1", status="success", message="User deleted successfully") ] - batch_data = { - "results": results, - "total_users": 1, - "success_count": 1, - "failed_count": 0 - } + batch_data = {"results": results, "total_users": 1, "success_count": 1, "failed_count": 0} batch_response = UserDeleteBatchResponse(**batch_data) serialized = batch_response.model_dump() @@ -714,27 +661,16 @@ def test_batch_delete_workflow(self): # 1. Create batch response with mixed results results = [ UserDeleteResult( - user_id="user1", - status="success", - message="User deleted successfully" + user_id="user1", status="success", message="User deleted successfully" ), + UserDeleteResult(user_id="user2", status="failed", message="User not found"), UserDeleteResult( - user_id="user2", - status="failed", - message="User not found" + user_id="user3", status="success", message="User deleted successfully" ), - UserDeleteResult( - user_id="user3", - status="success", - message="User deleted successfully" - ) ] batch_response = UserDeleteBatchResponse( - results=results, - total_users=3, - success_count=2, - failed_count=1 + results=results, total_users=3, success_count=2, failed_count=1 ) # 2. Verify response structure @@ -765,11 +701,11 @@ def test_batch_delete_response_codes(self): all_success = UserDeleteBatchResponse( results=[ UserDeleteResult(user_id="user1", status="success", message="Success"), - UserDeleteResult(user_id="user2", status="success", message="Success") + UserDeleteResult(user_id="user2", status="success", message="Success"), ], total_users=2, success_count=2, - failed_count=0 + failed_count=0, ) assert all_success.failed_count == 0 # Should return 200 @@ -777,22 +713,24 @@ def test_batch_delete_response_codes(self): partial_success = UserDeleteBatchResponse( results=[ UserDeleteResult(user_id="user1", status="success", message="Success"), - UserDeleteResult(user_id="user2", status="failed", message="Failed") + UserDeleteResult(user_id="user2", status="failed", message="Failed"), ], total_users=2, success_count=1, - failed_count=1 + failed_count=1, ) - assert partial_success.success_count > 0 and partial_success.failed_count > 0 # Should return 207 + assert ( + partial_success.success_count > 0 and partial_success.failed_count > 0 + ) # Should return 207 # All failed all_failed = UserDeleteBatchResponse( results=[ UserDeleteResult(user_id="user1", status="failed", message="Failed"), - UserDeleteResult(user_id="user2", status="failed", message="Failed") + UserDeleteResult(user_id="user2", status="failed", message="Failed"), ], total_users=2, success_count=0, - failed_count=2 + failed_count=2, ) - assert all_failed.success_count == 0 # Should return 400 \ No newline at end of file + assert all_failed.success_count == 0 # Should return 400 diff --git a/backend/tests/api/users/test_service.py b/backend/tests/api/users/test_service.py index a2fab1e..dd8f407 100644 --- a/backend/tests/api/users/test_service.py +++ b/backend/tests/api/users/test_service.py @@ -1,40 +1,42 @@ -import pytest -from unittest.mock import AsyncMock, patch from datetime import datetime, timedelta +from unittest.mock import AsyncMock, patch + +import pytest from sqlalchemy import text from sqlalchemy.ext.asyncio import AsyncSession -from models.users import Users -from models.roles import Roles -from models.role_mapper import RoleMapper -from models.login_logs import LoginLogs -from models.user_sessions import UserSessions -from models.password_reset_tokens import PasswordResetTokens -from models.email_verification_tokens import EmailVerificationTokens -from utils.custom_exception import ( - ConflictException, - NotFoundException, - ServerException, - AuthorizationException, + +from api.users.schema import ( + UserCreate, + UserDeleteBatchResponse, + UserPagination, + UserResponse, + UserUpdate, ) from api.users.services import ( - get_all_users, - create_user, - update_user, - delete_users, - reset_user_password, - _assign_user_role, - _update_user_role, _assert_can_manage_user_role, + _assign_user_role, + _delete_user_related_records, _get_user_role_name, _get_user_roles_map, - _delete_user_related_records, + _update_user_role, + create_user, + delete_users, + get_all_users, + reset_user_password, + update_user, ) -from api.users.schema import ( - UserCreate, - UserUpdate, - UserResponse, - UserPagination, - UserDeleteBatchResponse +from models.email_verification_tokens import EmailVerificationTokens +from models.login_logs import LoginLogs +from models.password_reset_tokens import PasswordResetTokens +from models.role_mapper import RoleMapper +from models.roles import Roles +from models.user_sessions import UserSessions +from models.users import Users +from utils.custom_exception import ( + AuthorizationException, + ConflictException, + NotFoundException, + ServerException, ) @@ -53,7 +55,7 @@ async def test_get_all_users_success(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -63,18 +65,14 @@ async def test_get_all_users_success(self, test_db_session: AsyncSession): phone="+1234567891", hash_password="hashed_password", status=False, - created_at=datetime.now() + created_at=datetime.now(), ) - + test_db_session.add(user1) test_db_session.add(user2) await test_db_session.commit() - result = await get_all_users( - db=test_db_session, - page=1, - per_page=10 - ) + result = await get_all_users(db=test_db_session, page=1, per_page=10) assert isinstance(result, UserPagination) assert result.total == 2 @@ -108,9 +106,7 @@ async def test_get_all_users_hides_super_admin_when_disabled( created_at=datetime.now(), ) user_role = Roles(id="role-user-list", name="user", description="", level=1) - super_role = Roles( - id="role-super-list", name="super-admin", description="", level=100 - ) + super_role = Roles(id="role-super-list", name="super-admin", description="", level=100) test_db_session.add_all([regular, super_user, user_role, super_role]) await test_db_session.commit() test_db_session.add_all( @@ -154,9 +150,7 @@ async def test_get_all_users_shows_super_admin_when_enabled( created_at=datetime.now(), ) user_role = Roles(id="role-user-list-2", name="user", description="", level=1) - super_role = Roles( - id="role-super-list-2", name="super-admin", description="", level=100 - ) + super_role = Roles(id="role-super-list-2", name="super-admin", description="", level=100) test_db_session.add_all([regular, super_user, user_role, super_role]) await test_db_session.commit() test_db_session.add_all( @@ -185,7 +179,7 @@ async def test_get_all_users_with_keyword(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -195,19 +189,14 @@ async def test_get_all_users_with_keyword(self, test_db_session: AsyncSession): phone="+1234567891", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) - + test_db_session.add(user1) test_db_session.add(user2) await test_db_session.commit() - result = await get_all_users( - db=test_db_session, - keyword="john", - page=1, - per_page=10 - ) + result = await get_all_users(db=test_db_session, keyword="john", page=1, per_page=10) assert result.total == 1 assert len(result.users) == 1 @@ -224,7 +213,7 @@ async def test_get_all_users_with_status_filter(self, test_db_session: AsyncSess phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -234,23 +223,18 @@ async def test_get_all_users_with_status_filter(self, test_db_session: AsyncSess phone="+1234567891", hash_password="hashed_password", status=False, - created_at=datetime.now() + created_at=datetime.now(), ) - + test_db_session.add(user1) test_db_session.add(user2) await test_db_session.commit() - result = await get_all_users( - db=test_db_session, - status="true", - page=1, - per_page=10 - ) + result = await get_all_users(db=test_db_session, status="true", page=1, per_page=10) assert result.total == 1 assert len(result.users) == 1 - assert result.users[0].status == True + assert result.users[0].status @pytest.mark.asyncio async def test_get_all_users_with_role_filter(self, test_db_session: AsyncSession): @@ -264,34 +248,22 @@ async def test_get_all_users_with_role_filter(self, test_db_session: AsyncSessio phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() # Create test role - role = Roles( - id="role1", - name="admin", - description="Administrator role" - ) + role = Roles(id="role1", name="admin", description="Administrator role") test_db_session.add(role) await test_db_session.commit() # Create role mapping - role_mapping = RoleMapper( - user_id="user1", - role_id="role1" - ) + role_mapping = RoleMapper(user_id="user1", role_id="role1") test_db_session.add(role_mapping) await test_db_session.commit() - result = await get_all_users( - db=test_db_session, - role="admin", - page=1, - per_page=10 - ) + result = await get_all_users(db=test_db_session, role="admin", page=1, per_page=10) assert result.total == 1 assert len(result.users) == 1 @@ -300,11 +272,7 @@ async def test_get_all_users_with_role_filter(self, test_db_session: AsyncSessio @pytest.mark.asyncio async def test_get_all_users_empty_result(self, test_db_session: AsyncSession): """Test users retrieval with no results""" - result = await get_all_users( - db=test_db_session, - page=1, - per_page=10 - ) + result = await get_all_users(db=test_db_session, page=1, per_page=10) assert result.total == 0 assert len(result.users) == 0 @@ -325,17 +293,13 @@ async def test_get_all_users_pagination(self, test_db_session: AsyncSession): phone=f"+123456789{i}", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) - + await test_db_session.commit() - result = await get_all_users( - db=test_db_session, - page=2, - per_page=10 - ) + result = await get_all_users(db=test_db_session, page=2, per_page=10) assert result.total == 15 assert len(result.users) == 5 # Second page should have 5 users @@ -354,7 +318,7 @@ async def test_get_all_users_sorting(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -364,19 +328,15 @@ async def test_get_all_users_sorting(self, test_db_session: AsyncSession): phone="+1234567891", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) - + test_db_session.add(user1) test_db_session.add(user2) await test_db_session.commit() result = await get_all_users( - db=test_db_session, - sort_by="email", - desc=False, - page=1, - per_page=10 + db=test_db_session, sort_by="email", desc=False, page=1, per_page=10 ) assert len(result.users) == 2 @@ -403,7 +363,7 @@ async def test_get_all_users_sort_by_role_desc(self, test_db_session: AsyncSessi phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -413,7 +373,7 @@ async def test_get_all_users_sort_by_role_desc(self, test_db_session: AsyncSessi phone="+1234567891", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) role_admin = Roles(id="role1", name="admin", description="Admin role") role_user = Roles(id="role2", name="user", description="User role") @@ -425,11 +385,7 @@ async def test_get_all_users_sort_by_role_desc(self, test_db_session: AsyncSessi await test_db_session.commit() result = await get_all_users( - db=test_db_session, - sort_by="role", - desc=True, - page=1, - per_page=10 + db=test_db_session, sort_by="role", desc=True, page=1, per_page=10 ) assert result.users[0].role == "user" @@ -472,9 +428,7 @@ async def test_get_all_users_sort_by_role_asc(self, test_db_session: AsyncSessio assert result.users[1].role == "user" @pytest.mark.asyncio - async def test_get_all_users_unknown_sort_falls_back( - self, test_db_session: AsyncSession - ): + async def test_get_all_users_unknown_sort_falls_back(self, test_db_session: AsyncSession): user = Users( id="user1", email="a@example.com", @@ -494,9 +448,7 @@ async def test_get_all_users_unknown_sort_falls_back( assert len(result.users) == 1 @pytest.mark.asyncio - async def test_get_all_users_multiple_status_values( - self, test_db_session: AsyncSession - ): + async def test_get_all_users_multiple_status_values(self, test_db_session: AsyncSession): active = Users( id="user1", email="active@example.com", @@ -520,9 +472,7 @@ async def test_get_all_users_multiple_status_values( test_db_session.add_all([active, inactive]) await test_db_session.commit() - result = await get_all_users( - db=test_db_session, status="true,false", page=1, per_page=10 - ) + result = await get_all_users(db=test_db_session, status="true,false", page=1, per_page=10) assert result.total == 2 """Test create_user service function""" @@ -537,14 +487,12 @@ async def test_create_user_success(self, test_db_session: AsyncSession): phone="+1234567890", password="TestPassword123!", status=True, - role="admin" + role="admin", ) with patch("api.users.services._assert_can_manage_user_role") as mock_assert_role: with patch("api.users.services._assign_user_role") as mock_assign_role: - result = await create_user( - test_db_session, user_data, actor_user_id="actor1" - ) + result = await create_user(test_db_session, user_data, actor_user_id="actor1") assert isinstance(result, UserResponse) assert result.email == user_data.email @@ -568,7 +516,7 @@ async def test_create_user_email_exists(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(existing_user) await test_db_session.commit() @@ -579,12 +527,12 @@ async def test_create_user_email_exists(self, test_db_session: AsyncSession): email="existing@example.com", phone="+1234567890", password="TestPassword123!", - status=True + status=True, ) with pytest.raises(ConflictException) as exc_info: await create_user(test_db_session, user_data, actor_user_id="actor1") - + assert "Email already exists" in str(exc_info.value) @pytest.mark.asyncio @@ -596,7 +544,7 @@ async def test_create_user_without_role(self, test_db_session: AsyncSession): email="john.doe@example.com", phone="+1234567890", password="TestPassword123!", - status=True + status=True, ) result = await create_user(test_db_session, user_data, actor_user_id="actor1") @@ -620,15 +568,13 @@ async def test_update_user_success(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() update_data = UserUpdate( - first_name="Updated", - last_name="Name", - email="updated@example.com" + first_name="Updated", last_name="Name", email="updated@example.com" ) with patch("api.users.services._update_user_role") as mock_update_role: @@ -643,7 +589,9 @@ async def test_update_user_success(self, test_db_session: AsyncSession): mock_update_role.assert_not_called() @pytest.mark.asyncio - async def test_update_user_email_clears_pending_verification(self, test_db_session: AsyncSession): + async def test_update_user_email_clears_pending_verification( + self, test_db_session: AsyncSession + ): """Admin email update should confirm the new address and clear pending verification""" user = Users( id="user1", @@ -655,7 +603,7 @@ async def test_update_user_email_clears_pending_verification(self, test_db_sessi phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() @@ -679,10 +627,8 @@ async def test_update_user_not_found(self, test_db_session: AsyncSession): update_data = UserUpdate(first_name="Updated") with pytest.raises(NotFoundException) as exc_info: - await update_user( - test_db_session, "nonexistent", update_data, actor_user_id="actor1" - ) - + await update_user(test_db_session, "nonexistent", update_data, actor_user_id="actor1") + assert "User not found" in str(exc_info.value) @pytest.mark.asyncio @@ -697,7 +643,7 @@ async def test_update_user_email_exists(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -707,7 +653,7 @@ async def test_update_user_email_exists(self, test_db_session: AsyncSession): phone="+1234567891", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user1) test_db_session.add(user2) @@ -716,10 +662,8 @@ async def test_update_user_email_exists(self, test_db_session: AsyncSession): update_data = UserUpdate(email="user2@example.com") with pytest.raises(ConflictException) as exc_info: - await update_user( - test_db_session, "user1", update_data, actor_user_id="actor1" - ) - + await update_user(test_db_session, "user1", update_data, actor_user_id="actor1") + assert "Email already exists" in str(exc_info.value) @pytest.mark.asyncio @@ -750,9 +694,7 @@ async def test_update_user_with_role(self, test_db_session: AsyncSession): mock_update_role.assert_awaited_once() @pytest.mark.asyncio - async def test_update_user_rejects_higher_role_level( - self, test_db_session: AsyncSession - ): + async def test_update_user_rejects_higher_role_level(self, test_db_session: AsyncSession): actor = Users( id="actor-upd", email="actor-upd@example.com", @@ -808,7 +750,7 @@ async def test_delete_users_success(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) user2 = Users( id="user2", @@ -818,16 +760,18 @@ async def test_delete_users_success(self, test_db_session: AsyncSession): phone="+1234567891", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user1) test_db_session.add(user2) await test_db_session.commit() mock_redis = AsyncMock() - - with patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions, \ - patch("api.users.services._delete_user_related_records") as mock_delete_related: + + with ( + patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions, + patch("api.users.services._delete_user_related_records") as mock_delete_related, + ): result = await delete_users(test_db_session, mock_redis, ["user1", "user2"]) assert isinstance(result, UserDeleteBatchResponse) @@ -840,9 +784,7 @@ async def test_delete_users_success(self, test_db_session: AsyncSession): mock_delete_related.assert_called() @pytest.mark.asyncio - async def test_delete_users_rejects_super_admin( - self, test_db_session: AsyncSession - ): + async def test_delete_users_rejects_super_admin(self, test_db_session: AsyncSession): """Test deleting a system super-admin user is rejected""" user = Users( id="super1", @@ -854,9 +796,7 @@ async def test_delete_users_rejects_super_admin( status=True, created_at=datetime.now(), ) - role = Roles( - id="role-super", name="super-admin", description="System super-admin" - ) + role = Roles(id="role-super", name="super-admin", description="System super-admin") mapping = RoleMapper(user_id="super1", role_id="role-super") test_db_session.add(user) test_db_session.add(role) @@ -887,15 +827,17 @@ async def test_delete_users_partial_success(self, test_db_session: AsyncSession) phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() mock_redis = AsyncMock() - - with patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions, \ - patch("api.users.services._delete_user_related_records") as mock_delete_related: + + with ( + patch("api.users.services.clear_user_all_sessions"), + patch("api.users.services._delete_user_related_records"), + ): result = await delete_users(test_db_session, mock_redis, ["user1", "nonexistent"]) assert isinstance(result, UserDeleteBatchResponse) @@ -903,11 +845,11 @@ async def test_delete_users_partial_success(self, test_db_session: AsyncSession) assert result.success_count == 1 assert result.failed_count == 1 assert len(result.results) == 2 - + # Check individual results success_results = [r for r in result.results if r.status == "success"] failed_results = [r for r in result.results if r.status == "failed"] - + assert len(success_results) == 1 assert len(failed_results) == 1 assert success_results[0].user_id == "user1" @@ -918,7 +860,7 @@ async def test_delete_users_partial_success(self, test_db_session: AsyncSession) async def test_delete_users_all_failed(self, test_db_session: AsyncSession): """Test users deletion with all failed""" mock_redis = AsyncMock() - + result = await delete_users(test_db_session, mock_redis, ["nonexistent1", "nonexistent2"]) assert isinstance(result, UserDeleteBatchResponse) @@ -941,17 +883,19 @@ async def test_delete_users_with_session_clear_failure(self, test_db_session: As phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() mock_redis = AsyncMock() - - with patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions, \ - patch("api.users.services._delete_user_related_records") as mock_delete_related: + + with ( + patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions, + patch("api.users.services._delete_user_related_records"), + ): mock_clear_sessions.side_effect = Exception("Redis connection failed") - + result = await delete_users(test_db_session, mock_redis, ["user1"]) assert isinstance(result, UserDeleteBatchResponse) @@ -973,15 +917,17 @@ async def test_delete_users_with_foreign_key_constraints(self, test_db_session: phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() mock_redis = AsyncMock() - - with patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions, \ - patch("api.users.services._delete_user_related_records") as mock_delete_related: + + with ( + patch("api.users.services.clear_user_all_sessions"), + patch("api.users.services._delete_user_related_records") as mock_delete_related, + ): result = await delete_users(test_db_session, mock_redis, ["user1"]) assert isinstance(result, UserDeleteBatchResponse) @@ -990,7 +936,7 @@ async def test_delete_users_with_foreign_key_constraints(self, test_db_session: assert result.failed_count == 0 assert result.results[0].status == "success" assert result.results[0].user_id == "user1" - + # Verify that related records deletion was called mock_delete_related.assert_called_once_with(test_db_session, "user1") @@ -1005,7 +951,7 @@ async def test_delete_users_skip_own_account(self, test_db_session: AsyncSession phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() @@ -1019,9 +965,7 @@ async def test_delete_users_skip_own_account(self, test_db_session: AsyncSession assert result.results[0].message == "Cannot delete your own account" @pytest.mark.asyncio - async def test_delete_users_rejects_higher_role_level( - self, test_db_session: AsyncSession - ): + async def test_delete_users_rejects_higher_role_level(self, test_db_session: AsyncSession): actor = Users( id="actor-del", email="actor-del@example.com", @@ -1081,22 +1025,19 @@ async def test_reset_password_success(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="old_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() mock_redis = AsyncMock() - + with patch("api.users.services.clear_user_all_sessions") as mock_clear_sessions: result = await reset_user_password( - test_db_session, - mock_redis, - "user1", - "NewPassword123!" + test_db_session, mock_redis, "user1", "NewPassword123!" ) - assert result == True + assert result mock_clear_sessions.assert_called_once() @pytest.mark.asyncio @@ -1105,13 +1046,8 @@ async def test_reset_password_user_not_found(self, test_db_session: AsyncSession mock_redis = AsyncMock() with pytest.raises(NotFoundException) as exc_info: - await reset_user_password( - test_db_session, - mock_redis, - "nonexistent", - "NewPassword123!" - ) - + await reset_user_password(test_db_session, mock_redis, "nonexistent", "NewPassword123!") + assert "User not found" in str(exc_info.value) @pytest.mark.asyncio @@ -1125,26 +1061,24 @@ async def test_reset_password_server_error(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="old_password", status=True, - created_at=datetime.now() + created_at=datetime.now(), ) test_db_session.add(user) await test_db_session.commit() mock_redis = AsyncMock() - with patch("api.users.services.clear_user_all_sessions", side_effect=Exception("Redis error")): + with patch( + "api.users.services.clear_user_all_sessions", side_effect=Exception("Redis error") + ): with pytest.raises(ServerException): - await reset_user_password( - test_db_session, mock_redis, "user1", "NewPassword123!" - ) + await reset_user_password(test_db_session, mock_redis, "user1", "NewPassword123!") class TestGetUserRoleName: """Test _get_user_role_name helper""" @pytest.mark.asyncio - async def test_get_user_role_name_returns_role( - self, test_db_session: AsyncSession - ): + async def test_get_user_role_name_returns_role(self, test_db_session: AsyncSession): user = Users( id="user1", email="user1@example.com", @@ -1190,9 +1124,7 @@ async def test_nobody_can_assign_super_admin_role(self, test_db_session: AsyncSe assert "Cannot assign the system super-admin role" in str(exc_info.value) @pytest.mark.asyncio - async def test_cannot_change_role_of_super_admin_user( - self, test_db_session: AsyncSession - ): + async def test_cannot_change_role_of_super_admin_user(self, test_db_session: AsyncSession): with patch( "api.users.services._get_user_role_name", new_callable=AsyncMock, @@ -1205,117 +1137,129 @@ async def test_cannot_change_role_of_super_admin_user( "user", target_user_id="super-user", ) - assert "Cannot change the role of a system super-admin user" in str( - exc_info.value - ) + assert "Cannot change the role of a system super-admin user" in str(exc_info.value) @pytest.mark.asyncio - async def test_super_admin_can_assign_non_system_role( - self, test_db_session: AsyncSession - ): - with patch( - "api.users.services.check_user_has_super_role", - new_callable=AsyncMock, - return_value=True, - ), patch( - "api.users.services._get_user_role_name", - new_callable=AsyncMock, - return_value="user", + async def test_super_admin_can_assign_non_system_role(self, test_db_session: AsyncSession): + with ( + patch( + "api.users.services.check_user_has_super_role", + new_callable=AsyncMock, + return_value=True, + ), + patch( + "api.users.services._get_user_role_name", + new_callable=AsyncMock, + return_value="user", + ), ): await _assert_can_manage_user_role( test_db_session, "actor1", "user", target_user_id="user1" ) @pytest.mark.asyncio - async def test_manage_roles_required_for_role_change( - self, test_db_session: AsyncSession - ): - with patch( - "api.users.services.check_user_has_super_role", - new_callable=AsyncMock, - return_value=False, - ), patch( - "api.users.services.get_user_attributes", - new_callable=AsyncMock, - return_value={"manage-users": True}, + async def test_manage_roles_required_for_role_change(self, test_db_session: AsyncSession): + with ( + patch( + "api.users.services.check_user_has_super_role", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "api.users.services.get_user_attributes", + new_callable=AsyncMock, + return_value={"manage-users": True}, + ), ): with pytest.raises(AuthorizationException) as exc_info: await _assert_can_manage_user_role(test_db_session, "actor1", "user") assert "Permission denied to assign roles" in str(exc_info.value) @pytest.mark.asyncio - async def test_manage_roles_can_assign_non_super_role( - self, test_db_session: AsyncSession - ): - with patch( - "api.users.services.check_user_has_super_role", - new_callable=AsyncMock, - return_value=False, - ), patch( - "api.users.services.get_user_attributes", - new_callable=AsyncMock, - return_value={"manage-roles": True}, - ), patch( - "api.users.services.get_user_role_level", - new_callable=AsyncMock, - return_value=50, - ), patch( - "api.users.services._get_role_level_by_name", - new_callable=AsyncMock, - return_value=10, + async def test_manage_roles_can_assign_non_super_role(self, test_db_session: AsyncSession): + with ( + patch( + "api.users.services.check_user_has_super_role", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "api.users.services.get_user_attributes", + new_callable=AsyncMock, + return_value={"manage-roles": True}, + ), + patch( + "api.users.services.get_user_role_level", + new_callable=AsyncMock, + return_value=50, + ), + patch( + "api.users.services._get_role_level_by_name", + new_callable=AsyncMock, + return_value=10, + ), ): await _assert_can_manage_user_role(test_db_session, "actor1", "user") @pytest.mark.asyncio async def test_cannot_assign_higher_role_level(self, test_db_session: AsyncSession): - with patch( - "api.users.services.check_user_has_super_role", - new_callable=AsyncMock, - return_value=False, - ), patch( - "api.users.services.get_user_attributes", - new_callable=AsyncMock, - return_value={"manage-roles": True}, - ), patch( - "api.users.services.get_user_role_level", - new_callable=AsyncMock, - return_value=20, - ), patch( - "api.users.services._get_role_level_by_name", - new_callable=AsyncMock, - return_value=50, + with ( + patch( + "api.users.services.check_user_has_super_role", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "api.users.services.get_user_attributes", + new_callable=AsyncMock, + return_value={"manage-roles": True}, + ), + patch( + "api.users.services.get_user_role_level", + new_callable=AsyncMock, + return_value=20, + ), + patch( + "api.users.services._get_role_level_by_name", + new_callable=AsyncMock, + return_value=50, + ), ): with pytest.raises(AuthorizationException) as exc_info: await _assert_can_manage_user_role(test_db_session, "actor1", "boss") assert "higher level" in str(exc_info.value) @pytest.mark.asyncio - async def test_cannot_manage_user_with_higher_role_level( - self, test_db_session: AsyncSession - ): + async def test_cannot_manage_user_with_higher_role_level(self, test_db_session: AsyncSession): async def fake_level(user_id, _db): return 80 if user_id == "target-high" else 20 - with patch( - "api.users.services.check_user_has_super_role", - new_callable=AsyncMock, - return_value=False, - ), patch( - "api.users.services.get_user_attributes", - new_callable=AsyncMock, - return_value={"manage-roles": True}, - ), patch( - "api.users.services._get_user_role_name", - new_callable=AsyncMock, - return_value="manager", - ), patch( - "api.users.services.get_user_role_level", - new_callable=AsyncMock, - side_effect=fake_level, - ), patch( - "api.users.services._get_role_level_by_name", - new_callable=AsyncMock, - return_value=10, + with ( + patch( + "api.users.services.check_user_has_super_role", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "api.users.services.get_user_attributes", + new_callable=AsyncMock, + return_value={"manage-roles": True}, + ), + patch( + "api.users.services._get_user_role_name", + new_callable=AsyncMock, + return_value="manager", + ), + patch( + "api.users.services.get_user_role_level", + new_callable=AsyncMock, + side_effect=fake_level, + ), + patch( + "api.users.services._get_role_level_by_name", + new_callable=AsyncMock, + return_value=10, + ), ): with pytest.raises(AuthorizationException) as exc_info: await _assert_can_manage_user_role( @@ -1331,9 +1275,7 @@ class TestUpdateUserOwnRole: """Test users cannot change their own role""" @pytest.mark.asyncio - async def test_update_user_rejects_own_role_change( - self, test_db_session: AsyncSession - ): + async def test_update_user_rejects_own_role_change(self, test_db_session: AsyncSession): user = Users( id="self-user", email="self@example.com", @@ -1372,13 +1314,9 @@ async def test_assign_user_role_success(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() - ) - role = Roles( - id="role1", - name="admin", - description="Administrator role" + created_at=datetime.now(), ) + role = Roles(id="role1", name="admin", description="Administrator role") test_db_session.add(user) test_db_session.add(role) await test_db_session.commit() @@ -1397,7 +1335,7 @@ async def test_assign_user_role_role_not_found(self, test_db_session: AsyncSessi """Test role assignment with non-existent role""" with pytest.raises(NotFoundException) as exc_info: await _assign_user_role(test_db_session, "user1", "nonexistent_role") - + assert "Role 'nonexistent_role' not found" in str(exc_info.value) @pytest.mark.asyncio @@ -1411,13 +1349,9 @@ async def test_assign_user_role_existing_mapping(self, test_db_session: AsyncSes phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() - ) - role = Roles( - id="role1", - name="admin", - description="Administrator role" + created_at=datetime.now(), ) + role = Roles(id="role1", name="admin", description="Administrator role") test_db_session.add(user) test_db_session.add(role) await test_db_session.commit() @@ -1444,28 +1378,17 @@ async def test_update_user_role_success(self, test_db_session: AsyncSession): phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() - ) - old_role = Roles( - id="role1", - name="old_role", - description="Old role" - ) - new_role = Roles( - id="role2", - name="new_role", - description="New role" + created_at=datetime.now(), ) + old_role = Roles(id="role1", name="old_role", description="Old role") + new_role = Roles(id="role2", name="new_role", description="New role") test_db_session.add(user) test_db_session.add(old_role) test_db_session.add(new_role) await test_db_session.commit() # Create existing role mapping - role_mapping = RoleMapper( - user_id="user1", - role_id="role1" - ) + role_mapping = RoleMapper(user_id="user1", role_id="role1") test_db_session.add(role_mapping) await test_db_session.commit() @@ -1485,22 +1408,15 @@ async def test_update_user_role_remove_only(self, test_db_session: AsyncSession) phone="+1234567890", hash_password="hashed_password", status=True, - created_at=datetime.now() - ) - role = Roles( - id="role1", - name="admin", - description="Administrator role" + created_at=datetime.now(), ) + role = Roles(id="role1", name="admin", description="Administrator role") test_db_session.add(user) test_db_session.add(role) await test_db_session.commit() # Create existing role mapping - role_mapping = RoleMapper( - user_id="user1", - role_id="role1" - ) + role_mapping = RoleMapper(user_id="user1", role_id="role1") test_db_session.add(role_mapping) await test_db_session.commit() @@ -1601,4 +1517,4 @@ async def test_delete_user_related_records_success(self, test_db_session: AsyncS class TestGetUserRolesMap: @pytest.mark.asyncio async def test_get_user_roles_map_empty(self, test_db_session: AsyncSession): - assert await _get_user_roles_map(test_db_session, []) == {} \ No newline at end of file + assert await _get_user_roles_map(test_db_session, []) == {} diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index b41fe9b..76b2d5b 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -3,16 +3,18 @@ os.environ.setdefault("OTEL_ENABLE", "false") os.environ.setdefault("LOG_HTTP_BODY", "false") -import pytest import asyncio -import pytest_asyncio -from uuid_utils import uuid7 from datetime import datetime, timedelta from unittest.mock import AsyncMock, patch -from httpx import AsyncClient, ASGITransport + +import pytest +import pytest_asyncio +from httpx import ASGITransport, AsyncClient from sqlalchemy import text -from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.pool import StaticPool +from uuid_utils import uuid7 + from core.config import settings from core.database import Base, make_async_url from core.dependencies import get_db @@ -22,12 +24,12 @@ with patch("core.config.setup_logging"): from main import app from models.password_reset_tokens import PasswordResetTokens -from models.user_sessions import UserSessions -from models.users import Users -from models.roles import Roles from models.role_attributes import RoleAttributes from models.role_attributes_mapper import RoleAttributesMapper from models.role_mapper import RoleMapper +from models.roles import Roles +from models.user_sessions import UserSessions +from models.users import Users mock_redis = AsyncMock() mock_redis.evalsha.return_value = 1000 # Indicates 1000ms remaining @@ -136,10 +138,12 @@ def event_loop(): @pytest.fixture(autouse=True) def mock_other_components(): """Mock startup components that are not needed for testing""" - with patch("schedule.register_schedules"), patch("schedule.scheduler.start"), patch( - "schedule.scheduler.shutdown" - ), patch("core.redis.init_redis", new=AsyncMock()), patch( - "core.config.setup_logging" + with ( + patch("schedule.register_schedules"), + patch("schedule.scheduler.start"), + patch("schedule.scheduler.shutdown"), + patch("core.redis.init_redis", new=AsyncMock()), + patch("core.config.setup_logging"), ): yield @@ -264,7 +268,6 @@ async def mock_get(key): } - @pytest_asyncio.fixture(scope="function") async def account_test_user(test_db_session: AsyncSession): """Create a dedicated test user for account module tests""" @@ -350,6 +353,7 @@ async def mock_get(key): "access_token": user_session.jwt_access_token, } + # Users module specific fixtures @pytest_asyncio.fixture(scope="function") async def users_test_user(test_db_session: AsyncSession): @@ -382,9 +386,7 @@ async def users_test_role(test_db_session: AsyncSession): """Create a non-system admin role for users/roles API tests""" role_id = str(uuid7()) role = Roles( - id=role_id, - name="admin", - description="Administrator role with management permissions" + id=role_id, name="admin", description="Administrator role with management permissions" ) test_db_session.add(role) await test_db_session.commit() @@ -395,45 +397,37 @@ async def users_test_role(test_db_session: AsyncSession): @pytest_asyncio.fixture(scope="function") async def users_test_role_attributes(test_db_session: AsyncSession, users_test_role: Roles): """Create role attributes and mappings for users/roles management permissions""" - + attribute_names = [ "view-users", "manage-users", "view-roles", "manage-roles", ] - attributes = [ - RoleAttributes(id=str(uuid7()), name=name) - for name in attribute_names - ] - + attributes = [RoleAttributes(id=str(uuid7()), name=name) for name in attribute_names] + test_db_session.add_all(attributes) await test_db_session.commit() - + attribute_mappings = [ - RoleAttributesMapper( - role_id=users_test_role.id, - attributes_id=attr.id, - value=True - ) + RoleAttributesMapper(role_id=users_test_role.id, attributes_id=attr.id, value=True) for attr in attributes ] - + for mapping in attribute_mappings: test_db_session.add(mapping) - + await test_db_session.commit() return attribute_mappings @pytest_asyncio.fixture(scope="function") -async def users_test_role_mapping(test_db_session: AsyncSession, users_test_user: Users, users_test_role: Roles): - """Create role mapping for test user""" - role_mapping = RoleMapper( - user_id=users_test_user.id, - role_id=users_test_role.id - ) - +async def users_test_role_mapping( + test_db_session: AsyncSession, users_test_user: Users, users_test_role: Roles +): + """Create role mapping for test user""" + role_mapping = RoleMapper(user_id=users_test_user.id, role_id=users_test_role.id) + test_db_session.add(role_mapping) await test_db_session.commit() return role_mapping @@ -472,10 +466,10 @@ async def users_test_session(test_db_session: AsyncSession, users_test_user: Use @pytest_asyncio.fixture(scope="function") async def users_auth_headers( - users_test_user: Users, - users_test_session: UserSessions, + users_test_user: Users, + users_test_session: UserSessions, users_test_role_mapping, - users_test_role_attributes + users_test_role_attributes, ): """Generate authentication headers for users API tests""" session_data = { @@ -513,37 +507,39 @@ async def mock_get(key): @pytest_asyncio.fixture async def create_password_reset_token_with_invalid_user(test_db_session: AsyncSession): """Create a password reset token with invalid user for testing edge cases""" - + nonexistent_user_id = "nonexistent_user_id" token_id = str(uuid7()) token_string = "test_token_123" expires_at = datetime.now() + timedelta(minutes=30) - + # Temporarily disable foreign key checks to insert invalid data await test_db_session.execute(text("SET FOREIGN_KEY_CHECKS = 0")) - await test_db_session.execute(text(""" - INSERT INTO password_reset_tokens (id, user_id, token, is_used, expires_at, created_at, updated_at) + await test_db_session.execute( + text(""" + INSERT INTO password_reset_tokens ( + id, user_id, token, is_used, expires_at, created_at, updated_at + ) VALUES (:id, :user_id, :token, :is_used, :expires_at, NOW(), NOW()) - """), { - "id": token_id, - "user_id": nonexistent_user_id, - "token": token_string, - "is_used": False, - "expires_at": expires_at - }) + """), + { + "id": token_id, + "user_id": nonexistent_user_id, + "token": token_string, + "is_used": False, + "expires_at": expires_at, + }, + ) # Re-enable foreign key checks await test_db_session.execute(text("SET FOREIGN_KEY_CHECKS = 1")) await test_db_session.commit() - + return { "token_id": token_id, "user_id": nonexistent_user_id, "token": token_string, "expires_at": expires_at, - "token_data": { - "sub": nonexistent_user_id, - "token": token_string - } + "token_data": {"sub": nonexistent_user_id, "token": token_string}, } @@ -552,7 +548,7 @@ async def create_password_reset_token_with_valid_user(test_db_session: AsyncSess """Create a password reset token with valid user for testing normal scenarios""" user_id = str(uuid7()) hashed_pwd = await hash_password("TestPassword123!") - + user = Users( id=user_id, email="tokenuser@example.com", @@ -561,25 +557,21 @@ async def create_password_reset_token_with_valid_user(test_db_session: AsyncSess phone="+1234567890", hash_password=hashed_pwd, status=True, - password_reset_required=True + password_reset_required=True, ) test_db_session.add(user) await test_db_session.commit() - + token_id = str(uuid7()) token_string = "valid_test_token_123" expires_at = datetime.now() + timedelta(minutes=30) - + reset_token_record = PasswordResetTokens( - id=token_id, - user_id=user_id, - token=token_string, - is_used=False, - expires_at=expires_at + id=token_id, user_id=user_id, token=token_string, is_used=False, expires_at=expires_at ) test_db_session.add(reset_token_record) await test_db_session.commit() - + return { "token_id": token_id, "user_id": user_id, @@ -587,8 +579,5 @@ async def create_password_reset_token_with_valid_user(test_db_session: AsyncSess "expires_at": expires_at, "user": user, "token_record": reset_token_record, - "token_data": { - "sub": user_id, - "token": token_string - } + "token_data": {"sub": user_id, "token": token_string}, } diff --git a/backend/tests/core/test_config.py b/backend/tests/core/test_config.py index 859688e..67b1dcc 100644 --- a/backend/tests/core/test_config.py +++ b/backend/tests/core/test_config.py @@ -12,9 +12,7 @@ class TestSkipConfig: def test_skip_paths_include_docs_and_health(self): - assert SKIP_PATHS == frozenset( - {"/", "/docs", "/redoc", "/openapi.json", "/healthz"} - ) + assert SKIP_PATHS == frozenset({"/", "/docs", "/redoc", "/openapi.json", "/healthz"}) def test_skip_methods_include_options(self): assert SKIP_METHODS == frozenset({"OPTIONS"}) diff --git a/backend/tests/core/test_init_db.py b/backend/tests/core/test_init_db.py index 5c23e21..16bdbf3 100644 --- a/backend/tests/core/test_init_db.py +++ b/backend/tests/core/test_init_db.py @@ -86,9 +86,7 @@ async def test_true_when_seeded(self, test_engine): hash_password="hashed", status=True, ) - db.add_all( - [role, user, RoleAttributes(id=str(uuid7()), name="view-users")] - ) + db.add_all([role, user, RoleAttributes(id=str(uuid7()), name="view-users")]) await db.commit() db.add(RoleMapper(user_id=user.id, role_id=role.id)) await db.commit() @@ -119,8 +117,9 @@ async def test_creates_and_updates_attributes(self, test_engine): @pytest.mark.asyncio async def test_rolls_back_on_error(self, test_engine): factory = _session_factory(test_engine) - with patch("core.init_db.AsyncSessionLocal", factory), patch( - "core.init_db.get_attributes", side_effect=RuntimeError("boom") + with ( + patch("core.init_db.AsyncSessionLocal", factory), + patch("core.init_db.get_attributes", side_effect=RuntimeError("boom")), ): with pytest.raises(RuntimeError, match="boom"): await create_role_attributes() @@ -142,9 +141,12 @@ async def test_creates_super_admin_and_user_roles(self, test_engine): @pytest.mark.asyncio async def test_rolls_back_on_error(self, test_engine): factory = _session_factory(test_engine) - with patch("core.init_db.AsyncSessionLocal", factory), patch( - "sqlalchemy.ext.asyncio.session.AsyncSession.add", - side_effect=RuntimeError("boom"), + with ( + patch("core.init_db.AsyncSessionLocal", factory), + patch( + "sqlalchemy.ext.asyncio.session.AsyncSession.add", + side_effect=RuntimeError("boom"), + ), ): with pytest.raises(RuntimeError, match="boom"): await create_default_roles() @@ -169,14 +171,10 @@ async def test_creates_admin_user(self, test_engine): async with factory() as db: user = ( - await db.execute( - select(Users).where(Users.email == settings.DEFAULT_ADMIN_EMAIL) - ) + await db.execute(select(Users).where(Users.email == settings.DEFAULT_ADMIN_EMAIL)) ).scalar_one() mapping = ( - await db.execute( - select(RoleMapper).where(RoleMapper.user_id == user.id) - ) + await db.execute(select(RoleMapper).where(RoleMapper.user_id == user.id)) ).scalar_one() assert user.email_verified is True assert mapping.role_id is not None @@ -212,9 +210,7 @@ async def test_assigns_role_to_existing_email(self, test_engine): db.add(user) await db.commit() - user_role = ( - await db.execute(select(Roles).where(Roles.name == "user")) - ).scalar_one() + user_role = (await db.execute(select(Roles).where(Roles.name == "user"))).scalar_one() db.add(RoleMapper(user_id=user.id, role_id=user_role.id)) await db.commit() @@ -223,9 +219,7 @@ async def test_assigns_role_to_existing_email(self, test_engine): async with factory() as db: user = ( - await db.execute( - select(Users).where(Users.email == settings.DEFAULT_ADMIN_EMAIL) - ) + await db.execute(select(Users).where(Users.email == settings.DEFAULT_ADMIN_EMAIL)) ).scalar_one() super_role = ( await db.execute( @@ -233,10 +227,10 @@ async def test_assigns_role_to_existing_email(self, test_engine): ) ).scalar_one() mappings = ( - await db.execute( - select(RoleMapper).where(RoleMapper.user_id == user.id) - ) - ).scalars().all() + (await db.execute(select(RoleMapper).where(RoleMapper.user_id == user.id))) + .scalars() + .all() + ) assert user.first_name == "Existing" assert len(mappings) == 1 assert mappings[0].role_id == super_role.id @@ -244,8 +238,9 @@ async def test_assigns_role_to_existing_email(self, test_engine): @pytest.mark.asyncio async def test_rolls_back_on_error(self, test_engine): factory = _session_factory(test_engine) - with patch("core.init_db.AsyncSessionLocal", factory), patch( - "core.init_db.select", side_effect=RuntimeError("boom") + with ( + patch("core.init_db.AsyncSessionLocal", factory), + patch("core.init_db.select", side_effect=RuntimeError("boom")), ): with pytest.raises(RuntimeError, match="boom"): await create_default_admin() @@ -269,12 +264,15 @@ async def test_skips_when_lock_not_acquired(self): lock_result = MagicMock() lock_result.scalar.return_value = 0 session.execute.return_value = lock_result - with patch( - "core.init_db.AsyncSessionLocal", - return_value=_LockSessionCM(session), - ), patch( - "core.init_db.is_already_initialized", new_callable=AsyncMock - ) as mock_initialized: + with ( + patch( + "core.init_db.AsyncSessionLocal", + return_value=_LockSessionCM(session), + ), + patch( + "core.init_db.is_already_initialized", new_callable=AsyncMock + ) as mock_initialized, + ): await init_database() mock_initialized.assert_not_called() @@ -284,16 +282,18 @@ async def test_skips_when_already_initialized(self): lock_result = MagicMock() lock_result.scalar.return_value = 1 session.execute.return_value = lock_result - with patch( - "core.init_db.AsyncSessionLocal", - return_value=_LockSessionCM(session), - ), patch( - "core.init_db.is_already_initialized", - new_callable=AsyncMock, - return_value=True, - ), patch( - "core.init_db.create_role_attributes", new_callable=AsyncMock - ) as mock_attrs: + with ( + patch( + "core.init_db.AsyncSessionLocal", + return_value=_LockSessionCM(session), + ), + patch( + "core.init_db.is_already_initialized", + new_callable=AsyncMock, + return_value=True, + ), + patch("core.init_db.create_role_attributes", new_callable=AsyncMock) as mock_attrs, + ): await init_database() mock_attrs.assert_not_called() assert session.execute.await_count >= 2 @@ -305,20 +305,20 @@ async def test_seeds_when_lock_acquired(self): lock_result = MagicMock() lock_result.scalar.return_value = 1 session.execute.return_value = lock_result - with patch( - "core.init_db.AsyncSessionLocal", - return_value=_LockSessionCM(session), - ), patch( - "core.init_db.is_already_initialized", - new_callable=AsyncMock, - return_value=False, - ), patch( - "core.init_db.create_role_attributes", new_callable=AsyncMock - ) as mock_attrs, patch( - "core.init_db.create_default_roles", new_callable=AsyncMock - ) as mock_roles, patch( - "core.init_db.create_default_admin", new_callable=AsyncMock - ) as mock_admin: + with ( + patch( + "core.init_db.AsyncSessionLocal", + return_value=_LockSessionCM(session), + ), + patch( + "core.init_db.is_already_initialized", + new_callable=AsyncMock, + return_value=False, + ), + patch("core.init_db.create_role_attributes", new_callable=AsyncMock) as mock_attrs, + patch("core.init_db.create_default_roles", new_callable=AsyncMock) as mock_roles, + patch("core.init_db.create_default_admin", new_callable=AsyncMock) as mock_admin, + ): await init_database() mock_attrs.assert_awaited_once() mock_roles.assert_awaited_once() @@ -330,17 +330,21 @@ async def test_releases_lock_on_error(self): lock_result = MagicMock() lock_result.scalar.return_value = 1 session.execute.return_value = lock_result - with patch( - "core.init_db.AsyncSessionLocal", - return_value=_LockSessionCM(session), - ), patch( - "core.init_db.is_already_initialized", - new_callable=AsyncMock, - return_value=False, - ), patch( - "core.init_db.create_role_attributes", - new_callable=AsyncMock, - side_effect=RuntimeError("seed failed"), + with ( + patch( + "core.init_db.AsyncSessionLocal", + return_value=_LockSessionCM(session), + ), + patch( + "core.init_db.is_already_initialized", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "core.init_db.create_role_attributes", + new_callable=AsyncMock, + side_effect=RuntimeError("seed failed"), + ), ): with pytest.raises(RuntimeError, match="seed failed"): await init_database() diff --git a/backend/tests/core/test_rbac.py b/backend/tests/core/test_rbac.py index 93ffaaf..c70e4c7 100644 --- a/backend/tests/core/test_rbac.py +++ b/backend/tests/core/test_rbac.py @@ -92,12 +92,8 @@ async def test_merges_attribute_values_with_or( [ RoleMapper(user_id=test_user.id, role_id=role_true.id), RoleMapper(user_id=test_user.id, role_id=role_false.id), - RoleAttributesMapper( - role_id=role_true.id, attributes_id=attr.id, value=True - ), - RoleAttributesMapper( - role_id=role_false.id, attributes_id=attr.id, value=False - ), + RoleAttributesMapper(role_id=role_true.id, attributes_id=attr.id, value=True), + RoleAttributesMapper(role_id=role_false.id, attributes_id=attr.id, value=False), ] ) await test_db_session.commit() @@ -124,9 +120,7 @@ async def protected(*, token, db=None): assert exc.value.status_code == 500 @pytest.mark.asyncio - async def test_super_admin_bypasses( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_super_admin_bypasses(self, test_db_session: AsyncSession, test_user: Users): role = Roles( id=str(uuid7()), name=settings.DEFAULT_SUPER_ADMIN_ROLE, @@ -155,9 +149,7 @@ async def test_allows_when_user_has_permission( test_db_session.add_all( [ RoleMapper(user_id=test_user.id, role_id=role.id), - RoleAttributesMapper( - role_id=role.id, attributes_id=attr.id, value=True - ), + RoleAttributesMapper(role_id=role.id, attributes_id=attr.id, value=True), ] ) await test_db_session.commit() @@ -170,9 +162,7 @@ async def protected(*, token, db): assert result == "ok" @pytest.mark.asyncio - async def test_denies_without_permission( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_denies_without_permission(self, test_db_session: AsyncSession, test_user: Users): @require_permission(["view-users"]) async def protected(*, token, db): return "ok" diff --git a/backend/tests/core/test_security.py b/backend/tests/core/test_security.py index 39ac0bc..96035d6 100644 --- a/backend/tests/core/test_security.py +++ b/backend/tests/core/test_security.py @@ -81,9 +81,7 @@ async def test_create_csrf_token_server_error(self): @pytest.mark.asyncio async def test_create_email_verification_token(self): - token = await create_email_verification_token( - "user-1", "a@b.c", "registration" - ) + token = await create_email_verification_token("user-1", "a@b.c", "registration") payload = _decode(token) assert payload["token_type"] == "email_verification" assert payload["verification_type"] == "registration" @@ -92,9 +90,7 @@ async def test_create_email_verification_token(self): @pytest.mark.asyncio async def test_create_email_verification_token_server_error(self): with patch("core.security.jwt.encode", side_effect=RuntimeError("encode failed")): - with pytest.raises( - ServerException, match="Failed to create email verification token" - ): + with pytest.raises(ServerException, match="Failed to create email verification token"): await create_email_verification_token("user-1", "a@b.c", "email_change") @@ -107,9 +103,7 @@ async def test_get_token_missing_credentials(self): @pytest.mark.asyncio async def test_get_token_returns_credentials(self): - credentials = HTTPAuthorizationCredentials( - scheme="Bearer", credentials="raw-token" - ) + credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="raw-token") assert await get_token(credentials=credentials) == "raw-token" @@ -117,9 +111,7 @@ class TestVerifySession: @pytest.mark.asyncio async def test_verify_session_success(self): redis_client = AsyncMock() - redis_client.get.return_value = str( - {"access_token": "tok", "user_id": "u1"} - ) + redis_client.get.return_value = str({"access_token": "tok", "user_id": "u1"}) data = await verify_session("sid-1", "tok", redis_client) assert data["access_token"] == "tok" redis_client.get.assert_awaited_once_with("session:sid-1") @@ -154,15 +146,11 @@ async def test_verify_session_token_mismatch(self): class TestVerifyToken: @pytest.mark.asyncio - async def test_verify_token_success( - self, test_db_session: AsyncSession, test_user: Users - ): + async def test_verify_token_success(self, test_db_session: AsyncSession, test_user: Users): token = await create_access_token({"sub": test_user.id, "sid": "sess-1"}) redis_client = AsyncMock() redis_client.get.return_value = str({"access_token": token}) - payload = await verify_token( - token=token, redis_client=redis_client, db=test_db_session - ) + payload = await verify_token(token=token, redis_client=redis_client, db=test_db_session) assert payload["sub"] == test_user.id assert payload["sid"] == "sess-1" @@ -170,9 +158,7 @@ async def test_verify_token_success( async def test_verify_token_missing_session_id(self, test_db_session: AsyncSession): token = await create_access_token({"sub": "user-1"}) with pytest.raises(HTTPException) as exc: - await verify_token( - token=token, redis_client=AsyncMock(), db=test_db_session - ) + await verify_token(token=token, redis_client=AsyncMock(), db=test_db_session) assert exc.value.status_code == 401 @pytest.mark.asyncio @@ -185,17 +171,13 @@ async def test_verify_token_disabled_account( redis_client = AsyncMock() redis_client.get.return_value = str({"access_token": token}) with pytest.raises(HTTPException) as exc: - await verify_token( - token=token, redis_client=redis_client, db=test_db_session - ) + await verify_token(token=token, redis_client=redis_client, db=test_db_session) assert exc.value.status_code == 403 @pytest.mark.asyncio async def test_verify_token_invalid_jwt(self, test_db_session: AsyncSession): with pytest.raises(HTTPException) as exc: - await verify_token( - token="not-a-jwt", redis_client=AsyncMock(), db=test_db_session - ) + await verify_token(token="not-a-jwt", redis_client=AsyncMock(), db=test_db_session) assert exc.value.status_code == 401 @@ -271,9 +253,7 @@ async def test_invalid_jwt(self): class TestVerifyEmailVerificationToken: @pytest.mark.asyncio async def test_valid_token(self): - token = await create_email_verification_token( - "user-1", "a@b.c", "email_change" - ) + token = await create_email_verification_token("user-1", "a@b.c", "email_change") result = await verify_email_verification_token(token=token) assert result["verification_type"] == "email_change" assert result["sub"] == "user-1" @@ -365,9 +345,7 @@ async def test_clear_user_all_sessions( test_user_session: UserSessions, ): redis_client = AsyncMock() - result = await clear_user_all_sessions( - test_db_session, redis_client, test_user.id - ) + result = await clear_user_all_sessions(test_db_session, redis_client, test_user.id) assert result is True redis_client.delete.assert_awaited() keys = redis_client.delete.await_args.args diff --git a/backend/tests/core/test_telemetry.py b/backend/tests/core/test_telemetry.py index 6b7f08b..104dd3b 100644 --- a/backend/tests/core/test_telemetry.py +++ b/backend/tests/core/test_telemetry.py @@ -74,9 +74,7 @@ async def test_skip_paths_are_not_traced_and_api_log_has_trace_id(): stream = io.StringIO() handler = logging.StreamHandler(stream) - handler.setFormatter( - TraceIdFormatter("%(message)s traceID=%(otelTraceID)s") - ) + handler.setFormatter(TraceIdFormatter("%(message)s traceID=%(otelTraceID)s")) api_logger = logging.getLogger("api_logger") api_logger.addHandler(handler) api_logger.setLevel(logging.INFO) @@ -217,26 +215,28 @@ def test_uses_project_settings(self): class TestSetupAndShutdownTelemetry: def test_setup_returns_when_disabled(self): test_app = FastAPI() - with patch("core.telemetry.settings.OTEL_ENABLE", False), patch( - "core.telemetry.LoggingInstrumentor" - ) as logging_instrumentor, patch( - "core.telemetry.FastAPIInstrumentor.instrument_app" - ) as instrument_app: + with ( + patch("core.telemetry.settings.OTEL_ENABLE", False), + patch("core.telemetry.LoggingInstrumentor") as logging_instrumentor, + patch("core.telemetry.FastAPIInstrumentor.instrument_app") as instrument_app, + ): setup_telemetry(test_app) logging_instrumentor.return_value.instrument.assert_called_once() instrument_app.assert_not_called() def test_shutdown_returns_when_disabled(self): - with patch("core.telemetry.settings.OTEL_ENABLE", False), patch( - "core.telemetry.trace.get_tracer_provider" - ) as get_provider: + with ( + patch("core.telemetry.settings.OTEL_ENABLE", False), + patch("core.telemetry.trace.get_tracer_provider") as get_provider, + ): shutdown_telemetry() get_provider.assert_not_called() def test_shutdown_calls_provider_shutdown(self): provider = MagicMock() - with patch("core.telemetry.settings.OTEL_ENABLE", True), patch( - "core.telemetry.trace.get_tracer_provider", return_value=provider + with ( + patch("core.telemetry.settings.OTEL_ENABLE", True), + patch("core.telemetry.trace.get_tracer_provider", return_value=provider), ): shutdown_telemetry() provider.shutdown.assert_called_once() diff --git a/backend/tests/extensions/test_smtp.py b/backend/tests/extensions/test_smtp.py index dd66386..c888f4d 100644 --- a/backend/tests/extensions/test_smtp.py +++ b/backend/tests/extensions/test_smtp.py @@ -1,15 +1,17 @@ -import pytest from unittest.mock import MagicMock, patch + +import pytest from fastapi import FastAPI + import extensions.smtp -from utils.custom_exception import SMTPNotConfiguredException from extensions.smtp import ( - SMTPSettings, SMTPMailer, + SMTPSettings, + add_smtp, build_smtp_settings, get_mailer, - add_smtp, ) +from utils.custom_exception import SMTPNotConfiguredException class TestSMTPSettings: diff --git a/backend/tests/utils/test_response.py b/backend/tests/utils/test_response.py index 5f8cd5a..7942015 100644 --- a/backend/tests/utils/test_response.py +++ b/backend/tests/utils/test_response.py @@ -4,7 +4,6 @@ from utils.response import ( generate_example_from_schema, - generate_property_example, is_openapi_examples, make_error_examples, make_response_doc, diff --git a/backend/utils/__init__.py b/backend/utils/__init__.py index 41f2565..741179b 100644 --- a/backend/utils/__init__.py +++ b/backend/utils/__init__.py @@ -1,3 +1,5 @@ # Add new utils imports below. from .get_real_ip import get_real_ip -from .response import APIResponse, parse_responses, common_responses \ No newline at end of file +from .response import APIResponse, common_responses, parse_responses + +__all__ = ["APIResponse", "common_responses", "get_real_ip", "parse_responses"] diff --git a/backend/utils/custom_exception.py b/backend/utils/custom_exception.py index 92695eb..d941575 100644 --- a/backend/utils/custom_exception.py +++ b/backend/utils/custom_exception.py @@ -1,16 +1,17 @@ import logging -from typing import Optional, Dict, Any +from typing import Any logger = logging.getLogger("exception") + class BaseServiceException(Exception): def __init__( - self, - message: str, + self, + message: str, error_code: str = None, - details: Optional[Dict[str, Any]] = None, + details: dict[str, Any] | None = None, status_code: int = None, - log_level: str = None + log_level: str = None, ): self.message = message self.error_code = error_code @@ -28,57 +29,106 @@ def __init__( super().__init__(self.message) + class ServerException(BaseServiceException): """Server exception""" - def __init__(self, message: str = "Server error", status_code: int = 500, details: Dict[str, Any] = None): - super().__init__(message=message, error_code="SERVER_ERROR", details=details, status_code=status_code, log_level="error") + + def __init__( + self, message: str = "Server error", status_code: int = 500, details: dict[str, Any] = None + ): + super().__init__( + message=message, + error_code="SERVER_ERROR", + details=details, + status_code=status_code, + log_level="error", + ) + class AuthenticationException(BaseServiceException): """Authentication related exceptions""" - def __init__(self, message: str = "Authentication failed", details: Dict[str, Any] = None): + + def __init__(self, message: str = "Authentication failed", details: dict[str, Any] = None): super().__init__(message=message, error_code="AUTH_ERROR", details=details, status_code=401) + class PasswordResetRequiredException(BaseServiceException): """Password reset required exception""" - def __init__(self, message: str = "Password reset required", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="PASSWORD_RESET_REQUIRED", details=details, status_code=202) + + def __init__(self, message: str = "Password reset required", details: dict[str, Any] = None): + super().__init__( + message=message, error_code="PASSWORD_RESET_REQUIRED", details=details, status_code=202 + ) + class EmailVerificationRequiredException(BaseServiceException): """Email verification required exception""" - def __init__(self, message: str = "Email verification required", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="EMAIL_VERIFICATION_REQUIRED", details=details, status_code=202) + + def __init__( + self, message: str = "Email verification required", details: dict[str, Any] = None + ): + super().__init__( + message=message, + error_code="EMAIL_VERIFICATION_REQUIRED", + details=details, + status_code=202, + ) + class AuthorizationException(BaseServiceException): """Authorization related exceptions""" - def __init__(self, message: str = "Permission denied", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="PERMISSION_ERROR", details=details, status_code=403) + + def __init__(self, message: str = "Permission denied", details: dict[str, Any] = None): + super().__init__( + message=message, error_code="PERMISSION_ERROR", details=details, status_code=403 + ) + class ValidationException(BaseServiceException): """Validation related exceptions""" - def __init__(self, message: str = "Validation failed", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="VALIDATION_ERROR", details=details, status_code=400) + + def __init__(self, message: str = "Validation failed", details: dict[str, Any] = None): + super().__init__( + message=message, error_code="VALIDATION_ERROR", details=details, status_code=400 + ) + class NotFoundException(BaseServiceException): """Resource not found exceptions""" - def __init__(self, message: str = "Resource not found", details: Dict[str, Any] = None): + + def __init__(self, message: str = "Resource not found", details: dict[str, Any] = None): super().__init__(message=message, error_code="NOT_FOUND", details=details, status_code=404) + class ConflictException(BaseServiceException): """Resource conflict exceptions""" - def __init__(self, message: str = "Resource conflict", details: Dict[str, Any] = None): + + def __init__(self, message: str = "Resource conflict", details: dict[str, Any] = None): super().__init__(message=message, error_code="CONFLICT", details=details, status_code=409) + class TokenException(BaseServiceException): """Token related exceptions""" - def __init__(self, message: str = "Token error", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="TOKEN_ERROR", details=details, status_code=401) + + def __init__(self, message: str = "Token error", details: dict[str, Any] = None): + super().__init__( + message=message, error_code="TOKEN_ERROR", details=details, status_code=401 + ) + class SMTPNotConfiguredException(BaseServiceException): """SMTP configuration related exceptions""" - def __init__(self, message: str = "SMTP is not configured", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="SMTP_NOT_CONFIGURED", details=details, status_code=503) + + def __init__(self, message: str = "SMTP is not configured", details: dict[str, Any] = None): + super().__init__( + message=message, error_code="SMTP_NOT_CONFIGURED", details=details, status_code=503 + ) + class RegistrationDisabledException(BaseServiceException): """Registration disabled exception""" - def __init__(self, message: str = "Registration is disabled", details: Dict[str, Any] = None): - super().__init__(message=message, error_code="REGISTRATION_DISABLED", details=details, status_code=503) + + def __init__(self, message: str = "Registration is disabled", details: dict[str, Any] = None): + super().__init__( + message=message, error_code="REGISTRATION_DISABLED", details=details, status_code=503 + ) diff --git a/backend/utils/email_templates.py b/backend/utils/email_templates.py index bc0857f..be44dfc 100644 --- a/backend/utils/email_templates.py +++ b/backend/utils/email_templates.py @@ -1,25 +1,21 @@ -from typing import Dict, Any, Optional +from typing import Any + class EmailTemplate: """Email template with subject, plain text body, and optional HTML body""" - - def __init__( - self, - subject: str, - body: str, - html_body: Optional[str] = None - ): + + def __init__(self, subject: str, body: str, html_body: str | None = None): self.subject = subject self.body = body self.html_body = html_body - - def render(self, **kwargs: Any) -> Dict[str, str]: + + def render(self, **kwargs: Any) -> dict[str, str]: """ Render template with variables. - + Args: **kwargs: Variables to substitute in template - + Returns: Dict with 'subject', 'body', and optionally 'html_body' keys """ @@ -27,10 +23,10 @@ def render(self, **kwargs: Any) -> Dict[str, str]: "subject": self.subject.format(**kwargs), "body": self.body.format(**kwargs), } - + if self.html_body: result["html_body"] = self.html_body.format(**kwargs) - + return result @@ -48,28 +44,28 @@ def render(self, **kwargs: Any) -> Dict[str, str]: "" "" "
" - "" - "" + "" + "" "" - "" - ""
- "You requested a password reset for your {app_name} account.
Click the button below to set a new password. This link will expire in 30 minutes."
- "
" - "" - "Reset Password" - "" - "
" - "" - "If you did not request this password reset, you can safely ignore this email. Your account remains secure." - "
" - ""
+ "You requested a password reset for your {app_name} account.
Click the button below to set a new password. This link will expire in 30 minutes."
+ "
" + "" + "Reset Password" + "" + "
" + "" + "If you did not request this password reset, you can safely ignore this email. Your account remains secure." + "
" + ""
- "Please verify your email address for your {app_name} account.
Please click the button below to verify your email address. This link will expire in {expire_minutes} minutes."
- "
" - "" - "Verify Email" - "" - "
" - "" - "If you did not request this email, you can safely ignore it." - "
" - ""
+ "Please verify your email address for your {app_name} account.
Please click the button below to verify your email address. This link will expire in {expire_minutes} minutes."
+ "
" + "" + "Verify Email" + "" + "
" + "" + "If you did not request this email, you can safely ignore it." + "
" + "