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]: "" "" "" - "" - "" + "" + "" "" - "" - "
" - "

Hi {user_name},

" - "

" - "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." - "

" - "
" + '' + "
" + "

Hi {user_name},

" + "

" + "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." + "

" + "
" "" "" ), @@ -89,29 +85,29 @@ def render(self, **kwargs: Any) -> Dict[str, str]: "" "" "" - "" - "" + "" + "" "" - "" - "
" - "

Hi {user_name},

" - "

" - "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." - "

" - "
" + '' + "
" + "

Hi {user_name},

" + "

" + "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." + "

" + "
" "" "" ), -) \ No newline at end of file +) diff --git a/backend/utils/get_real_ip.py b/backend/utils/get_real_ip.py index 59346fa..6b6723c 100644 --- a/backend/utils/get_real_ip.py +++ b/backend/utils/get_real_ip.py @@ -1,5 +1,6 @@ from fastapi import Request + def get_real_ip(request: Request) -> str: """ Get the client IP address from the X-Real-IP header set by nginx. diff --git a/backend/utils/response.py b/backend/utils/response.py index 91763de..c5e6008 100644 --- a/backend/utils/response.py +++ b/backend/utils/response.py @@ -1,19 +1,25 @@ from datetime import datetime +from typing import TypeVar + from pydantic import BaseModel, RootModel -from typing import Optional, Type, TypeVar, Generic T = TypeVar("T") -class APIResponse(BaseModel, Generic[T]): + +class APIResponse[T](BaseModel): """Generic API response wrapper that maintains consistent response structure""" + code: int message: str - data: Optional[T] = None + data: T | None = None + class ValidationErrorData(RootModel[dict[str, str]]): """Model for validation error data structure""" + pass + def is_openapi_examples(example: dict) -> bool: """ Detect OpenAPI Media Type ``examples`` (named map with dropdown in Swagger UI). @@ -28,26 +34,22 @@ def is_openapi_examples(example: dict) -> bool: return False if "code" in example or "message" in example or "data" in example: return False - return all( - isinstance(item, dict) and "value" in item - for item in example.values() - ) + return all(isinstance(item, dict) and "value" in item for item in example.values()) -def make_response_doc(description: str, model: Optional[Type] = None, example: Optional[dict] = None) -> dict: +def make_response_doc( + description: str, model: type | None = None, example: dict | None = None +) -> dict: """Create OpenAPI response documentation with model and example(s)""" doc = {"description": description} if model: doc["model"] = APIResponse[model] if example: - media = ( - {"examples": example} - if is_openapi_examples(example) - else {"example": example} - ) + media = {"examples": example} if is_openapi_examples(example) else {"example": example} doc["content"] = {"application/json": media} return doc + def make_error_examples(code: int, cases: dict[str, str]) -> dict: """ Build OpenAPI named ``examples`` for simple error responses. @@ -93,9 +95,9 @@ def parse_responses(custom: dict, default: dict = None) -> dict: try: schema = model.model_json_schema() data_example = generate_example_from_schema(schema) - except: + except Exception: data_example = None - + example = {"code": code, "message": desc, "data": data_example} result[code] = make_response_doc(desc, model, example) elif len(val) == 3: @@ -114,6 +116,7 @@ def parse_responses(custom: dict, default: dict = None) -> dict: result[code] = val return result + def generate_example_from_schema(schema: dict) -> dict: """Generate example data from JSON schema object properties""" if schema.get("type") == "object": @@ -124,6 +127,7 @@ def generate_example_from_schema(schema: dict) -> dict: return example return None + def generate_property_example(prop: dict, key: str = "", full_schema: dict = None): """Generate example value for a single property based on its type and field name""" # Handle $ref references first (for nested objects) @@ -132,9 +136,9 @@ def generate_property_example(prop: dict, key: str = "", full_schema: dict = Non if referenced_schema: return generate_example_from_schema(referenced_schema) return None - + prop_type = prop.get("type") - + if prop_type == "string": if key == "id": return "123e4567-e89b-12d3-a456-426614174000" @@ -151,7 +155,7 @@ def generate_property_example(prop: dict, key: str = "", full_schema: dict = Non else: return f"Example {key.replace('_', ' ').title()}" elif prop_type == "integer": - if key in ["per_page", "total_pages"]: + if key in ["per_page", "total_pages"]: return 10 elif key == "page": return 1 @@ -183,17 +187,18 @@ def generate_property_example(prop: dict, key: str = "", full_schema: dict = Non else: return None + def resolve_ref(ref_path: str, schema: dict) -> dict: """ Resolve JSON Schema $ref references to actual schema definitions """ if not ref_path.startswith("#/"): return None - + # Parse reference path: "#/$defs/UserRead" -> ["$defs", "UserRead"] path_parts = ref_path[2:].split("/") current = schema - + # Navigate through nested dict structure following the path for part in path_parts: if isinstance(current, dict) and part in current: @@ -201,56 +206,37 @@ def resolve_ref(ref_path: str, schema: dict) -> dict: else: # Reference not found return None - + if isinstance(current, dict): return current else: return None + common_responses = { 401: ( "Invalid or expired token", APIResponse[None], - { - "code": 401, - "message": "Invalid or expired token", - "data": None - } + {"code": 401, "message": "Invalid or expired token", "data": None}, ), 403: ( "Permission denied", APIResponse[None], - { - "code": 403, - "message": "Permission denied", - "data": None - } + {"code": 403, "message": "Permission denied", "data": None}, ), 422: ( "Validation Error", APIResponse[ValidationErrorData], - { - "code": 422, - "message": "Validation Error", - "data": {"body.params": "field required"} - } + {"code": 422, "message": "Validation Error", "data": {"body.params": "field required"}}, ), 429: ( "Too many requests. Try again later.", APIResponse[None], - { - "code": 429, - "message": "Too many requests. Try again later.", - "data": None - } + {"code": 429, "message": "Too many requests. Try again later.", "data": None}, ), 500: ( "Internal Server Error", APIResponse[None], - { - "code": 500, - "message": "Internal Server Error", - "data": None - } - ) -} \ No newline at end of file + {"code": 500, "message": "Internal Server Error", "data": None}, + ), +} diff --git a/frontend/.prettierignore b/frontend/.prettierignore new file mode 100644 index 0000000..91a3983 --- /dev/null +++ b/frontend/.prettierignore @@ -0,0 +1,3 @@ +dist +node_modules +package-lock.json diff --git a/frontend/README.md b/frontend/README.md index b1f5e88..8537127 100755 --- a/frontend/README.md +++ b/frontend/README.md @@ -8,6 +8,7 @@ This frontend project is built with modern web technologies to provide a fast, m - **Vite**: A lightning-fast frontend build tool and development server, enabling instant HMR and optimized production builds. - **Tailwind CSS**: A utility-first CSS framework for rapid UI development with a modern, responsive design. - **ESLint**: A pluggable JavaScript linter to maintain code quality and consistency. +- **Prettier**: An opinionated code formatter for consistent layout across the project. ## Features @@ -15,18 +16,46 @@ This frontend project is built with modern web technologies to provide a fast, m - 🎨 Modern, fully responsive UI styled with Tailwind CSS - 🧩 Modular, component-based architecture using React - 🛡️ Code quality enforced by ESLint +- ✨ Consistent formatting with Prettier -## Lint +## Lint & format -Requires `npm install` in this directory. +### Standards + +| Item | Value | +|------|--------| +| Lint config | [`eslint.config.js`](./eslint.config.js) | +| Format config | [`prettier.config.js`](./prettier.config.js) | +| Linter | [ESLint](https://eslint.org/) 9 (flat config) | +| Formatter | [Prettier](https://prettier.io/) 3 | +| Line length | 100 | +| Indent | 2 spaces | +| Quotes | Double quotes | + +### Manual commands + +> Run from the `frontend/` directory. Requires `npm install` in this directory. + +Lint the project; report issues without changing files: ```bash -cd frontend npm run lint ``` -Auto-fix: +Check formatting only; report mismatches without writing: + +```bash +npm run format:check +``` + +Lint and auto-fix what ESLint can: ```bash npm run lint:fix ``` + +Apply formatting to all supported files: + +```bash +npm run format +``` diff --git a/frontend/eslint.config.js b/frontend/eslint.config.js index 8b39716..08a2181 100755 --- a/frontend/eslint.config.js +++ b/frontend/eslint.config.js @@ -1,41 +1,47 @@ -import js from '@eslint/js' -import globals from 'globals' -import reactHooks from 'eslint-plugin-react-hooks' -import reactRefresh from 'eslint-plugin-react-refresh' +import js from "@eslint/js"; +import eslintConfigPrettier from "eslint-config-prettier"; +import globals from "globals"; +import reactHooks from "eslint-plugin-react-hooks"; +import reactRefresh from "eslint-plugin-react-refresh"; export default [ - { ignores: ['dist'] }, + { ignores: ["dist"] }, { - files: ['vite.config.js', 'script/**/*.js'], + files: ["vite.config.js", "script/**/*.js", "src/config/env.node.js"], languageOptions: { - ecmaVersion: 'latest', + ecmaVersion: "latest", globals: globals.node, - sourceType: 'module', + sourceType: "module", }, }, { - files: ['**/*.{js,jsx}'], + files: ["**/*.{js,jsx}"], languageOptions: { ecmaVersion: 2020, globals: globals.browser, parserOptions: { - ecmaVersion: 'latest', + ecmaVersion: "latest", ecmaFeatures: { jsx: true }, - sourceType: 'module', + sourceType: "module", }, }, plugins: { - 'react-hooks': reactHooks, - 'react-refresh': reactRefresh, + "react-hooks": reactHooks, + "react-refresh": reactRefresh, }, rules: { ...js.configs.recommended.rules, ...reactHooks.configs.recommended.rules, - 'no-unused-vars': ['error', { varsIgnorePattern: '^[A-Z_]' }], - 'react-refresh/only-export-components': [ - 'warn', - { allowConstantExport: true }, + "no-unused-vars": [ + "error", + { + varsIgnorePattern: "^[A-Z_]|^motion$", + argsIgnorePattern: "^_|^[A-Z]", + caughtErrorsIgnorePattern: "^_", + }, ], + "react-refresh/only-export-components": ["warn", { allowConstantExport: true }], }, }, -] + eslintConfigPrettier, +]; diff --git a/frontend/index.html b/frontend/index.html index fb046c3..0feb246 100755 --- a/frontend/index.html +++ b/frontend/index.html @@ -6,11 +6,14 @@ - + %VITE_PROJECT_NAME%
- \ No newline at end of file + diff --git a/frontend/jsconfig.json b/frontend/jsconfig.json index 7d585e8..5890108 100644 --- a/frontend/jsconfig.json +++ b/frontend/jsconfig.json @@ -10,11 +10,10 @@ "noEmit": true, "jsx": "react-jsx", "allowSyntheticDefaultImports": true, - "baseUrl": ".", "paths": { "@/*": ["./src/*"] } }, "include": ["src/**/*", "vite.config.js"], "exclude": ["node_modules", "dist"] -} \ No newline at end of file +} diff --git a/frontend/package-lock.json b/frontend/package-lock.json index 0c7a7e7..37e4d59 100644 --- a/frontend/package-lock.json +++ b/frontend/package-lock.json @@ -62,10 +62,12 @@ "@vitejs/plugin-react": "^4.4.1", "autoprefixer": "^10.4.21", "eslint": "^9.25.0", + "eslint-config-prettier": "^10.1.8", "eslint-plugin-react-hooks": "^5.2.0", "eslint-plugin-react-refresh": "^0.4.19", "globals": "^16.0.0", "postcss": "^8.5.4", + "prettier": "^3.9.6", "sitemap": "^8.0.0", "tailwindcss": "^4.1.8", "tw-animate-css": "^1.3.4", @@ -4414,6 +4416,22 @@ } } }, + "node_modules/eslint-config-prettier": { + "version": "10.1.8", + "resolved": "https://registry.npmjs.org/eslint-config-prettier/-/eslint-config-prettier-10.1.8.tgz", + "integrity": "sha512-82GZUjRS0p/jganf6q1rEO25VSoHH0hKPCTrgillPjdI/3bgBhAE1QzHrHTizjpRvy6pGAvKjDJtk2pF9NDq8w==", + "dev": true, + "license": "MIT", + "bin": { + "eslint-config-prettier": "bin/cli.js" + }, + "funding": { + "url": "https://opencollective.com/eslint-config-prettier" + }, + "peerDependencies": { + "eslint": ">=7.0.0" + } + }, "node_modules/eslint-plugin-react-hooks": { "version": "5.2.0", "resolved": "https://registry.npmjs.org/eslint-plugin-react-hooks/-/eslint-plugin-react-hooks-5.2.0.tgz", @@ -5684,6 +5702,22 @@ "node": ">= 0.8.0" } }, + "node_modules/prettier": { + "version": "3.9.6", + "resolved": "https://registry.npmjs.org/prettier/-/prettier-3.9.6.tgz", + "integrity": "sha512-OpN0zzVdiaiAhxpuuj5efpIS4sY9j7bY6uR5mnj5yPzGkdkjNKSJeUThPb60Jw29QuAZgA4o+/iB49kFiaBX6g==", + "dev": true, + "license": "MIT", + "bin": { + "prettier": "bin/prettier.cjs" + }, + "engines": { + "node": ">=14" + }, + "funding": { + "url": "https://github.com/prettier/prettier?sponsor=1" + } + }, "node_modules/proxy-from-env": { "version": "2.1.0", "resolved": "https://registry.npmjs.org/proxy-from-env/-/proxy-from-env-2.1.0.tgz", diff --git a/frontend/package.json b/frontend/package.json index a3b2727..76cbd74 100755 --- a/frontend/package.json +++ b/frontend/package.json @@ -8,6 +8,9 @@ "build": "vite build", "lint": "eslint .", "lint:fix": "eslint . --fix", + "format": "prettier --write .", + "format:check": "prettier --check .", + "fix": "npm run format && npm run lint:fix", "preview": "vite preview --port 3000", "generate-sitemap": "node script/generate-sitemap.js" }, @@ -66,10 +69,12 @@ "@vitejs/plugin-react": "^4.4.1", "autoprefixer": "^10.4.21", "eslint": "^9.25.0", + "eslint-config-prettier": "^10.1.8", "eslint-plugin-react-hooks": "^5.2.0", "eslint-plugin-react-refresh": "^0.4.19", "globals": "^16.0.0", "postcss": "^8.5.4", + "prettier": "^3.9.6", "sitemap": "^8.0.0", "tailwindcss": "^4.1.8", "tw-animate-css": "^1.3.4", diff --git a/frontend/prettier.config.js b/frontend/prettier.config.js new file mode 100644 index 0000000..bc1acef --- /dev/null +++ b/frontend/prettier.config.js @@ -0,0 +1,8 @@ +/** @type {import('prettier').Config} */ +export default { + semi: true, + singleQuote: false, + tabWidth: 2, + trailingComma: "es5", + printWidth: 100, +}; diff --git a/frontend/script/generate-sitemap.js b/frontend/script/generate-sitemap.js index be59fad..7f0c957 100755 --- a/frontend/script/generate-sitemap.js +++ b/frontend/script/generate-sitemap.js @@ -1,15 +1,15 @@ -import { createWriteStream } from 'fs'; -import { routes } from '../src/router/routes.js'; -import { ENV_NODE } from '../src/config/env.node.js'; -import { SitemapStream, streamToPromise } from 'sitemap'; +import { createWriteStream } from "fs"; +import { routes } from "../src/router/routes.js"; +import { ENV_NODE } from "../src/config/env.node.js"; +import { SitemapStream, streamToPromise } from "sitemap"; -const sitemap = new SitemapStream({ hostname: ENV_NODE.SITE_URL || 'http://localhost' }); +const sitemap = new SitemapStream({ hostname: ENV_NODE.SITE_URL || "http://localhost" }); -routes.forEach(route => { - sitemap.write({ url: route.path, changefreq: 'weekly', priority: 0.8 }); +routes.forEach((route) => { + sitemap.write({ url: route.path, changefreq: "weekly", priority: 0.8 }); }); sitemap.end(); -streamToPromise(sitemap).then(data => { - createWriteStream('public/sitemap.xml').end(data); -}); \ No newline at end of file +streamToPromise(sitemap).then((data) => { + createWriteStream("public/sitemap.xml").end(data); +}); diff --git a/frontend/src/components/animate-ui/icons/arrow-left.jsx b/frontend/src/components/animate-ui/icons/arrow-left.jsx index 71a8368..693605b 100644 --- a/frontend/src/components/animate-ui/icons/arrow-left.jsx +++ b/frontend/src/components/animate-ui/icons/arrow-left.jsx @@ -1,37 +1,41 @@ -import * as React from 'react'; -import { motion } from 'motion/react'; -import { getVariants, useAnimateIconContext, IconWrapper } from '@/components/animate-ui/icons/icon'; +import * as React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const animations = { default: { group: { initial: { x: 0, - transition: { ease: 'easeInOut', duration: 0.3 }, + transition: { ease: "easeInOut", duration: 0.3 }, }, animate: { - x: '-25%', - transition: { ease: 'easeInOut', duration: 0.3 }, + x: "-25%", + transition: { ease: "easeInOut", duration: 0.3 }, }, }, path1: {}, - path2: {} + path2: {}, }, - 'default-loop': { + "default-loop": { group: { initial: { x: 0, }, animate: { - x: [0, '-25%', 0], - transition: { ease: 'easeInOut', duration: 0.6 }, + x: [0, "-25%", 0], + transition: { ease: "easeInOut", duration: 0.6 }, }, }, path1: {}, - path2: {} + path2: {}, }, pointing: { @@ -39,49 +43,49 @@ const animations = { path1: { initial: { - d: 'M19 12H5', - transition: { ease: 'easeInOut', duration: 0.3 }, + d: "M19 12H5", + transition: { ease: "easeInOut", duration: 0.3 }, }, animate: { - d: 'M19 12H10', - transition: { ease: 'easeInOut', duration: 0.3 }, + d: "M19 12H10", + transition: { ease: "easeInOut", duration: 0.3 }, }, }, path2: { initial: { - d: 'm12 19-7-7 7-7', - transition: { ease: 'easeInOut', duration: 0.3 }, + d: "m12 19-7-7 7-7", + transition: { ease: "easeInOut", duration: 0.3 }, }, animate: { - d: 'm15.5 19-7-7 7-7', - transition: { ease: 'easeInOut', duration: 0.3 }, + d: "m15.5 19-7-7 7-7", + transition: { ease: "easeInOut", duration: 0.3 }, }, - } + }, }, - 'pointing-loop': { + "pointing-loop": { group: {}, path1: { initial: { - d: 'M19 12H5', + d: "M19 12H5", }, animate: { - d: ['M19 12H5', 'M19 12H10', 'M19 12H5'], - transition: { ease: 'easeInOut', duration: 0.6 }, + d: ["M19 12H5", "M19 12H10", "M19 12H5"], + transition: { ease: "easeInOut", duration: 0.6 }, }, }, path2: { initial: { - d: 'm12 19-7-7 7-7', + d: "m12 19-7-7 7-7", }, animate: { - d: ['m12 19-7-7 7-7', 'm15.5 19-7-7 7-7', 'm12 19-7-7 7-7'], - transition: { ease: 'easeInOut', duration: 0.6 }, + d: ["m12 19-7-7 7-7", "m15.5 19-7-7 7-7", "m12 19-7-7 7-7"], + transition: { ease: "easeInOut", duration: 0.6 }, }, - } + }, }, out: { @@ -90,11 +94,11 @@ const animations = { x: 0, }, animate: { - x: [0, '-150%', '150%', 0], + x: [0, "-150%", "150%", 0], transition: { - default: { ease: 'easeInOut', duration: 0.6 }, + default: { ease: "easeInOut", duration: 0.6 }, x: { - ease: 'easeInOut', + ease: "easeInOut", duration: 0.6, times: [0, 0.5, 0.5, 1], }, @@ -103,14 +107,11 @@ const animations = { }, path1: {}, - path2: {} - } + path2: {}, + }, }; -function IconComponent({ - size, - ...props -}) { +function IconComponent({ size, ...props }) { const { controls } = useAnimateIconContext(); const variants = getVariants(animations); @@ -125,18 +126,16 @@ function IconComponent({ strokeWidth={2} strokeLinecap="round" strokeLinejoin="round" - {...props}> + {...props} + > - + + animate={controls} + /> ); diff --git a/frontend/src/components/animate-ui/icons/icon.jsx b/frontend/src/components/animate-ui/icons/icon.jsx index 09e85bd..f35af05 100644 --- a/frontend/src/components/animate-ui/icons/icon.jsx +++ b/frontend/src/components/animate-ui/icons/icon.jsx @@ -1,8 +1,8 @@ -import * as React from 'react'; -import { motion, useAnimation } from 'motion/react'; -import { cn } from '@/lib/utils'; -import { useIsInView } from '@/hooks/useIsInView'; -import { Slot } from '@/components/animate-ui/primitives/animate/slot'; +import * as React from "react"; +import { motion, useAnimation } from "motion/react"; +import { cn } from "@/lib/utils"; +import { useIsInView } from "@/hooks/useIsInView"; +import { Slot } from "@/components/animate-ui/primitives/animate/slot"; const staticAnimations = { path: { @@ -12,22 +12,22 @@ const staticAnimations = { pathLength: [0.05, 1], transition: { duration: 0.8, - ease: 'easeInOut', + ease: "easeInOut", }, - } + }, }, - 'path-loop': { + "path-loop": { initial: { pathLength: 1 }, animate: { pathLength: [1, 0.05, 1], transition: { duration: 1.6, - ease: 'easeInOut', + ease: "easeInOut", }, - } - } + }, + }, }; const AnimateIconContext = React.createContext(null); @@ -37,7 +37,7 @@ function useAnimateIconContext() { if (!context) return { controls: undefined, - animation: 'default', + animation: "default", loop: undefined, loopDelay: undefined, active: undefined, @@ -64,9 +64,9 @@ function AnimateIcon({ animateOnHover = false, animateOnTap = false, animateOnView = false, - animateOnViewMargin = '0px', + animateOnViewMargin = "0px", animateOnViewOnce = true, - animation = 'default', + animation = "default", loop = false, loopDelay = 0, initialOnAnimateEnd = false, @@ -82,8 +82,10 @@ function AnimateIcon({ if (animate === undefined || animate === false) return false; return delay <= 0; }); - const [currentAnimation, setCurrentAnimation] = React.useState(typeof animate === 'string' ? animate : animation); - const [status, setStatus] = React.useState('initial'); + const [currentAnimation, setCurrentAnimation] = React.useState( + typeof animate === "string" ? animate : animation + ); + const [status, setStatus] = React.useState("initial"); const delayRef = React.useRef(null); const loopDelayRef = React.useRef(null); @@ -99,23 +101,26 @@ function AnimateIcon({ runGenRef.current++; }, []); - const startAnimation = React.useCallback((trigger) => { - const next = typeof trigger === 'string' ? trigger : animation; - bumpGeneration(); - if (delayRef.current) { - clearTimeout(delayRef.current); - delayRef.current = null; - } - setCurrentAnimation(next); - if (delay > 0) { - setLocalAnimate(false); - delayRef.current = setTimeout(() => { + const startAnimation = React.useCallback( + (trigger) => { + const next = typeof trigger === "string" ? trigger : animation; + bumpGeneration(); + if (delayRef.current) { + clearTimeout(delayRef.current); + delayRef.current = null; + } + setCurrentAnimation(next); + if (delay > 0) { + setLocalAnimate(false); + delayRef.current = setTimeout(() => { + setLocalAnimate(true); + }, delay); + } else { setLocalAnimate(true); - }, delay); - } else { - setLocalAnimate(true); - } - }, [animation, delay, bumpGeneration]); + } + }, + [animation, delay, bumpGeneration] + ); const stopAnimation = React.useCallback(() => { bumpGeneration(); @@ -136,7 +141,7 @@ function AnimateIcon({ React.useEffect(() => { if (animate === undefined) return; - setCurrentAnimation(typeof animate === 'string' ? animate : animation); + setCurrentAnimation(typeof animate === "string" ? animate : animation); if (animate) startAnimation(animate); else stopAnimation(); // eslint-disable-next-line react-hooks/exhaustive-deps @@ -156,14 +161,17 @@ function AnimateIcon({ inViewMargin: animateOnViewMargin, }); - const startAnim = React.useCallback(async (anim, method = 'start') => { - try { - await controls[method](anim); - setStatus(anim); - } catch { - return; - } - }, [controls]); + const startAnim = React.useCallback( + async (anim, method = "start") => { + try { + await controls[method](anim); + setStatus(anim); + } catch { + return; + } + }, + [controls] + ); React.useEffect(() => { if (!animateOnView) return; @@ -177,16 +185,12 @@ function AnimateIcon({ async function run() { if (cancelledRef.current || gen !== runGenRef.current) { - await startAnim('initial'); + await startAnim("initial"); return; } if (!localAnimate) { - if ( - completeOnStop && - isAnimateInProgressRef.current && - animateEndPromiseRef.current - ) { + if (completeOnStop && isAnimateInProgressRef.current && animateEndPromiseRef.current) { try { await animateEndPromiseRef.current; } catch { @@ -195,20 +199,20 @@ function AnimateIcon({ } if (!persistOnAnimateEnd) { if (cancelledRef.current || gen !== runGenRef.current) { - await startAnim('initial'); + await startAnim("initial"); return; } - await startAnim('initial'); + await startAnim("initial"); } return; } if (loop) { if (cancelledRef.current || gen !== runGenRef.current) { - await startAnim('initial'); + await startAnim("initial"); return; } - await startAnim('initial', 'set'); + await startAnim("initial", "set"); } isAnimateInProgressRef.current = true; @@ -221,18 +225,18 @@ function AnimateIcon({ resolveAnimateEndRef.current?.(); resolveAnimateEndRef.current = null; animateEndPromiseRef.current = null; - await startAnim('initial'); + await startAnim("initial"); return; } - await startAnim('animate'); + await startAnim("animate"); if (cancelledRef.current || gen !== runGenRef.current) { isAnimateInProgressRef.current = false; resolveAnimateEndRef.current?.(); resolveAnimateEndRef.current = null; animateEndPromiseRef.current = null; - await startAnim('initial'); + await startAnim("initial"); return; } @@ -243,10 +247,10 @@ function AnimateIcon({ if (initialOnAnimateEnd) { if (cancelledRef.current || gen !== runGenRef.current) { - await startAnim('initial'); + await startAnim("initial"); return; } - await startAnim('initial', 'set'); + await startAnim("initial", "set"); } if (loop) { @@ -259,23 +263,21 @@ function AnimateIcon({ }); if (cancelledRef.current || gen !== runGenRef.current) { - await startAnim('initial'); + await startAnim("initial"); return; } if (!activeRef.current) { - if (status !== 'initial' && !persistOnAnimateEnd) - await startAnim('initial'); + if (status !== "initial" && !persistOnAnimateEnd) await startAnim("initial"); return; } } else { if (!activeRef.current) { - if (status !== 'initial' && !persistOnAnimateEnd) - await startAnim('initial'); + if (status !== "initial" && !persistOnAnimateEnd) await startAnim("initial"); return; } } if (cancelledRef.current || gen !== runGenRef.current) { - await startAnim('initial'); + await startAnim("initial"); return; } await run(); @@ -298,7 +300,7 @@ function AnimateIcon({ // eslint-disable-next-line react-hooks/exhaustive-deps }, [localAnimate, controls]); - const childProps = (React.isValidElement(children) ? (children).props : {}); + const childProps = React.isValidElement(children) ? children.props : {}; const handleMouseEnter = composeEventHandlers(childProps.onMouseEnter, () => { if (animateOnHover) startAnimation(animateOnHover); @@ -323,7 +325,8 @@ function AnimateIcon({ onMouseLeave={handleMouseLeave} onPointerDown={handlePointerDown} onPointerUp={handlePointerUp} - {...props}> + {...props} + > {children} ) : ( @@ -333,7 +336,8 @@ function AnimateIcon({ onMouseLeave={handleMouseLeave} onPointerDown={handlePointerDown} onPointerUp={handlePointerUp} - {...props}> + {...props} + > {children} ); @@ -357,40 +361,44 @@ function AnimateIcon({ value.speedMultiplier = inheritedSpeedMultiplier; } return value; - }, [controls, currentAnimation, loop, loopDelay, localAnimate, animate, initialOnAnimateEnd, completeOnStop, delay, inheritedSpeedMultiplier]); - - return ( - - {content} - - ); -} - -const pathClassName = - "[&_[stroke-dasharray='1px_1px']]:![stroke-dasharray:1px_0px]"; - -function IconWrapper( - { - size = 28, - animation: animationProp, - animate, - animateOnHover, - animateOnTap, - animateOnView, - animateOnViewMargin, - animateOnViewOnce, - icon: IconComponent, + }, [ + controls, + currentAnimation, loop, loopDelay, - persistOnAnimateEnd, + localAnimate, + animate, initialOnAnimateEnd, - delay, completeOnStop, - speedMultiplier = 0.7, - className, - ...props - } -) { + delay, + inheritedSpeedMultiplier, + ]); + + return {content}; +} + +const pathClassName = "[&_[stroke-dasharray='1px_1px']]:![stroke-dasharray:1px_0px]"; + +function IconWrapper({ + size = 28, + animation: animationProp, + animate, + animateOnHover, + animateOnTap, + animateOnView, + animateOnViewMargin, + animateOnViewOnce, + icon: IconComponent, + loop, + loopDelay, + persistOnAnimateEnd, + initialOnAnimateEnd, + delay, + completeOnStop, + speedMultiplier = 0.7, + className, + ...props +}) { const context = React.useContext(AnimateIconContext); if (context) { @@ -439,35 +447,36 @@ function IconWrapper( delay: parentDelay, completeOnStop: parentCompleteOnStop, speedMultiplier: finalSpeedMultiplier, - }}> + }} + > + {...props} + /> ); } if (hasAnimationOverrides) { const inheritedAnimate = parentActive - ? (animationProp ?? parentAnimation ?? 'default') + ? (animationProp ?? parentAnimation ?? "default") : false; - const finalAnimate = (animate ?? - parentAnimate ?? inheritedAnimate); + const finalAnimate = animate ?? parentAnimate ?? inheritedAnimate; const finalSpeedMultiplier = speedMultiplier ?? parentSpeedMultiplier ?? 0.7; - + return ( + }} + > + asChild + > + className={cn( + className, + ((animationProp ?? parentAnimation) === "path" || + (animationProp ?? parentAnimation) === "path-loop") && + pathClassName + )} + {...props} + /> ); @@ -512,15 +526,16 @@ function IconWrapper( delay: parentDelay, completeOnStop: parentCompleteOnStop, speedMultiplier: finalSpeedMultiplier, - }}> + }} + > + {...props} + /> ); } @@ -546,7 +561,8 @@ function IconWrapper( delay, completeOnStop, speedMultiplier: speedMultiplier ?? 0.7, - }}> + }} + > + asChild + > + className={cn( + className, + (animationProp === "path" || animationProp === "path-loop") && pathClassName + )} + {...props} + /> ); @@ -573,9 +593,12 @@ function IconWrapper( return ( + className={cn( + className, + (animationProp === "path" || animationProp === "path-loop") && pathClassName + )} + {...props} + /> ); } @@ -589,15 +612,12 @@ function getVariants(animations) { const variant = staticAnimations[animationType]; result = {}; for (const key in animations.default) { - if ( - (animationType === 'path' || animationType === 'path-loop') && - key.includes('group') - ) + if ((animationType === "path" || animationType === "path-loop") && key.includes("group")) continue; result[key] = variant; } } else { - result = (animations[animationType]) ?? animations.default; + result = animations[animationType] ?? animations.default; } if (speedMultiplier !== undefined && speedMultiplier !== 1) { @@ -608,51 +628,52 @@ function getVariants(animations) { } function adjustAnimationSpeed(animations, multiplier) { - if (!animations || typeof animations !== 'object') { + if (!animations || typeof animations !== "object") { return animations; } if (Array.isArray(animations)) { - return animations.map(item => adjustAnimationSpeed(item, multiplier)); + return animations.map((item) => adjustAnimationSpeed(item, multiplier)); } const adjusted = {}; for (const key in animations) { const anim = animations[key]; - + if (anim === null || anim === undefined) { adjusted[key] = anim; continue; } - if (typeof anim === 'object') { - if ('initial' in anim || 'animate' in anim) { + if (typeof anim === "object") { + if ("initial" in anim || "animate" in anim) { adjusted[key] = { ...anim, }; - - if (anim.initial && typeof anim.initial === 'object' && anim.initial.transition) { + + if (anim.initial && typeof anim.initial === "object" && anim.initial.transition) { adjusted[key].initial = { ...anim.initial, transition: adjustTransition(anim.initial.transition, multiplier), }; } - - if (anim.animate && typeof anim.animate === 'object') { + + if (anim.animate && typeof anim.animate === "object") { adjusted[key].animate = { ...anim.animate }; - + if (anim.animate.transition) { - adjusted[key].animate.transition = adjustTransition(anim.animate.transition, multiplier); + adjusted[key].animate.transition = adjustTransition( + anim.animate.transition, + multiplier + ); } } - } - else if (anim.transition) { + } else if (anim.transition) { adjusted[key] = { ...anim, transition: adjustTransition(anim.transition, multiplier), }; - } - else { + } else { adjusted[key] = adjustAnimationSpeed(anim, multiplier); } } else { @@ -664,12 +685,12 @@ function adjustAnimationSpeed(animations, multiplier) { function adjustTransition(transition, multiplier) { if (!transition) return transition; - if (typeof transition === 'object') { + if (typeof transition === "object") { const adjusted = { ...transition }; - if (typeof transition.duration === 'number') { + if (typeof transition.duration === "number") { adjusted.duration = transition.duration * multiplier; } - if (typeof transition.delay === 'number') { + if (typeof transition.delay === "number") { adjusted.delay = transition.delay * multiplier; } return adjusted; @@ -677,4 +698,11 @@ function adjustTransition(transition, multiplier) { return transition; } -export { pathClassName, staticAnimations, AnimateIcon, IconWrapper, useAnimateIconContext, getVariants }; +export { + pathClassName, + staticAnimations, + AnimateIcon, + IconWrapper, + useAnimateIconContext, + getVariants, +}; diff --git a/frontend/src/components/animate-ui/icons/languages.jsx b/frontend/src/components/animate-ui/icons/languages.jsx index a26405d..f913460 100644 --- a/frontend/src/components/animate-ui/icons/languages.jsx +++ b/frontend/src/components/animate-ui/icons/languages.jsx @@ -1,10 +1,14 @@ -import React from "react" -import { motion } from "motion/react" -import { getVariants, useAnimateIconContext, IconWrapper } from "@/components/animate-ui/icons/icon" +import React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const animations = { default: (() => { - const animation = {} + const animation = {}; for (let i = 0; i <= 5; i++) { animation[`path${i + 1}`] = { initial: { pathLength: 1, opacity: 1 }, @@ -17,12 +21,12 @@ const animations = { ease: [0.4, 0, 0.2, 1], }, }, - } + }; } - return animation + return animation; })(), "draw-stroke": (() => { - const animation = {} + const animation = {}; for (let i = 0; i <= 5; i++) { animation[`path${i + 1}`] = { initial: { pathLength: 1, opacity: 1 }, @@ -35,12 +39,12 @@ const animations = { ease: [0.4, 0, 0.2, 1], }, }, - } + }; } - return animation + return animation; })(), "scale-group": (() => { - const animation = {} + const animation = {}; animation.path1 = { initial: { scale: 1, opacity: 1 }, animate: { @@ -52,7 +56,7 @@ const animations = { ease: [0.34, 1.56, 0.64, 1], }, }, - } + }; animation.path2 = { initial: { scale: 1, opacity: 1 }, animate: { @@ -64,7 +68,7 @@ const animations = { ease: [0.34, 1.56, 0.64, 1], }, }, - } + }; animation.path3 = { initial: { scale: 1, opacity: 1 }, animate: { @@ -76,7 +80,7 @@ const animations = { ease: [0.34, 1.56, 0.64, 1], }, }, - } + }; animation.path4 = { initial: { scale: 1, opacity: 1 }, animate: { @@ -88,7 +92,7 @@ const animations = { ease: [0.34, 1.56, 0.64, 1], }, }, - } + }; animation.path5 = { initial: { scale: 1, opacity: 1 }, animate: { @@ -100,7 +104,7 @@ const animations = { ease: [0.34, 1.56, 0.64, 1], }, }, - } + }; animation.path6 = { initial: { scale: 1, opacity: 1 }, animate: { @@ -112,8 +116,8 @@ const animations = { ease: [0.34, 1.56, 0.64, 1], }, }, - } - return animation + }; + return animation; })(), bounce: { group: { @@ -133,12 +137,12 @@ const animations = { path5: {}, path6: {}, }, -} +}; function IconComponent({ size, ...props }) { - const { controls, animation: animationType } = useAnimateIconContext() - const variants = getVariants(animations) - const variant = animationType || "draw-stroke" + const { controls, animation: animationType } = useAnimateIconContext(); + const variants = getVariants(animations); + const variant = animationType || "draw-stroke"; if (variant === "draw-stroke") { return ( @@ -157,13 +161,23 @@ function IconComponent({ size, ...props }) { {...props} > - + - + - ) + ); } if (variant === "scale-group") { @@ -225,7 +239,7 @@ function IconComponent({ size, ...props }) { style={{ originX: 0.5, originY: 0.5 }} /> - ) + ); } if (variant === "bounce") { @@ -252,7 +266,7 @@ function IconComponent({ size, ...props }) { - ) + ); } return ( @@ -271,17 +285,27 @@ function IconComponent({ size, ...props }) { {...props} > - + - + - ) + ); } function LanguagesIcon(props) { - return + return ; } -export { animations, LanguagesIcon, LanguagesIcon as LanguagesIconIcon } +export { animations, LanguagesIcon, LanguagesIcon as LanguagesIconIcon }; diff --git a/frontend/src/components/animate-ui/icons/loader-circle.jsx b/frontend/src/components/animate-ui/icons/loader-circle.jsx index b646271..95aa28b 100644 --- a/frontend/src/components/animate-ui/icons/loader-circle.jsx +++ b/frontend/src/components/animate-ui/icons/loader-circle.jsx @@ -1,6 +1,10 @@ -import * as React from 'react'; -import { motion } from 'motion/react'; -import { getVariants, useAnimateIconContext, IconWrapper } from '@/components/animate-ui/icons/icon'; +import * as React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const animations = { default: { @@ -10,21 +14,18 @@ const animations = { rotate: 360, transition: { duration: 1, - ease: 'linear', + ease: "linear", repeat: Infinity, - repeatType: 'loop', + repeatType: "loop", }, }, }, - path: {} - } + path: {}, + }, }; -function IconComponent({ - size, - ...props -}) { +function IconComponent({ size, ...props }) { const { controls } = useAnimateIconContext(); const variants = getVariants(animations); @@ -42,12 +43,14 @@ function IconComponent({ variants={variants.group} initial="initial" animate={controls} - {...props}> + {...props} + > + animate={controls} + /> ); } diff --git a/frontend/src/components/animate-ui/icons/loader.jsx b/frontend/src/components/animate-ui/icons/loader.jsx index 8e61d22..53a53c3 100644 --- a/frontend/src/components/animate-ui/icons/loader.jsx +++ b/frontend/src/components/animate-ui/icons/loader.jsx @@ -1,6 +1,10 @@ -import * as React from 'react'; -import { motion } from 'motion/react'; -import { getVariants, useAnimateIconContext, IconWrapper } from '@/components/animate-ui/icons/icon'; +import * as React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const SEGMENT_COUNT = 8; const DURATION = 1.2; @@ -22,9 +26,9 @@ const animations = { opacity: [1, BASE_OPACITY], transition: { duration: DURATION, - ease: 'linear', + ease: "linear", repeat: Infinity, - repeatType: 'loop', + repeatType: "loop", delay, }, }, @@ -41,9 +45,9 @@ const animations = { rotate: 360, transition: { duration: 1.5, - ease: 'linear', + ease: "linear", repeat: Infinity, - repeatType: 'loop', + repeatType: "loop", }, }, }, @@ -55,14 +59,11 @@ const animations = { path5: {}, path6: {}, path7: {}, - path8: {} - } + path8: {}, + }, }; -function IconComponent({ - size, - ...props -}) { +function IconComponent({ size, ...props }) { const { controls } = useAnimateIconContext(); const variants = getVariants(animations); @@ -80,47 +81,36 @@ function IconComponent({ variants={variants.group} initial="initial" animate={controls} - {...props}> - + {...props} + > + - + animate={controls} + /> + - + animate={controls} + /> + - + animate={controls} + /> + + animate={controls} + /> ); } diff --git a/frontend/src/components/animate-ui/icons/moon.jsx b/frontend/src/components/animate-ui/icons/moon.jsx index ecd78ad..734c1ca 100644 --- a/frontend/src/components/animate-ui/icons/moon.jsx +++ b/frontend/src/components/animate-ui/icons/moon.jsx @@ -1,6 +1,10 @@ -import * as React from 'react'; -import { motion } from 'motion/react'; -import { getVariants, useAnimateIconContext, IconWrapper } from '@/components/animate-ui/icons/icon'; +import * as React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const animations = { default: { @@ -9,7 +13,7 @@ const animations = { rotate: 0, transition: { duration: 0.5, - ease: 'easeInOut', + ease: "easeInOut", }, }, animate: { @@ -17,10 +21,10 @@ const animations = { transition: { duration: 1.2, times: [0, 0.25, 0.75, 1], - ease: ['easeInOut', 'easeInOut', 'easeInOut'], + ease: ["easeInOut", "easeInOut", "easeInOut"], }, }, - } + }, }, balancing: { @@ -29,24 +33,21 @@ const animations = { rotate: 0, transition: { duration: 0.5, - ease: 'easeInOut', + ease: "easeInOut", }, }, animate: { rotate: [0, -30, 25, -15, 10, -5, 0], transition: { duration: 1.2, - ease: 'easeInOut', + ease: "easeInOut", }, }, - } - } + }, + }, }; -function IconComponent({ - size, - ...props -}) { +function IconComponent({ size, ...props }) { const { controls } = useAnimateIconContext(); const variants = getVariants(animations); @@ -63,12 +64,14 @@ function IconComponent({ strokeLinejoin="round" initial="initial" animate={controls} - {...props}> + {...props} + > + animate={controls} + /> ); } diff --git a/frontend/src/components/animate-ui/icons/sun-moon.jsx b/frontend/src/components/animate-ui/icons/sun-moon.jsx index fb28654..69dd362 100644 --- a/frontend/src/components/animate-ui/icons/sun-moon.jsx +++ b/frontend/src/components/animate-ui/icons/sun-moon.jsx @@ -1,6 +1,10 @@ -import * as React from 'react'; -import { motion } from 'motion/react'; -import { getVariants, useAnimateIconContext, IconWrapper } from '@/components/animate-ui/icons/icon'; +import * as React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const animations = { default: (() => { @@ -13,7 +17,7 @@ const animations = { rotate: [0, -10, 10, 0], transition: { duration: 0.6, - ease: 'easeInOut', + ease: "easeInOut", }, }, }, @@ -28,7 +32,7 @@ const animations = { pathLength: [0, 1], transition: { duration: 0.6, - ease: 'easeInOut', + ease: "easeInOut", delay: (i - 1) * 0.15, }, }, @@ -36,13 +40,10 @@ const animations = { } return animation; - })() + })(), }; -function IconComponent({ - size, - ...props -}) { +function IconComponent({ size, ...props }) { const { controls } = useAnimateIconContext(); const variants = getVariants(animations); @@ -59,17 +60,20 @@ function IconComponent({ strokeLinejoin="round" initial="initial" animate={controls} - {...props}> + {...props} + > + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> ); } diff --git a/frontend/src/components/animate-ui/icons/sun.jsx b/frontend/src/components/animate-ui/icons/sun.jsx index 51e61be..14ca017 100644 --- a/frontend/src/components/animate-ui/icons/sun.jsx +++ b/frontend/src/components/animate-ui/icons/sun.jsx @@ -1,6 +1,10 @@ -import * as React from 'react'; -import { motion } from 'motion/react'; -import { getVariants, useAnimateIconContext, IconWrapper } from '@/components/animate-ui/icons/icon'; +import * as React from "react"; +import { motion } from "motion/react"; +import { + getVariants, + useAnimateIconContext, + IconWrapper, +} from "@/components/animate-ui/icons/icon"; const animations = { default: (() => { @@ -16,7 +20,7 @@ const animations = { pathLength: [0, 1], transition: { duration: 0.6, - ease: 'easeInOut', + ease: "easeInOut", delay: (i - 1) * 0.15, }, }, @@ -24,13 +28,10 @@ const animations = { } return animation; - })() + })(), }; -function IconComponent({ - size, - ...props -}) { +function IconComponent({ size, ...props }) { const { controls } = useAnimateIconContext(); const variants = getVariants(animations); @@ -47,14 +48,16 @@ function IconComponent({ strokeLinejoin="round" initial="initial" animate={controls} - {...props}> + {...props} + > + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> + animate={controls} + /> ); } diff --git a/frontend/src/components/animate-ui/primitives/animate/slot.jsx b/frontend/src/components/animate-ui/primitives/animate/slot.jsx index 2994c87..12dc05f 100644 --- a/frontend/src/components/animate-ui/primitives/animate/slot.jsx +++ b/frontend/src/components/animate-ui/primitives/animate/slot.jsx @@ -1,15 +1,15 @@ -import * as React from 'react'; -import { motion, isMotionComponent } from 'motion/react'; -import { cn } from '@/lib/utils'; +import * as React from "react"; +import { motion, isMotionComponent } from "motion/react"; +import { cn } from "@/lib/utils"; function mergeRefs(...refs) { return (node) => { refs.forEach((ref) => { if (!ref) return; - if (typeof ref === 'function') { + if (typeof ref === "function") { ref(node); } else { - (ref).current = node; + ref.current = node; } }); }; @@ -24,30 +24,22 @@ function mergeProps(childProps, slotProps) { if (childProps.style || slotProps.style) { merged.style = { - ...(childProps.style), - ...(slotProps.style), + ...childProps.style, + ...slotProps.style, }; } return merged; } -function Slot( - { - children, - ref, - ...props - } -) { +function Slot({ children, ref, ...props }) { const isAlreadyMotion = - typeof children.type === 'object' && - children.type !== null && - isMotionComponent(children.type); + typeof children.type === "object" && children.type !== null && isMotionComponent(children.type); - const Base = React.useMemo(() => - isAlreadyMotion - ? (children.type) - : motion.create(children.type), [isAlreadyMotion, children.type]); + const Base = React.useMemo( + () => (isAlreadyMotion ? children.type : motion.create(children.type)), + [isAlreadyMotion, children.type] + ); if (!React.isValidElement(children)) return null; @@ -55,7 +47,7 @@ function Slot( const mergedProps = mergeProps(childProps, props); - return (); + return ; } export { Slot }; diff --git a/frontend/src/components/animate-ui/primitives/effects/highlight.jsx b/frontend/src/components/animate-ui/primitives/effects/highlight.jsx index 9750909..6015e93 100644 --- a/frontend/src/components/animate-ui/primitives/effects/highlight.jsx +++ b/frontend/src/components/animate-ui/primitives/effects/highlight.jsx @@ -1,41 +1,35 @@ -import * as React from 'react'; -import { AnimatePresence, motion } from 'motion/react'; +import * as React from "react"; +import { AnimatePresence, motion } from "motion/react"; -import { cn } from '@/lib/utils'; +import { cn } from "@/lib/utils"; -const HighlightContext = React.createContext(// eslint-disable-next-line @typescript-eslint/no-explicit-any -undefined); +const HighlightContext = React.createContext(undefined); function useHighlight() { const context = React.useContext(HighlightContext); if (!context) { - throw new Error('useHighlight must be used within a HighlightProvider'); + throw new Error("useHighlight must be used within a HighlightProvider"); } return context; } -function Highlight( - { - ref, - ...props - } -) { +function Highlight({ ref, ...props }) { const { - as: Component = 'div', + as: Component = "div", children, value, defaultValue, onValueChange, className, style, - transition = { type: 'spring', stiffness: 350, damping: 35 }, + transition = { type: "spring", stiffness: 350, damping: 35 }, hover = false, click = true, enabled = true, controlledItems, disabled = false, exitDelay = 200, - mode = 'children', + mode = "children", } = props; const localRef = React.useRef(null); @@ -43,46 +37,50 @@ function Highlight( const [activeValue, setActiveValue] = React.useState(value ?? defaultValue ?? null); const [boundsState, setBoundsState] = React.useState(null); - const [activeClassNameState, setActiveClassNameState] = - React.useState(''); - - const safeSetActiveValue = React.useCallback((id) => { - setActiveValue((prev) => (prev === id ? prev : id)); - if (id !== activeValue) onValueChange?.(id); - }, [activeValue, onValueChange]); - - const safeSetBounds = React.useCallback((bounds) => { - if (!localRef.current) return; - - const boundsOffset = (props) - ?.boundsOffset ?? { - top: 0, - left: 0, - width: 0, - height: 0, - }; + const [activeClassNameState, setActiveClassNameState] = React.useState(""); + + const safeSetActiveValue = React.useCallback( + (id) => { + setActiveValue((prev) => (prev === id ? prev : id)); + if (id !== activeValue) onValueChange?.(id); + }, + [activeValue, onValueChange] + ); - const containerRect = localRef.current.getBoundingClientRect(); - const newBounds = { - top: bounds.top - containerRect.top + (boundsOffset.top ?? 0), - left: bounds.left - containerRect.left + (boundsOffset.left ?? 0), - width: bounds.width + (boundsOffset.width ?? 0), - height: bounds.height + (boundsOffset.height ?? 0), - }; + const safeSetBounds = React.useCallback( + (bounds) => { + if (!localRef.current) return; - setBoundsState((prev) => { - if ( - prev && - prev.top === newBounds.top && - prev.left === newBounds.left && - prev.width === newBounds.width && - prev.height === newBounds.height - ) { - return prev; - } - return newBounds; - }); - }, [props]); + const boundsOffset = props?.boundsOffset ?? { + top: 0, + left: 0, + width: 0, + height: 0, + }; + + const containerRect = localRef.current.getBoundingClientRect(); + const newBounds = { + top: bounds.top - containerRect.top + (boundsOffset.top ?? 0), + left: bounds.left - containerRect.left + (boundsOffset.left ?? 0), + width: bounds.width + (boundsOffset.width ?? 0), + height: bounds.height + (boundsOffset.height ?? 0), + }; + + setBoundsState((prev) => { + if ( + prev && + prev.top === newBounds.top && + prev.left === newBounds.left && + prev.width === newBounds.width && + prev.height === newBounds.height + ) { + return prev; + } + return newBounds; + }); + }, + [props] + ); const clearBounds = React.useCallback(() => { setBoundsState((prev) => (prev === null ? prev : null)); @@ -96,75 +94,82 @@ function Highlight( const id = React.useId(); React.useEffect(() => { - if (mode !== 'parent') return; + if (mode !== "parent") return; const container = localRef.current; if (!container) return; const onScroll = () => { if (!activeValue) return; - const activeEl = container.querySelector(`[data-value="${activeValue}"][data-highlight="true"]`); + const activeEl = container.querySelector( + `[data-value="${activeValue}"][data-highlight="true"]` + ); if (activeEl) safeSetBounds(activeEl.getBoundingClientRect()); }; - container.addEventListener('scroll', onScroll, { passive: true }); - return () => container.removeEventListener('scroll', onScroll); + container.addEventListener("scroll", onScroll, { passive: true }); + return () => container.removeEventListener("scroll", onScroll); }, [mode, activeValue, safeSetBounds]); - const render = React.useCallback((children) => { - if (mode === 'parent') { - return ( - - - {boundsState && ( - - )} - - {children} - - ); - } + const render = React.useCallback( + (children) => { + if (mode === "parent") { + return ( + + + {boundsState && ( + + )} + + {children} + + ); + } - return children; - }, [ - mode, - Component, - props, - boundsState, - transition, - exitDelay, - style, - className, - activeClassNameState, - ]); + return children; + }, + [ + mode, + Component, + props, + boundsState, + transition, + exitDelay, + style, + className, + activeClassNameState, + ] + ); return ( + forceUpdateBounds: props?.forceUpdateBounds, + }} + > {enabled ? controlledItems ? render(children) - : render(React.Children.map(children, (child, index) => ( - - {child} - - ))) + : render( + React.Children.map(children, (child, index) => ( + + {child} + + )) + ) : children} ); @@ -203,31 +210,29 @@ function Highlight( function getNonOverridingDataAttributes(element, dataAttributes) { return Object.keys(dataAttributes).reduce((acc, key) => { - if ((element.props)[key] === undefined) { + if (element.props[key] === undefined) { acc[key] = dataAttributes[key]; } return acc; }, {}); } -function HighlightItem( - { - ref, - as, - children, - id, - value, - className, - style, - transition, - disabled = false, - activeClassName, - exitDelay, - asChild = false, - forceUpdateBounds, - ...props - } -) { +function HighlightItem({ + ref, + as, + children, + id, + value, + className, + style, + transition, + disabled = false, + activeClassName, + exitDelay, + asChild = false, + forceUpdateBounds, + ...props +}) { const itemId = React.useId(); const { activeValue, @@ -248,10 +253,9 @@ function HighlightItem( setActiveClassName, } = useHighlight(); - const Component = as ?? 'div'; + const Component = as ?? "div"; const element = children; - const childValue = - id ?? value ?? element.props?.['data-value'] ?? element.props?.id ?? itemId; + const childValue = id ?? value ?? element.props?.["data-value"] ?? element.props?.id ?? itemId; const isActive = activeValue === childValue; const isDisabled = disabled === undefined ? contextDisabled : disabled; const itemTransition = transition ?? contextTransition; @@ -260,12 +264,11 @@ function HighlightItem( React.useImperativeHandle(ref, () => localRef.current); React.useEffect(() => { - if (mode !== 'parent') return; + if (mode !== "parent") return; let rafId; let previousBounds = null; const shouldUpdateBounds = - forceUpdateBounds === true || - (contextForceUpdateBounds && forceUpdateBounds !== false); + forceUpdateBounds === true || (contextForceUpdateBounds && forceUpdateBounds !== false); const updateBounds = () => { if (!localRef.current) return; @@ -292,7 +295,7 @@ function HighlightItem( if (isActive) { updateBounds(); - setActiveClassName(activeClassName ?? ''); + setActiveClassName(activeClassName ?? ""); } else if (!activeValue) clearBounds(); if (shouldUpdateBounds) return () => cancelAnimationFrame(rafId); @@ -311,11 +314,11 @@ function HighlightItem( if (!React.isValidElement(children)) return children; const dataAttributes = { - 'data-active': isActive ? 'true' : 'false', - 'aria-selected': isActive, - 'data-disabled': isDisabled, - 'data-value': childValue, - 'data-highlight': true, + "data-active": isActive ? "true" : "false", + "aria-selected": isActive, + "data-disabled": isDisabled, + "data-value": childValue, + "data-highlight": true, }; const commonHandlers = hover @@ -339,61 +342,66 @@ function HighlightItem( : {}; if (asChild) { - if (mode === 'children') { - return React.cloneElement(element, { - key: childValue, - ref: localRef, - className: cn('relative', element.props.className), - ...getNonOverridingDataAttributes(element, { - ...dataAttributes, - 'data-slot': 'motion-highlight-item-container', - }), - ...commonHandlers, - ...props, - }, <> - - {isActive && !isDisabled && ( - - )} - + if (mode === "children") { + return React.cloneElement( + element, + { + key: childValue, + ref: localRef, + className: cn("relative", element.props.className), + ...getNonOverridingDataAttributes(element, { + ...dataAttributes, + "data-slot": "motion-highlight-item-container", + }), + ...commonHandlers, + ...props, + }, + <> + + {isActive && !isDisabled && ( + + )} + - - {children} - - ); + + {children} + + + ); } return React.cloneElement(element, { ref: localRef, ...getNonOverridingDataAttributes(element, { ...dataAttributes, - 'data-slot': 'motion-highlight-item', + "data-slot": "motion-highlight-item", }), ...commonHandlers, }); @@ -404,18 +412,19 @@ function HighlightItem( key={childValue} ref={localRef} data-slot="motion-highlight-item-container" - className={cn(mode === 'children' && 'relative', className)} + className={cn(mode === "children" && "relative", className)} {...dataAttributes} {...props} - {...commonHandlers}> - {mode === 'children' && ( + {...commonHandlers} + > + {mode === "children" && ( {isActive && !isDisabled && ( + {...dataAttributes} + /> )} )} {React.cloneElement(element, { - style: { position: 'relative', zIndex: 1 }, + style: { position: "relative", zIndex: 1 }, className: element.props.className, ...getNonOverridingDataAttributes(element, { ...dataAttributes, - 'data-slot': 'motion-highlight-item', + "data-slot": "motion-highlight-item", }), })} diff --git a/frontend/src/components/auth/forgot-password-form.jsx b/frontend/src/components/auth/forgot-password-form.jsx index 25437b5..f2828c3 100644 --- a/frontend/src/components/auth/forgot-password-form.jsx +++ b/frontend/src/components/auth/forgot-password-form.jsx @@ -1,12 +1,12 @@ -import React, { useState, useCallback, useMemo, useEffect, useRef } from 'react' -import { useForm } from 'react-hook-form' -import { zodResolver } from '@hookform/resolvers/zod' -import { z } from 'zod' -import { Link, useLocation } from 'react-router-dom' -import { useTranslation } from 'react-i18next' -import { cn, debugWarn } from '@/lib/utils' -import { Button } from '@/components/ui/button' -import { MailCheck } from 'lucide-react' +import React, { useState, useCallback, useMemo, useEffect, useRef } from "react"; +import { useForm } from "react-hook-form"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { z } from "zod"; +import { Link, useLocation } from "react-router-dom"; +import { useTranslation } from "react-i18next"; +import { cn, debugWarn } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; +import { MailCheck } from "lucide-react"; import { Form, FormControl, @@ -14,75 +14,79 @@ import { FormItem, FormLabel, FormMessage, -} from '@/components/ui/form' -import { Input } from '@/components/ui/input' -import { authService } from '@/services/auth.service' -import { debugError } from '@/lib/utils' -import { Spinner } from '@/components/ui/spinner' -import { useIsMobile } from '@/hooks/useMobile' +} from "@/components/ui/form"; +import { Input } from "@/components/ui/input"; +import { authService } from "@/services/auth.service"; +import { debugError } from "@/lib/utils"; +import { Spinner } from "@/components/ui/spinner"; +import { useIsMobile } from "@/hooks/useMobile"; const SubmitButton = React.memo(({ onSubmit, t, className, isSubmitting }) => { return ( - - ) -}) + ); +}); -SubmitButton.displayName = 'SubmitButton' +SubmitButton.displayName = "SubmitButton"; export const ForgotPasswordForm = ({ className, onStateChange, ...props }) => { - const location = useLocation() - const { t } = useTranslation() - const isMobile = useIsMobile() - const [isSubmitting, setIsSubmitting] = useState(false) - const [isResending, setIsResending] = useState(false) - const [cooldownSeconds, setCooldownSeconds] = useState(0) - const [isLoadingCooldown, setIsLoadingCooldown] = useState(false) - const [showConfirmation, setShowConfirmation] = useState(false) - const [email, setEmail] = useState('') - const intervalRef = useRef(null) - + const location = useLocation(); + const { t } = useTranslation(); + const isMobile = useIsMobile(); + const [isSubmitting, setIsSubmitting] = useState(false); + const [isResending, setIsResending] = useState(false); + const [cooldownSeconds, setCooldownSeconds] = useState(0); + const [isLoadingCooldown, setIsLoadingCooldown] = useState(false); + const [showConfirmation, setShowConfirmation] = useState(false); + const [email, setEmail] = useState(""); + const intervalRef = useRef(null); + const fetchCooldown = useCallback(async (emailToCheck) => { - if (!emailToCheck) return - - setIsLoadingCooldown(true) + if (!emailToCheck) return; + + setIsLoadingCooldown(true); try { - const result = await authService.getPasswordResetCooldown(emailToCheck) - if (result.status === 'success' && result.data?.data) { - setCooldownSeconds(result.data.data.cooldown_seconds || 0) - } else if (result.status === 'success' && result.data?.cooldown_seconds !== undefined) { - setCooldownSeconds(result.data.cooldown_seconds || 0) + const result = await authService.getPasswordResetCooldown(emailToCheck); + if (result.status === "success" && result.data?.data) { + setCooldownSeconds(result.data.data.cooldown_seconds || 0); + } else if (result.status === "success" && result.data?.cooldown_seconds !== undefined) { + setCooldownSeconds(result.data.cooldown_seconds || 0); } } catch (error) { - debugError('Failed to fetch cooldown:', error) + debugError("Failed to fetch cooldown:", error); } finally { - setIsLoadingCooldown(false) + setIsLoadingCooldown(false); } - }, []) - + }, []); + // Initialize email from location state (only when coming from other pages) useEffect(() => { - const stateEmail = location.state?.email - + const stateEmail = location.state?.email; + if (stateEmail) { - setEmail(stateEmail) - setShowConfirmation(true) + setEmail(stateEmail); + setShowConfirmation(true); // Notify parent component about confirmation state if (onStateChange) { - onStateChange(true) + onStateChange(true); } // Fetch cooldown status - fetchCooldown(stateEmail) + fetchCooldown(stateEmail); } - }, [fetchCooldown, location.state?.email, onStateChange]) + }, [fetchCooldown, location.state?.email, onStateChange]); // Countdown timer - only countdown, stop at 0 useEffect(() => { @@ -91,207 +95,229 @@ export const ForgotPasswordForm = ({ className, onStateChange, ...props }) => { setCooldownSeconds((prev) => { if (prev <= 1) { if (intervalRef.current) { - clearInterval(intervalRef.current) - intervalRef.current = null + clearInterval(intervalRef.current); + intervalRef.current = null; } - return 0 + return 0; } - return prev - 1 - }) - }, 1000) + return prev - 1; + }); + }, 1000); } else { if (intervalRef.current) { - clearInterval(intervalRef.current) - intervalRef.current = null + clearInterval(intervalRef.current); + intervalRef.current = null; } } return () => { if (intervalRef.current) { - clearInterval(intervalRef.current) - intervalRef.current = null + clearInterval(intervalRef.current); + intervalRef.current = null; } - } - }, [cooldownSeconds]) + }; + }, [cooldownSeconds]); const formSchema = useMemo(() => { return z.object({ email: z .string() - .min(1, t('pages.auth.forgotPassword.fields.email.validation.required', { defaultValue: 'Please enter your email' })) - .refine((val) => { - const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/ - return emailRegex.test(val) - }, { - message: t('pages.auth.forgotPassword.fields.email.validation.invalid', { defaultValue: 'Please enter a valid email format' }), - }), - }) - }, [t]) - - const stableDefaultValues = useMemo(() => ({ - email: email || '', - }), [email]) + .min( + 1, + t("pages.auth.forgotPassword.fields.email.validation.required", { + defaultValue: "Please enter your email", + }) + ) + .refine( + (val) => { + const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/; + return emailRegex.test(val); + }, + { + message: t("pages.auth.forgotPassword.fields.email.validation.invalid", { + defaultValue: "Please enter a valid email format", + }), + } + ), + }); + }, [t]); + + const stableDefaultValues = useMemo( + () => ({ + email: email || "", + }), + [email] + ); const form = useForm({ resolver: zodResolver(formSchema), defaultValues: stableDefaultValues, - }) + }); // Update form when email changes useEffect(() => { - if (email && form.getValues('email') !== email) { - form.setValue('email', email, { shouldValidate: false }) + if (email && form.getValues("email") !== email) { + form.setValue("email", email, { shouldValidate: false }); } - }, [email, form]) + }, [email, form]); useEffect(() => { try { - form.clearErrors() - const newResolver = zodResolver(formSchema) - - if (form._options && typeof form._options === 'object' && 'resolver' in form._options) { - form._options.resolver = newResolver + form.clearErrors(); + const newResolver = zodResolver(formSchema); + + if (form._options && typeof form._options === "object" && "resolver" in form._options) { + form._options.resolver = newResolver; } - - if ('_resolver' in form && form._resolver !== undefined) { - form._resolver = newResolver + + if ("_resolver" in form && form._resolver !== undefined) { + form._resolver = newResolver; } - + if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } catch (error) { - debugWarn('Failed to update form resolver:', error) - form.clearErrors() + debugWarn("Failed to update form resolver:", error); + form.clearErrors(); if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } - }, [formSchema, form]) + }, [formSchema, form]); - const formMethodsRef = useRef(form) + const formMethodsRef = useRef(form); useEffect(() => { - formMethodsRef.current = form - }, [form]) + formMethodsRef.current = form; + }, [form]); + + const handleSubmit = useCallback( + async (formValues) => { + setIsSubmitting(true); + + try { + await authService.forgotPassword(formValues.email, { + showErrorToast: true, + showSuccessToast: true, + }); + + // Success - switch to confirmation mode + const submittedEmail = formValues.email; + setEmail(submittedEmail); + setShowConfirmation(true); - const handleSubmit = useCallback(async (formValues) => { - setIsSubmitting(true) - - try { - await authService.forgotPassword(formValues.email, { - showErrorToast: true, - showSuccessToast: true, - }) - - // Success - switch to confirmation mode - const submittedEmail = formValues.email - setEmail(submittedEmail) - setShowConfirmation(true) - - // Notify parent component about confirmation state - if (onStateChange) { - onStateChange(true) - } - - // Fetch cooldown status - await fetchCooldown(submittedEmail) - - setIsSubmitting(false) - } catch (error) { - debugError('Forgot password error:', error) - // If error is due to cooldown, switch to confirmation mode and fetch cooldown - if (error.response?.status === 400 && formValues.email) { - const submittedEmail = formValues.email - setEmail(submittedEmail) - setShowConfirmation(true) - // Notify parent component about confirmation state if (onStateChange) { - onStateChange(true) + onStateChange(true); + } + + // Fetch cooldown status + await fetchCooldown(submittedEmail); + + setIsSubmitting(false); + } catch (error) { + debugError("Forgot password error:", error); + // If error is due to cooldown, switch to confirmation mode and fetch cooldown + if (error.response?.status === 400 && formValues.email) { + const submittedEmail = formValues.email; + setEmail(submittedEmail); + setShowConfirmation(true); + + // Notify parent component about confirmation state + if (onStateChange) { + onStateChange(true); + } + + await fetchCooldown(submittedEmail); } - - await fetchCooldown(submittedEmail) + setIsSubmitting(false); } - setIsSubmitting(false) - } - }, [fetchCooldown, onStateChange]) + }, + [fetchCooldown, onStateChange] + ); const handleResend = useCallback(async () => { if (cooldownSeconds > 0 || isResending || !email) { - return + return; } - setIsResending(true) - + setIsResending(true); + try { await authService.forgotPassword(email, { showErrorToast: true, showSuccessToast: true, - }) - + }); + // Fetch new cooldown after successful send - await fetchCooldown(email) + await fetchCooldown(email); } catch (error) { - debugError('Resend email error:', error) + debugError("Resend email error:", error); // If error is due to cooldown, fetch the current cooldown status if (error.response?.status === 400) { - await fetchCooldown(email) + await fetchCooldown(email); } } finally { - setIsResending(false) + setIsResending(false); } - }, [email, cooldownSeconds, isResending, fetchCooldown]) + }, [email, cooldownSeconds, isResending, fetchCooldown]); const formatTime = (seconds) => { - const mins = Math.floor(seconds / 60) - const secs = seconds % 60 - return `${mins}:${secs.toString().padStart(2, '0')}` - } + const mins = Math.floor(seconds / 60); + const secs = seconds % 60; + return `${mins}:${secs.toString().padStart(2, "0")}`; + }; - const onSubmitHandler = useCallback((e) => { - e?.preventDefault?.() - formMethodsRef.current.handleSubmit(handleSubmit)() - }, [handleSubmit]) + const onSubmitHandler = useCallback( + (e) => { + e?.preventDefault?.(); + formMethodsRef.current.handleSubmit(handleSubmit)(); + }, + [handleSubmit] + ); - const handleKeyDown = useCallback((e) => { - if (e.key === 'Enter' && !isSubmitting && !showConfirmation) { - e.preventDefault() - onSubmitHandler(e) - } - }, [onSubmitHandler, isSubmitting, showConfirmation]) + const handleKeyDown = useCallback( + (e) => { + if (e.key === "Enter" && !isSubmitting && !showConfirmation) { + e.preventDefault(); + onSubmitHandler(e); + } + }, + [onSubmitHandler, isSubmitting, showConfirmation] + ); - const emailInputRef = useRef(null) + const emailInputRef = useRef(null); useEffect(() => { if (!isMobile && emailInputRef.current && !showConfirmation) { - emailInputRef.current.focus() + emailInputRef.current.focus(); } - }, [isMobile, showConfirmation]) - + }, [isMobile, showConfirmation]); + // Confirmation view if (showConfirmation && email) { return ( -
+
- +

- {t('pages.auth.forgotPasswordConfirmation.description', { - defaultValue: 'We\'ve sent a password reset link to {{email}}', - email + {t("pages.auth.forgotPasswordConfirmation.description", { + defaultValue: "We've sent a password reset link to {{email}}", + email, })}

- ) + ); } // Input form view return (
- - + ( - {t('pages.auth.forgotPassword.fields.email.label', { defaultValue: 'Email' })} + + {t("pages.auth.forgotPassword.fields.email.label", { defaultValue: "Email" })} + - { - emailInputRef.current = e - field.ref(e) + emailInputRef.current = e; + field.ref(e); }} - disabled={isSubmitting} + disabled={isSubmitting} /> )} /> - - - +
- - {t('pages.auth.forgotPassword.actions.backToLogin', { defaultValue: 'Back to Login' })} + {t("pages.auth.forgotPassword.actions.backToLogin", { defaultValue: "Back to Login" })}
- ) -} + ); +}; -export default ForgotPasswordForm +export default ForgotPasswordForm; diff --git a/frontend/src/components/auth/login-form.jsx b/frontend/src/components/auth/login-form.jsx index 45a71c7..4309d65 100644 --- a/frontend/src/components/auth/login-form.jsx +++ b/frontend/src/components/auth/login-form.jsx @@ -1,12 +1,12 @@ -import React, { useState, useCallback, useMemo, useEffect, useLayoutEffect, useRef } from 'react' -import ENV from '@/config/env.config' -import { useForm } from 'react-hook-form' -import { zodResolver } from '@hookform/resolvers/zod' -import { z } from 'zod' -import { useNavigate, Link, useLocation } from 'react-router-dom' -import { useTranslation } from 'react-i18next' -import { cn, debugWarn } from '@/lib/utils' -import { Button } from '@/components/ui/button' +import React, { useState, useCallback, useMemo, useEffect, useLayoutEffect, useRef } from "react"; +import ENV from "@/config/env.config"; +import { useForm } from "react-hook-form"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { z } from "zod"; +import { useNavigate, Link, useLocation } from "react-router-dom"; +import { useTranslation } from "react-i18next"; +import { cn, debugWarn } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; import { Form, FormControl, @@ -14,318 +14,371 @@ import { FormItem, FormLabel, FormMessage, -} from '@/components/ui/form' -import { Input } from '@/components/ui/input' -import { useAuth } from '@/hooks/useAuth' -import { debugError } from '@/lib/utils' -import { Spinner } from '@/components/ui/spinner' -import { useIsMobile } from '@/hooks/useMobile' +} from "@/components/ui/form"; +import { Input } from "@/components/ui/input"; +import { useAuth } from "@/hooks/useAuth"; +import { debugError } from "@/lib/utils"; +import { Spinner } from "@/components/ui/spinner"; +import { useIsMobile } from "@/hooks/useMobile"; const LoginButton = React.memo(({ onSubmit, t, className, isSubmitting }) => { return ( - - ) -}) - -LoginButton.displayName = 'LoginButton' - -const emailPersistRef = { current: '' } -const formInstanceRef = { current: null } - -export const LoginForm = ({ className, redirectTo = '/', ...props }) => { - const navigate = useNavigate() - const location = useLocation() - const { t } = useTranslation() - const authContext = useAuth() - const isMobile = useIsMobile() - const [isSubmitting, setIsSubmitting] = useState(false) - - const login = useMemo(() => authContext.login, [authContext.login]) - const clearError = useMemo(() => authContext.clearError, [authContext.clearError]) - - const handleLogin = useCallback(async (credentials) => { - clearError() - return await login(credentials) - }, [login, clearError]) - - const emailRef = useRef(emailPersistRef.current) - const emailInputRef = useRef(null) + ); +}); + +LoginButton.displayName = "LoginButton"; + +const emailPersistRef = { current: "" }; +const formInstanceRef = { current: null }; + +export const LoginForm = ({ className, redirectTo = "/", ...props }) => { + const navigate = useNavigate(); + const location = useLocation(); + const { t } = useTranslation(); + const authContext = useAuth(); + const isMobile = useIsMobile(); + const [isSubmitting, setIsSubmitting] = useState(false); + + const login = useMemo(() => authContext.login, [authContext.login]); + const clearError = useMemo(() => authContext.clearError, [authContext.clearError]); + + const handleLogin = useCallback( + async (credentials) => { + clearError(); + return await login(credentials); + }, + [login, clearError] + ); + + const emailRef = useRef(emailPersistRef.current); + const emailInputRef = useRef(null); const [emailValue, setEmailValue] = useState(() => { - const storedEmail = emailPersistRef.current + const storedEmail = emailPersistRef.current; if (storedEmail) { - emailRef.current = storedEmail + emailRef.current = storedEmail; } - return storedEmail || '' - }) - - const handleLoginRef = useRef(handleLogin) + return storedEmail || ""; + }); + + const handleLoginRef = useRef(handleLogin); useEffect(() => { - handleLoginRef.current = handleLogin - }, [handleLogin]) + handleLoginRef.current = handleLogin; + }, [handleLogin]); const formSchema = useMemo(() => { return z.object({ email: z .string() - .min(1, t('pages.auth.login.fields.email.validation.required', { defaultValue: 'Please enter your email' })) - .refine((val) => { - const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/ - return emailRegex.test(val) - }, { - message: t('pages.auth.login.fields.email.validation.invalid', { defaultValue: 'Please enter a valid email format' }), - }), + .min( + 1, + t("pages.auth.login.fields.email.validation.required", { + defaultValue: "Please enter your email", + }) + ) + .refine( + (val) => { + const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/; + return emailRegex.test(val); + }, + { + message: t("pages.auth.login.fields.email.validation.invalid", { + defaultValue: "Please enter a valid email format", + }), + } + ), password: z .string() - .min(1, t('pages.auth.login.fields.password.validation.required', { defaultValue: 'Please enter your password' })) - .min(6, t('pages.auth.login.fields.password.validation.minLength', { defaultValue: 'Password must be at least 6 characters' })), - }) - }, [t]) + .min( + 1, + t("pages.auth.login.fields.password.validation.required", { + defaultValue: "Please enter your password", + }) + ) + .min( + 6, + t("pages.auth.login.fields.password.validation.minLength", { + defaultValue: "Password must be at least 6 characters", + }) + ), + }); + }, [t]); - const stableDefaultValues = useMemo(() => ({ - email: emailPersistRef.current || '', - password: '', - }), []) + const stableDefaultValues = useMemo( + () => ({ + email: emailPersistRef.current || "", + password: "", + }), + [] + ); const form = useForm({ resolver: zodResolver(formSchema), defaultValues: stableDefaultValues, - }) + }); useEffect(() => { try { - form.clearErrors() - const newResolver = zodResolver(formSchema) - - if (form._options && typeof form._options === 'object' && 'resolver' in form._options) { - form._options.resolver = newResolver + form.clearErrors(); + const newResolver = zodResolver(formSchema); + + if (form._options && typeof form._options === "object" && "resolver" in form._options) { + form._options.resolver = newResolver; } - - if ('_resolver' in form && form._resolver !== undefined) { - form._resolver = newResolver + + if ("_resolver" in form && form._resolver !== undefined) { + form._resolver = newResolver; } - + if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } catch (error) { - debugWarn('Failed to update form resolver:', error) - form.clearErrors() + debugWarn("Failed to update form resolver:", error); + form.clearErrors(); if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } - }, [formSchema, form]) + }, [formSchema, form]); useEffect(() => { if (!formInstanceRef.current) { - formInstanceRef.current = form + formInstanceRef.current = form; } - }, [form]) + }, [form]); - const formMethodsRef = useRef(form) + const formMethodsRef = useRef(form); useEffect(() => { - formMethodsRef.current = form - }, [form]) - + formMethodsRef.current = form; + }, [form]); + useEffect(() => { if (emailPersistRef.current) { - const currentEmail = form.getValues('email') + const currentEmail = form.getValues("email"); if (currentEmail !== emailPersistRef.current) { - form.setValue('email', emailPersistRef.current, { shouldValidate: false, shouldDirty: false, shouldTouch: false }) + form.setValue("email", emailPersistRef.current, { + shouldValidate: false, + shouldDirty: false, + shouldTouch: false, + }); } } - }, [form]) + }, [form]); useLayoutEffect(() => { - const storedEmail = emailPersistRef.current + const storedEmail = emailPersistRef.current; if (storedEmail) { if (emailRef.current !== storedEmail || emailValue !== storedEmail) { - emailRef.current = storedEmail - setEmailValue(storedEmail) + emailRef.current = storedEmail; + setEmailValue(storedEmail); if (formMethodsRef.current) { - formMethodsRef.current.setValue('email', storedEmail, { shouldValidate: false, shouldDirty: false, shouldTouch: false }) + formMethodsRef.current.setValue("email", storedEmail, { + shouldValidate: false, + shouldDirty: false, + shouldTouch: false, + }); } } } - }, [emailValue]) + }, [emailValue]); useEffect(() => { const subscription = form.watch((value, { name, type }) => { - if (name === 'email' && value.email !== undefined && type === 'change') { - const currentEmail = value.email + if (name === "email" && value.email !== undefined && type === "change") { + const currentEmail = value.email; if (currentEmail !== emailRef.current && (currentEmail || !emailPersistRef.current)) { - emailRef.current = currentEmail - emailPersistRef.current = currentEmail - setEmailValue(currentEmail) + emailRef.current = currentEmail; + emailPersistRef.current = currentEmail; + setEmailValue(currentEmail); } } - }) + }); return () => { - subscription.unsubscribe() - } - }, [form]) + subscription.unsubscribe(); + }; + }, [form]); + + const navigateRef = useRef(navigate); + const tRef = useRef(t); + const redirectToRef = useRef(redirectTo); - const navigateRef = useRef(navigate) - const tRef = useRef(t) - const redirectToRef = useRef(redirectTo) - useEffect(() => { - navigateRef.current = navigate - tRef.current = t - redirectToRef.current = redirectTo - }, [navigate, t, redirectTo]) - - const setEmailValueRef = useRef(setEmailValue) + navigateRef.current = navigate; + tRef.current = t; + redirectToRef.current = redirectTo; + }, [navigate, t, redirectTo]); + + const setEmailValueRef = useRef(setEmailValue); useEffect(() => { - setEmailValueRef.current = setEmailValue - }, [setEmailValue]) - + setEmailValueRef.current = setEmailValue; + }, [setEmailValue]); + const handleSubmit = useCallback(async (formValues) => { - const currentEmail = formValues.email || emailRef.current || emailPersistRef.current - - emailRef.current = currentEmail - emailPersistRef.current = currentEmail - + const currentEmail = formValues.email || emailRef.current || emailPersistRef.current; + + emailRef.current = currentEmail; + emailPersistRef.current = currentEmail; + const data = { email: currentEmail, password: formValues.password, - } - - setIsSubmitting(true) - + }; + + setIsSubmitting(true); + try { - const result = await handleLoginRef.current(data) - + const result = await handleLoginRef.current(data); + if (result.success) { - emailRef.current = '' - emailPersistRef.current = '' - setEmailValueRef.current('') - formMethodsRef.current.reset({ email: '', password: '' }, { keepValues: false }) + emailRef.current = ""; + emailPersistRef.current = ""; + setEmailValueRef.current(""); + formMethodsRef.current.reset({ email: "", password: "" }, { keepValues: false }); setTimeout(() => { - navigateRef.current(redirectToRef.current, { replace: true }) - }, 0) + navigateRef.current(redirectToRef.current, { replace: true }); + }, 0); } else if (result.requiresPasswordReset && result.resetToken) { // Redirect to reset password page with token - emailRef.current = '' - emailPersistRef.current = '' - setEmailValueRef.current('') - formMethodsRef.current.reset({ email: '', password: '' }, { keepValues: false }) + emailRef.current = ""; + emailPersistRef.current = ""; + setEmailValueRef.current(""); + formMethodsRef.current.reset({ email: "", password: "" }, { keepValues: false }); setTimeout(() => { - navigateRef.current(`/auth/reset-password?token=${encodeURIComponent(result.resetToken)}`, { replace: true }) - }, 0) + navigateRef.current( + `/auth/reset-password?token=${encodeURIComponent(result.resetToken)}`, + { replace: true } + ); + }, 0); } else if (result.requiresEmailVerification) { // Email verification required - redirect to verification page with email - const emailToKeep = formValues.email || emailRef.current || emailPersistRef.current - - emailRef.current = emailToKeep - emailPersistRef.current = emailToKeep - - formMethodsRef.current.reset({ email: emailToKeep, password: '' }, { keepValues: false }) - - setEmailValueRef.current(emailToKeep) - - setIsSubmitting(false) - + const emailToKeep = formValues.email || emailRef.current || emailPersistRef.current; + + emailRef.current = emailToKeep; + emailPersistRef.current = emailToKeep; + + formMethodsRef.current.reset({ email: emailToKeep, password: "" }, { keepValues: false }); + + setEmailValueRef.current(emailToKeep); + + setIsSubmitting(false); + // Redirect to verification page with email in state requestAnimationFrame(() => { requestAnimationFrame(() => { - navigateRef.current('/auth/verify-email', { + navigateRef.current("/auth/verify-email", { replace: true, - state: { email: emailToKeep } + state: { email: emailToKeep }, }); }); }); } else { - const emailToKeep = formValues.email || emailRef.current || emailPersistRef.current - - emailRef.current = emailToKeep - emailPersistRef.current = emailToKeep - - formMethodsRef.current.setValue('password', '', { shouldValidate: false, shouldDirty: false, shouldTouch: false }) - - setEmailValueRef.current(prev => prev !== emailToKeep ? emailToKeep : prev) - - setIsSubmitting(false) + const emailToKeep = formValues.email || emailRef.current || emailPersistRef.current; + + emailRef.current = emailToKeep; + emailPersistRef.current = emailToKeep; + + formMethodsRef.current.setValue("password", "", { + shouldValidate: false, + shouldDirty: false, + shouldTouch: false, + }); + + setEmailValueRef.current((prev) => (prev !== emailToKeep ? emailToKeep : prev)); + + setIsSubmitting(false); } } catch (error) { - debugError('Login error:', error) - const emailToKeep = formValues.email || emailRef.current || emailPersistRef.current - - emailRef.current = emailToKeep - emailPersistRef.current = emailToKeep - - formMethodsRef.current.setValue('password', '', { shouldValidate: false, shouldDirty: false, shouldTouch: false }) - - setEmailValueRef.current(prev => prev !== emailToKeep ? emailToKeep : prev) - - setIsSubmitting(false) - } - }, []) + debugError("Login error:", error); + const emailToKeep = formValues.email || emailRef.current || emailPersistRef.current; + + emailRef.current = emailToKeep; + emailPersistRef.current = emailToKeep; - const onSubmitHandler = useCallback((e) => { - e?.preventDefault?.() - formMethodsRef.current.handleSubmit(handleSubmit)() - }, [handleSubmit]) + formMethodsRef.current.setValue("password", "", { + shouldValidate: false, + shouldDirty: false, + shouldTouch: false, + }); - const handleKeyDown = useCallback((e) => { - if (e.key === 'Enter' && !isSubmitting) { - e.preventDefault() - onSubmitHandler(e) + setEmailValueRef.current((prev) => (prev !== emailToKeep ? emailToKeep : prev)); + + setIsSubmitting(false); } - }, [onSubmitHandler, isSubmitting]) + }, []); + + const onSubmitHandler = useCallback( + (e) => { + e?.preventDefault?.(); + formMethodsRef.current.handleSubmit(handleSubmit)(); + }, + [handleSubmit] + ); + + const handleKeyDown = useCallback( + (e) => { + if (e.key === "Enter" && !isSubmitting) { + e.preventDefault(); + onSubmitHandler(e); + } + }, + [onSubmitHandler, isSubmitting] + ); useEffect(() => { if (!isMobile && emailInputRef.current) { - emailInputRef.current.focus() + emailInputRef.current.focus(); } - }, [isMobile]) + }, [isMobile]); return (
- + ( - {t('pages.auth.login.fields.email.label', { defaultValue: 'Email' })} + + {t("pages.auth.login.fields.email.label", { defaultValue: "Email" })} + - { - const newValue = e.target.value - emailRef.current = newValue - emailPersistRef.current = newValue - setEmailValue(newValue) - field.onChange(e) + const newValue = e.target.value; + emailRef.current = newValue; + emailPersistRef.current = newValue; + setEmailValue(newValue); + field.onChange(e); }} onBlur={field.onBlur} name={field.name} ref={(e) => { - emailInputRef.current = e - field.ref(e) + emailInputRef.current = e; + field.ref(e); }} - disabled={isSubmitting} + disabled={isSubmitting} /> @@ -338,54 +391,60 @@ export const LoginForm = ({ className, redirectTo = '/', ...props }) => { render={({ field }) => (
- {t('pages.auth.login.fields.password.label', { defaultValue: 'Password' })} + + {t("pages.auth.login.fields.password.label", { defaultValue: "Password" })} + {ENV.SMTP_ENABLE && ( - - {t('pages.auth.login.links.forgotPassword', { defaultValue: 'Forgot password?' })} + {t("pages.auth.login.links.forgotPassword", { + defaultValue: "Forgot password?", + })} )}
-
)} /> - - - + {ENV.REGISTRATION_ENABLE && (
- {t('pages.auth.login.links.newUser', { defaultValue: 'New user? ' })} - + {t("pages.auth.login.links.newUser", { defaultValue: "New user? " })} + + - {t('pages.auth.login.links.register', { defaultValue: 'Sign up' })} + {t("pages.auth.login.links.register", { defaultValue: "Sign up" })}
)} - ) -} + ); +}; -export default LoginForm \ No newline at end of file +export default LoginForm; diff --git a/frontend/src/components/auth/register-form.jsx b/frontend/src/components/auth/register-form.jsx index 8b0c33e..70f7f79 100644 --- a/frontend/src/components/auth/register-form.jsx +++ b/frontend/src/components/auth/register-form.jsx @@ -1,11 +1,11 @@ -import React, { useState, useCallback, useMemo, useEffect, useRef } from 'react' -import { useForm } from 'react-hook-form' -import { zodResolver } from '@hookform/resolvers/zod' -import { z } from 'zod' -import { useNavigate, Link, useLocation } from 'react-router-dom' -import { useTranslation } from 'react-i18next' -import { cn, debugWarn } from '@/lib/utils' -import { Button } from '@/components/ui/button' +import React, { useState, useCallback, useMemo, useEffect, useRef } from "react"; +import { useForm } from "react-hook-form"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { z } from "zod"; +import { useNavigate, Link, useLocation } from "react-router-dom"; +import { useTranslation } from "react-i18next"; +import { cn, debugWarn } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; import { Form, FormControl, @@ -13,224 +13,277 @@ import { FormItem, FormLabel, FormMessage, -} from '@/components/ui/form' -import { Input } from '@/components/ui/input' -import { useAuth } from '@/hooks/useAuth' -import { debugError } from '@/lib/utils' -import { Spinner } from '@/components/ui/spinner' -import { useIsMobile } from '@/hooks/useMobile' +} from "@/components/ui/form"; +import { Input } from "@/components/ui/input"; +import { useAuth } from "@/hooks/useAuth"; +import { debugError } from "@/lib/utils"; +import { Spinner } from "@/components/ui/spinner"; +import { useIsMobile } from "@/hooks/useMobile"; const RegisterButton = React.memo(({ onSubmit, t, className, isSubmitting }) => { return ( - - ) -}) + ); +}); -RegisterButton.displayName = 'RegisterButton' +RegisterButton.displayName = "RegisterButton"; -export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { - const navigate = useNavigate() - const location = useLocation() - const { t } = useTranslation() - const authContext = useAuth() - const isMobile = useIsMobile() - const [isSubmitting, setIsSubmitting] = useState(false) - - const register = useMemo(() => authContext.register, [authContext.register]) - const clearError = useMemo(() => authContext.clearError, [authContext.clearError]) - - const handleRegister = useCallback(async (userData) => { - clearError() - return await register(userData) - }, [register, clearError]) - - const handleRegisterRef = useRef(handleRegister) +export const RegisterForm = ({ className, redirectTo = "/", ...props }) => { + const navigate = useNavigate(); + const location = useLocation(); + const { t } = useTranslation(); + const authContext = useAuth(); + const isMobile = useIsMobile(); + const [isSubmitting, setIsSubmitting] = useState(false); + + const register = useMemo(() => authContext.register, [authContext.register]); + const clearError = useMemo(() => authContext.clearError, [authContext.clearError]); + + const handleRegister = useCallback( + async (userData) => { + clearError(); + return await register(userData); + }, + [register, clearError] + ); + + const handleRegisterRef = useRef(handleRegister); useEffect(() => { - handleRegisterRef.current = handleRegister - }, [handleRegister]) + handleRegisterRef.current = handleRegister; + }, [handleRegister]); const formSchema = useMemo(() => { - return z.object({ - first_name: z - .string() - .min(1, t('pages.auth.register.fields.firstName.validation.required', { defaultValue: 'Please enter your first name' })), - last_name: z - .string() - .min(1, t('pages.auth.register.fields.lastName.validation.required', { defaultValue: 'Please enter your last name' })), - email: z - .string() - .min(1, t('pages.auth.register.fields.email.validation.required', { defaultValue: 'Please enter your email' })) - .transform((val) => val.trim().toLowerCase()) - .refine((val) => { - const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/ - return emailRegex.test(val) - }, { - message: t('pages.auth.register.fields.email.validation.invalid', { defaultValue: 'Please enter a valid email format' }), + return z + .object({ + first_name: z.string().min( + 1, + t("pages.auth.register.fields.firstName.validation.required", { + defaultValue: "Please enter your first name", + }) + ), + last_name: z.string().min( + 1, + t("pages.auth.register.fields.lastName.validation.required", { + defaultValue: "Please enter your last name", + }) + ), + email: z + .string() + .min( + 1, + t("pages.auth.register.fields.email.validation.required", { + defaultValue: "Please enter your email", + }) + ) + .transform((val) => val.trim().toLowerCase()) + .refine( + (val) => { + const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/; + return emailRegex.test(val); + }, + { + message: t("pages.auth.register.fields.email.validation.invalid", { + defaultValue: "Please enter a valid email format", + }), + } + ), + phone: z.string().min( + 1, + t("pages.auth.register.fields.phone.validation.required", { + defaultValue: "Please enter your phone number", + }) + ), + password: z + .string() + .min( + 1, + t("pages.auth.register.fields.password.validation.required", { + defaultValue: "Please enter your password", + }) + ) + .min( + 6, + t("pages.auth.register.fields.password.validation.minLength", { + defaultValue: "Password must be at least 6 characters", + }) + ), + confirm_password: z.string().min( + 1, + t("pages.auth.register.fields.confirmPassword.validation.required", { + defaultValue: "Please confirm your password", + }) + ), + }) + .refine((data) => data.password === data.confirm_password, { + message: t("pages.auth.register.fields.confirmPassword.validation.notMatch", { + defaultValue: "Passwords do not match", }), - phone: z - .string() - .min(1, t('pages.auth.register.fields.phone.validation.required', { defaultValue: 'Please enter your phone number' })), - password: z - .string() - .min(1, t('pages.auth.register.fields.password.validation.required', { defaultValue: 'Please enter your password' })) - .min(6, t('pages.auth.register.fields.password.validation.minLength', { defaultValue: 'Password must be at least 6 characters' })), - confirm_password: z - .string() - .min(1, t('pages.auth.register.fields.confirmPassword.validation.required', { defaultValue: 'Please confirm your password' })), - }).refine((data) => data.password === data.confirm_password, { - message: t('pages.auth.register.fields.confirmPassword.validation.notMatch', { defaultValue: 'Passwords do not match' }), - path: ['confirm_password'], - }) - }, [t]) + path: ["confirm_password"], + }); + }, [t]); - const stableDefaultValues = useMemo(() => ({ - first_name: '', - last_name: '', - email: '', - phone: '', - password: '', - confirm_password: '', - }), []) + const stableDefaultValues = useMemo( + () => ({ + first_name: "", + last_name: "", + email: "", + phone: "", + password: "", + confirm_password: "", + }), + [] + ); const form = useForm({ resolver: zodResolver(formSchema), defaultValues: stableDefaultValues, - }) + }); useEffect(() => { try { - form.clearErrors() - const newResolver = zodResolver(formSchema) - - if (form._options && typeof form._options === 'object' && 'resolver' in form._options) { - form._options.resolver = newResolver + form.clearErrors(); + const newResolver = zodResolver(formSchema); + + if (form._options && typeof form._options === "object" && "resolver" in form._options) { + form._options.resolver = newResolver; } - - if ('_resolver' in form && form._resolver !== undefined) { - form._resolver = newResolver + + if ("_resolver" in form && form._resolver !== undefined) { + form._resolver = newResolver; } - + if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } catch (error) { - debugWarn('Failed to update form resolver:', error) - form.clearErrors() + debugWarn("Failed to update form resolver:", error); + form.clearErrors(); if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } - }, [formSchema, form]) + }, [formSchema, form]); - const formMethodsRef = useRef(form) + const formMethodsRef = useRef(form); useEffect(() => { - formMethodsRef.current = form - }, [form]) + formMethodsRef.current = form; + }, [form]); + + const navigateRef = useRef(navigate); + const tRef = useRef(t); + const redirectToRef = useRef(redirectTo); - const navigateRef = useRef(navigate) - const tRef = useRef(t) - const redirectToRef = useRef(redirectTo) - useEffect(() => { - navigateRef.current = navigate - tRef.current = t - redirectToRef.current = redirectTo - }, [navigate, t, redirectTo]) - - const handleSubmit = useCallback(async (formValues) => { - const data = { - first_name: formValues.first_name, - last_name: formValues.last_name, - email: formValues.email, - phone: formValues.phone, - password: formValues.password, - } - - setIsSubmitting(true) - - try { - const result = await handleRegisterRef.current(data) - - if (result.success) { - formMethodsRef.current.reset(stableDefaultValues, { keepValues: false }) - setTimeout(() => { - navigateRef.current(redirectToRef.current, { replace: true }) - }, 0) - } else if (result.requiresEmailVerification) { - const emailToKeep = formValues.email - formMethodsRef.current.reset(stableDefaultValues, { keepValues: false }) - setIsSubmitting(false) - requestAnimationFrame(() => { - navigateRef.current('/auth/verify-email', { - replace: true, - state: { email: emailToKeep }, - }) - }) - } else { - setIsSubmitting(false) + navigateRef.current = navigate; + tRef.current = t; + redirectToRef.current = redirectTo; + }, [navigate, t, redirectTo]); + + const handleSubmit = useCallback( + async (formValues) => { + const data = { + first_name: formValues.first_name, + last_name: formValues.last_name, + email: formValues.email, + phone: formValues.phone, + password: formValues.password, + }; + + setIsSubmitting(true); + + try { + const result = await handleRegisterRef.current(data); + + if (result.success) { + formMethodsRef.current.reset(stableDefaultValues, { keepValues: false }); + setTimeout(() => { + navigateRef.current(redirectToRef.current, { replace: true }); + }, 0); + } else if (result.requiresEmailVerification) { + const emailToKeep = formValues.email; + formMethodsRef.current.reset(stableDefaultValues, { keepValues: false }); + setIsSubmitting(false); + requestAnimationFrame(() => { + navigateRef.current("/auth/verify-email", { + replace: true, + state: { email: emailToKeep }, + }); + }); + } else { + setIsSubmitting(false); + } + } catch (error) { + debugError("Register error:", error); + setIsSubmitting(false); } - } catch (error) { - debugError('Register error:', error) - setIsSubmitting(false) - } - }, [stableDefaultValues]) + }, + [stableDefaultValues] + ); - const onSubmitHandler = useCallback((e) => { - e?.preventDefault?.() - formMethodsRef.current.handleSubmit(handleSubmit)() - }, [handleSubmit]) + const onSubmitHandler = useCallback( + (e) => { + e?.preventDefault?.(); + formMethodsRef.current.handleSubmit(handleSubmit)(); + }, + [handleSubmit] + ); - const handleKeyDown = useCallback((e) => { - if (e.key === 'Enter' && !isSubmitting) { - e.preventDefault() - onSubmitHandler(e) - } - }, [onSubmitHandler, isSubmitting]) + const handleKeyDown = useCallback( + (e) => { + if (e.key === "Enter" && !isSubmitting) { + e.preventDefault(); + onSubmitHandler(e); + } + }, + [onSubmitHandler, isSubmitting] + ); - const firstNameInputRef = useRef(null) + const firstNameInputRef = useRef(null); useEffect(() => { if (!isMobile && firstNameInputRef.current) { - firstNameInputRef.current.focus() + firstNameInputRef.current.focus(); } - }, [isMobile]) + }, [isMobile]); return (
- +
( - {t('pages.auth.register.fields.firstName.label', { defaultValue: 'First Name' })} + + {t("pages.auth.register.fields.firstName.label", { defaultValue: "First Name" })} + - { - firstNameInputRef.current = e - field.ref(e) + firstNameInputRef.current = e; + field.ref(e); }} - disabled={isSubmitting} + disabled={isSubmitting} /> @@ -242,14 +295,11 @@ export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { name="last_name" render={({ field }) => ( - {t('pages.auth.register.fields.lastName.label', { defaultValue: 'Last Name' })} + + {t("pages.auth.register.fields.lastName.label", { defaultValue: "Last Name" })} + - + @@ -261,14 +311,11 @@ export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { name="email" render={({ field }) => ( - {t('pages.auth.register.fields.email.label', { defaultValue: 'Email' })} + + {t("pages.auth.register.fields.email.label", { defaultValue: "Email" })} + - + @@ -279,14 +326,11 @@ export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { name="phone" render={({ field }) => ( - {t('pages.auth.register.fields.phone.label', { defaultValue: 'Phone' })} + + {t("pages.auth.register.fields.phone.label", { defaultValue: "Phone" })} + - + @@ -298,14 +342,16 @@ export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { name="password" render={({ field }) => ( - {t('pages.auth.register.fields.password.label', { defaultValue: 'Password' })} + + {t("pages.auth.register.fields.password.label", { defaultValue: "Password" })} + - @@ -317,14 +363,18 @@ export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { name="confirm_password" render={({ field }) => ( - {t('pages.auth.register.fields.confirmPassword.label', { defaultValue: 'Confirm Password' })} + + {t("pages.auth.register.fields.confirmPassword.label", { + defaultValue: "Confirm Password", + })} + - @@ -332,27 +382,31 @@ export const RegisterForm = ({ className, redirectTo = '/', ...props }) => { )} />
- - - +
- {t('pages.auth.register.links.existingUser', { defaultValue: 'Already have an account? ' })} - + {t("pages.auth.register.links.existingUser", { + defaultValue: "Already have an account? ", + })} + + - {t('pages.auth.register.links.login', { defaultValue: 'Sign in' })} + {t("pages.auth.register.links.login", { defaultValue: "Sign in" })}
- ) -} + ); +}; -export default RegisterForm \ No newline at end of file +export default RegisterForm; diff --git a/frontend/src/components/auth/reset-password-form.jsx b/frontend/src/components/auth/reset-password-form.jsx index 6905431..a0c5d93 100644 --- a/frontend/src/components/auth/reset-password-form.jsx +++ b/frontend/src/components/auth/reset-password-form.jsx @@ -1,12 +1,12 @@ -import React, { useState, useCallback, useMemo, useEffect, useRef } from 'react' -import { useForm } from 'react-hook-form' -import { zodResolver } from '@hookform/resolvers/zod' -import { z } from 'zod' -import { useNavigate, Link } from 'react-router-dom' -import { useTranslation } from 'react-i18next' -import { cn, debugWarn } from '@/lib/utils' -import { Button } from '@/components/ui/button' -import { LogIn } from 'lucide-react' +import React, { useState, useCallback, useMemo, useEffect, useRef } from "react"; +import { useForm } from "react-hook-form"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { z } from "zod"; +import { useNavigate, Link } from "react-router-dom"; +import { useTranslation } from "react-i18next"; +import { cn, debugWarn } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; +import { LogIn } from "lucide-react"; import { Form, FormControl, @@ -14,285 +14,330 @@ import { FormItem, FormLabel, FormMessage, -} from '@/components/ui/form' -import { Input } from '@/components/ui/input' -import { useAuth } from '@/hooks/useAuth' -import { debugError } from '@/lib/utils' -import { Spinner } from '@/components/ui/spinner' -import { useIsMobile } from '@/hooks/useMobile' +} from "@/components/ui/form"; +import { Input } from "@/components/ui/input"; +import { useAuth } from "@/hooks/useAuth"; +import { debugError } from "@/lib/utils"; +import { Spinner } from "@/components/ui/spinner"; +import { useIsMobile } from "@/hooks/useMobile"; const ResetPasswordButton = React.memo(({ onSubmit, t, className, isSubmitting }) => { return ( - - ) -}) + ); +}); -ResetPasswordButton.displayName = 'ResetPasswordButton' +ResetPasswordButton.displayName = "ResetPasswordButton"; export const ResetPasswordForm = ({ className, token, ...props }) => { - const navigate = useNavigate() - const { t } = useTranslation() - const authContext = useAuth() - const isMobile = useIsMobile() - const [isSubmitting, setIsSubmitting] = useState(false) - const [isValidatingToken, setIsValidatingToken] = useState(true) - const [tokenValid, setTokenValid] = useState(false) - - const resetPassword = useMemo(() => authContext.resetPassword, [authContext.resetPassword]) - const validateResetToken = useMemo(() => authContext.validateResetToken, [authContext.validateResetToken]) - const clearError = useMemo(() => authContext.clearError, [authContext.clearError]) - const isResettingPasswordRef = useMemo(() => authContext.isResettingPasswordRef, [authContext.isResettingPasswordRef]) + const navigate = useNavigate(); + const { t } = useTranslation(); + const authContext = useAuth(); + const isMobile = useIsMobile(); + const [isSubmitting, setIsSubmitting] = useState(false); + const [isValidatingToken, setIsValidatingToken] = useState(true); + const [tokenValid, setTokenValid] = useState(false); + + const resetPassword = useMemo(() => authContext.resetPassword, [authContext.resetPassword]); + const validateResetToken = useMemo( + () => authContext.validateResetToken, + [authContext.validateResetToken] + ); + const clearError = useMemo(() => authContext.clearError, [authContext.clearError]); + const isResettingPasswordRef = useMemo( + () => authContext.isResettingPasswordRef, + [authContext.isResettingPasswordRef] + ); useEffect(() => { return () => { if (isResettingPasswordRef) { - isResettingPasswordRef.current = false + isResettingPasswordRef.current = false; } - } - }, [isResettingPasswordRef]) - - const handleResetPassword = useCallback(async (newPassword, resetToken) => { - clearError() - return await resetPassword(newPassword, resetToken) - }, [resetPassword, clearError]) - - const handleResetPasswordRef = useRef(handleResetPassword) + }; + }, [isResettingPasswordRef]); + + const handleResetPassword = useCallback( + async (newPassword, resetToken) => { + clearError(); + return await resetPassword(newPassword, resetToken); + }, + [resetPassword, clearError] + ); + + const handleResetPasswordRef = useRef(handleResetPassword); useEffect(() => { - handleResetPasswordRef.current = handleResetPassword - }, [handleResetPassword]) + handleResetPasswordRef.current = handleResetPassword; + }, [handleResetPassword]); // Track if token validation has been initiated - const validationInitiatedRef = useRef(false) - const lastValidatedTokenRef = useRef(null) + const validationInitiatedRef = useRef(false); + const lastValidatedTokenRef = useRef(null); // Validate token on mount or when token changes useEffect(() => { // Skip if no token if (!token) { - setIsValidatingToken(false) - setTokenValid(false) - validationInitiatedRef.current = false - lastValidatedTokenRef.current = null - return + setIsValidatingToken(false); + setTokenValid(false); + validationInitiatedRef.current = false; + lastValidatedTokenRef.current = null; + return; } // Skip if already validating or already validated this token if (validationInitiatedRef.current && lastValidatedTokenRef.current === token) { - return + return; } // Mark as initiated and store the token being validated - validationInitiatedRef.current = true - lastValidatedTokenRef.current = token + validationInitiatedRef.current = true; + lastValidatedTokenRef.current = token; const validateToken = async () => { - setIsValidatingToken(true) + setIsValidatingToken(true); try { - const result = await validateResetToken(token) + const result = await validateResetToken(token); if (result.success) { - setTokenValid(true) + setTokenValid(true); } else { - setTokenValid(false) - validationInitiatedRef.current = false - lastValidatedTokenRef.current = null + setTokenValid(false); + validationInitiatedRef.current = false; + lastValidatedTokenRef.current = null; } } catch (error) { - debugError('Token validation error:', error) - setTokenValid(false) - validationInitiatedRef.current = false - lastValidatedTokenRef.current = null + debugError("Token validation error:", error); + setTokenValid(false); + validationInitiatedRef.current = false; + lastValidatedTokenRef.current = null; } finally { - setIsValidatingToken(false) + setIsValidatingToken(false); } - } + }; + + validateToken(); + }, [token, validateResetToken, t]); - validateToken() - }, [token, validateResetToken, t]) - const formSchema = useMemo(() => { - return z.object({ - password: z - .string() - .min(1, t('pages.auth.resetPassword.fields.password.validation.required', { defaultValue: 'Please enter your new password' })) - .min(6, t('pages.auth.resetPassword.fields.password.validation.minLength', { defaultValue: 'Password must be at least 6 characters' })), - confirm_password: z - .string() - .min(1, t('pages.auth.resetPassword.fields.confirmPassword.validation.required', { defaultValue: 'Please confirm your new password' })), - }).refine((data) => data.password === data.confirm_password, { - message: t('pages.auth.resetPassword.fields.confirmPassword.validation.notMatch', { defaultValue: 'Passwords do not match' }), - path: ['confirm_password'], - }) - }, [t]) - - const stableDefaultValues = useMemo(() => ({ - password: '', - confirm_password: '', - }), []) + return z + .object({ + password: z + .string() + .min( + 1, + t("pages.auth.resetPassword.fields.password.validation.required", { + defaultValue: "Please enter your new password", + }) + ) + .min( + 6, + t("pages.auth.resetPassword.fields.password.validation.minLength", { + defaultValue: "Password must be at least 6 characters", + }) + ), + confirm_password: z.string().min( + 1, + t("pages.auth.resetPassword.fields.confirmPassword.validation.required", { + defaultValue: "Please confirm your new password", + }) + ), + }) + .refine((data) => data.password === data.confirm_password, { + message: t("pages.auth.resetPassword.fields.confirmPassword.validation.notMatch", { + defaultValue: "Passwords do not match", + }), + path: ["confirm_password"], + }); + }, [t]); + + const stableDefaultValues = useMemo( + () => ({ + password: "", + confirm_password: "", + }), + [] + ); const form = useForm({ resolver: zodResolver(formSchema), defaultValues: stableDefaultValues, - }) + }); useEffect(() => { try { - form.clearErrors() - const newResolver = zodResolver(formSchema) - - if (form._options && typeof form._options === 'object' && 'resolver' in form._options) { - form._options.resolver = newResolver + form.clearErrors(); + const newResolver = zodResolver(formSchema); + + if (form._options && typeof form._options === "object" && "resolver" in form._options) { + form._options.resolver = newResolver; } - - if ('_resolver' in form && form._resolver !== undefined) { - form._resolver = newResolver + + if ("_resolver" in form && form._resolver !== undefined) { + form._resolver = newResolver; } - + if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } catch (error) { - debugWarn('Failed to update form resolver:', error) - form.clearErrors() + debugWarn("Failed to update form resolver:", error); + form.clearErrors(); if (form.formState.isSubmitted) { setTimeout(() => { - form.trigger() - }, 0) + form.trigger(); + }, 0); } } - }, [formSchema, form]) + }, [formSchema, form]); - const formMethodsRef = useRef(form) + const formMethodsRef = useRef(form); useEffect(() => { - formMethodsRef.current = form - }, [form]) + formMethodsRef.current = form; + }, [form]); + + const navigateRef = useRef(navigate); + const tRef = useRef(t); - const navigateRef = useRef(navigate) - const tRef = useRef(t) - useEffect(() => { - navigateRef.current = navigate - tRef.current = t - }, [navigate, t]) - - const handleSubmit = useCallback(async (formValues) => { - if (!token || !tokenValid) { - return - } - - setIsSubmitting(true) - - try { - const result = await handleResetPasswordRef.current(formValues.password, token) - - if (result.success) { - formMethodsRef.current.reset(stableDefaultValues, { keepValues: false }) - navigateRef.current('/', { replace: true }) - } else { + navigateRef.current = navigate; + tRef.current = t; + }, [navigate, t]); + + const handleSubmit = useCallback( + async (formValues) => { + if (!token || !tokenValid) { + return; + } + + setIsSubmitting(true); + + try { + const result = await handleResetPasswordRef.current(formValues.password, token); + + if (result.success) { + formMethodsRef.current.reset(stableDefaultValues, { keepValues: false }); + navigateRef.current("/", { replace: true }); + } else { + // Check if error is 401 (token expired/invalid) + if (result.status === 401) { + setTokenValid(false); + } + setIsSubmitting(false); + } + } catch (error) { + debugError("Reset password error:", error); // Check if error is 401 (token expired/invalid) - if (result.status === 401) { - setTokenValid(false) + if (error.response?.status === 401) { + setTokenValid(false); } - setIsSubmitting(false) + setIsSubmitting(false); } - } catch (error) { - debugError('Reset password error:', error) - // Check if error is 401 (token expired/invalid) - if (error.response?.status === 401) { - setTokenValid(false) - } - setIsSubmitting(false) - } - }, [token, tokenValid, stableDefaultValues]) + }, + [token, tokenValid, stableDefaultValues] + ); - const onSubmitHandler = useCallback((e) => { - e?.preventDefault?.() - if (!tokenValid) { - return - } - formMethodsRef.current.handleSubmit(handleSubmit)() - }, [handleSubmit, tokenValid]) + const onSubmitHandler = useCallback( + (e) => { + e?.preventDefault?.(); + if (!tokenValid) { + return; + } + formMethodsRef.current.handleSubmit(handleSubmit)(); + }, + [handleSubmit, tokenValid] + ); - const handleKeyDown = useCallback((e) => { - if (e.key === 'Enter' && !isSubmitting && tokenValid) { - e.preventDefault() - onSubmitHandler(e) - } - }, [onSubmitHandler, isSubmitting, tokenValid]) + const handleKeyDown = useCallback( + (e) => { + if (e.key === "Enter" && !isSubmitting && tokenValid) { + e.preventDefault(); + onSubmitHandler(e); + } + }, + [onSubmitHandler, isSubmitting, tokenValid] + ); - const passwordInputRef = useRef(null) + const passwordInputRef = useRef(null); useEffect(() => { if (!isMobile && passwordInputRef.current && tokenValid) { - passwordInputRef.current.focus() + passwordInputRef.current.focus(); } - }, [isMobile, tokenValid]) + }, [isMobile, tokenValid]); // Show loading state while validating token if (isValidatingToken) { return ( -
+
- ) + ); } // Show error if token is invalid if (!tokenValid) { return ( -
+

- {t('pages.auth.resetPassword.messages.invalidToken', { defaultValue: 'The reset password link has expired or is invalid. Please request a new one.' })} + {t("pages.auth.resetPassword.messages.invalidToken", { + defaultValue: + "The reset password link has expired or is invalid. Please request a new one.", + })}

- ) + ); } return (
- + ( - {t('pages.auth.resetPassword.fields.password.label', { defaultValue: 'New Password' })} + + {t("pages.auth.resetPassword.fields.password.label", { + defaultValue: "New Password", + })} + - { - passwordInputRef.current = e - field.ref(e) + passwordInputRef.current = e; + field.ref(e); }} - disabled={isSubmitting || !tokenValid} + disabled={isSubmitting || !tokenValid} /> @@ -304,30 +349,34 @@ export const ResetPasswordForm = ({ className, token, ...props }) => { name="confirm_password" render={({ field }) => ( - {t('pages.auth.resetPassword.fields.confirmPassword.label', { defaultValue: 'Confirm New Password' })} + + {t("pages.auth.resetPassword.fields.confirmPassword.label", { + defaultValue: "Confirm New Password", + })} + - )} /> - - - ) -} + ); +}; -export default ResetPasswordForm \ No newline at end of file +export default ResetPasswordForm; diff --git a/frontend/src/components/auth/verify-email-form.jsx b/frontend/src/components/auth/verify-email-form.jsx index bd9db60..6e339bd 100644 --- a/frontend/src/components/auth/verify-email-form.jsx +++ b/frontend/src/components/auth/verify-email-form.jsx @@ -1,223 +1,240 @@ -import React, { useState, useCallback, useEffect, useRef } from 'react' -import { useNavigate, Link, useLocation } from 'react-router-dom' -import { useTranslation } from 'react-i18next' -import { cn } from '@/lib/utils' -import { Button } from '@/components/ui/button' -import { MailCheck, Home, Mail, LogIn } from 'lucide-react' -import { authService } from '@/services/auth.service' -import accountService from '@/services/account.service' -import { useAuth as useAuthContext } from '@/contexts/authContext' -import { useAuth } from '@/hooks/useAuth' -import { getCsrfTokenFromCookie } from '@/lib/cookies' -import { debugError } from '@/lib/utils' -import { Spinner } from '@/components/ui/spinner' +import React, { useState, useCallback, useEffect, useRef } from "react"; +import { useNavigate, Link, useLocation } from "react-router-dom"; +import { useTranslation } from "react-i18next"; +import { cn } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; +import { MailCheck, Home, Mail, LogIn } from "lucide-react"; +import { authService } from "@/services/auth.service"; +import accountService from "@/services/account.service"; +import { useAuth as useAuthContext } from "@/contexts/authContext"; +import { useAuth } from "@/hooks/useAuth"; +import { getCsrfTokenFromCookie } from "@/lib/cookies"; +import { debugError } from "@/lib/utils"; +import { Spinner } from "@/components/ui/spinner"; export const VerifyEmailForm = ({ className, token }) => { - const navigate = useNavigate() - const location = useLocation() - const { t } = useTranslation() - const authContextDirect = useAuthContext() - + const navigate = useNavigate(); + const location = useLocation(); + const { t } = useTranslation(); + const authContextDirect = useAuthContext(); + // Calculate initial state based on token and email const initialState = React.useMemo(() => { - const stateEmail = location.state?.email + const stateEmail = location.state?.email; if (!token) { if (stateEmail) { - return { isVerifying: false, verificationStatus: 'pending', email: stateEmail } + return { isVerifying: false, verificationStatus: "pending", email: stateEmail }; } else { - return { isVerifying: false, verificationStatus: 'error', email: '' } + return { isVerifying: false, verificationStatus: "error", email: "" }; } } - return { isVerifying: true, verificationStatus: null, email: stateEmail || '' } - }, [token, location.state?.email]) - - const [isVerifying, setIsVerifying] = useState(() => initialState.isVerifying) - const [verificationStatus, setVerificationStatus] = useState(() => initialState.verificationStatus) - const [errorMessage, setErrorMessage] = useState('') - const [isResending, setIsResending] = useState(false) - const [cooldownSeconds, setCooldownSeconds] = useState(0) - const [email, setEmail] = useState(() => initialState.email) - const intervalRef = useRef(null) - const verificationInitiatedRef = useRef(false) - const lastVerifiedTokenRef = useRef(null) - const { isAuthenticated, isLoading: isAuthLoading, user } = useAuth() - const { loginSuccess, setToken } = authContextDirect + return { isVerifying: true, verificationStatus: null, email: stateEmail || "" }; + }, [token, location.state?.email]); + + const [isVerifying, setIsVerifying] = useState(() => initialState.isVerifying); + const [verificationStatus, setVerificationStatus] = useState( + () => initialState.verificationStatus + ); + const [errorMessage, setErrorMessage] = useState(""); + const [isResending, setIsResending] = useState(false); + const [cooldownSeconds, setCooldownSeconds] = useState(0); + const [email, setEmail] = useState(() => initialState.email); + const intervalRef = useRef(null); + const verificationInitiatedRef = useRef(false); + const lastVerifiedTokenRef = useRef(null); + const { isAuthenticated, isLoading: isAuthLoading, user } = useAuth(); + const { loginSuccess, setToken } = authContextDirect; // Pending page only (no ?token=): if already logged in, go home. // Skip when user is waiting to confirm a profile email change (pending_email). useEffect(() => { - if (token) return - if (isAuthLoading || !isAuthenticated) return - if (user?.pending_email) return - navigate('/', { replace: true }) - }, [token, isAuthLoading, isAuthenticated, user?.pending_email, navigate]) + if (token) return; + if (isAuthLoading || !isAuthenticated) return; + if (user?.pending_email) return; + navigate("/", { replace: true }); + }, [token, isAuthLoading, isAuthenticated, user?.pending_email, navigate]); // Pending page: detect login completed in another tab (session cookie) via focus / polling useEffect(() => { - if (token) return + if (token) return; - let cancelled = false + let cancelled = false; const checkSession = async () => { - if (cancelled || isAuthenticated) return + if (cancelled || isAuthenticated) return; try { const result = await authService.getToken({ showErrorToast: false, - csrfToken: getCsrfTokenFromCookie() ?? '', - }) + csrfToken: getCsrfTokenFromCookie() ?? "", + }); const accessToken = - typeof result === 'string' + typeof result === "string" ? result - : result?.access_token || result?.data?.access_token || null + : result?.access_token || result?.data?.access_token || null; if (!cancelled && accessToken) { if (loginSuccess) { - loginSuccess(null, accessToken) + loginSuccess(null, accessToken); } else if (setToken) { - setToken(accessToken) + setToken(accessToken); } - navigate('/', { replace: true }) + navigate("/", { replace: true }); } - } catch (error) { + } catch { // Still waiting for verification in another tab } - } + }; const onVisible = () => { - if (document.visibilityState === 'visible') { - checkSession() + if (document.visibilityState === "visible") { + checkSession(); } - } + }; - window.addEventListener('focus', checkSession) - document.addEventListener('visibilitychange', onVisible) - const intervalId = setInterval(checkSession, 4000) + window.addEventListener("focus", checkSession); + document.addEventListener("visibilitychange", onVisible); + const intervalId = setInterval(checkSession, 4000); return () => { - cancelled = true - window.removeEventListener('focus', checkSession) - document.removeEventListener('visibilitychange', onVisible) - clearInterval(intervalId) - } - }, [token, isAuthenticated, loginSuccess, setToken, navigate]) + cancelled = true; + window.removeEventListener("focus", checkSession); + document.removeEventListener("visibilitychange", onVisible); + clearInterval(intervalId); + }; + }, [token, isAuthenticated, loginSuccess, setToken, navigate]); // Handle email verification when token is provided, or set pending state when only email is available useEffect(() => { if (!token) { - setIsVerifying(false) - const stateEmail = location.state?.email + setIsVerifying(false); + const stateEmail = location.state?.email; if (stateEmail) { - setEmail(prev => prev !== stateEmail ? stateEmail : prev) - setVerificationStatus(prev => prev !== 'pending' ? 'pending' : prev) + setEmail((prev) => (prev !== stateEmail ? stateEmail : prev)); + setVerificationStatus((prev) => (prev !== "pending" ? "pending" : prev)); } else { - setVerificationStatus(prev => prev !== 'error' ? 'error' : prev) - setErrorMessage(t('pages.auth.verifyEmail.messages.noToken', { defaultValue: 'No verification token provided' })) + setVerificationStatus((prev) => (prev !== "error" ? "error" : prev)); + setErrorMessage( + t("pages.auth.verifyEmail.messages.noToken", { + defaultValue: "No verification token provided", + }) + ); } - verificationInitiatedRef.current = false - lastVerifiedTokenRef.current = null - return + verificationInitiatedRef.current = false; + lastVerifiedTokenRef.current = null; + return; } if (verificationInitiatedRef.current && lastVerifiedTokenRef.current === token) { - return + return; } - verificationInitiatedRef.current = true - lastVerifiedTokenRef.current = token + verificationInitiatedRef.current = true; + lastVerifiedTokenRef.current = token; const verifyToken = async () => { - setIsVerifying(true) - setVerificationStatus(null) - setErrorMessage('') + setIsVerifying(true); + setVerificationStatus(null); + setErrorMessage(""); try { - const result = await authService.verifyEmail(token, { showErrorToast: false }) + const result = await authService.verifyEmail(token, { showErrorToast: false }); if (result?.user && result?.access_token) { - const { user, access_token } = result - const loginSuccess = authContextDirect.loginSuccess - const setUser = authContextDirect.setUser - + const { user, access_token } = result; + const loginSuccess = authContextDirect.loginSuccess; + const setUser = authContextDirect.setUser; + if (loginSuccess) { - loginSuccess(user, access_token) - await new Promise(resolve => setTimeout(resolve, 50)) + loginSuccess(user, access_token); + await new Promise((resolve) => setTimeout(resolve, 50)); try { - const profileResult = await accountService.getProfile({ showErrorToast: false, showSuccessToast: false }) + const profileResult = await accountService.getProfile({ + showErrorToast: false, + showSuccessToast: false, + }); if (profileResult && setUser) { - setUser(profileResult) + setUser(profileResult); } } catch (error) { - debugError('Failed to fetch user profile:', error) + debugError("Failed to fetch user profile:", error); } } - setIsVerifying(false) - navigate('/', { replace: true }) - return + setIsVerifying(false); + navigate("/", { replace: true }); + return; } else { - setVerificationStatus('error') - setErrorMessage(t('pages.auth.verifyEmail.messages.verificationFailed', { defaultValue: 'Email verification failed' })) - verificationInitiatedRef.current = false - lastVerifiedTokenRef.current = null + setVerificationStatus("error"); + setErrorMessage( + t("pages.auth.verifyEmail.messages.verificationFailed", { + defaultValue: "Email verification failed", + }) + ); + verificationInitiatedRef.current = false; + lastVerifiedTokenRef.current = null; } } catch (error) { - debugError('Email verification error:', error) - setVerificationStatus('error') - const status = error.response?.status + debugError("Email verification error:", error); + setVerificationStatus("error"); + const status = error.response?.status; const errorMsg = status === 401 - ? t('pages.auth.verifyEmail.messages.invalidToken', { - defaultValue: 'The verification link has expired or is invalid. Please request a new one.', + ? t("pages.auth.verifyEmail.messages.invalidToken", { + defaultValue: + "The verification link has expired or is invalid. Please request a new one.", }) : status === 404 - ? t('pages.auth.verifyEmail.messages.userNotFound', { - defaultValue: 'User not found', + ? t("pages.auth.verifyEmail.messages.userNotFound", { + defaultValue: "User not found", }) : status === 409 - ? t('pages.auth.verifyEmail.messages.emailExists', { - defaultValue: 'Email already exists', + ? t("pages.auth.verifyEmail.messages.emailExists", { + defaultValue: "Email already exists", }) - : t('pages.auth.verifyEmail.messages.verificationFailed', { - defaultValue: 'Email verification failed', - }) - setErrorMessage(errorMsg) - verificationInitiatedRef.current = false - lastVerifiedTokenRef.current = null + : t("pages.auth.verifyEmail.messages.verificationFailed", { + defaultValue: "Email verification failed", + }); + setErrorMessage(errorMsg); + verificationInitiatedRef.current = false; + lastVerifiedTokenRef.current = null; } finally { - setIsVerifying(false) + setIsVerifying(false); } - } + }; - verifyToken() + verifyToken(); // eslint-disable-next-line react-hooks/exhaustive-deps - }, [token, t, location.state?.email, navigate]) + }, [token, t, location.state?.email, navigate]); // Fetch email verification cooldown status - const fetchCooldownRef = useRef(null) + const fetchCooldownRef = useRef(null); fetchCooldownRef.current = async (emailToCheck) => { - if (!emailToCheck) return - + if (!emailToCheck) return; + try { - const result = await authService.getEmailVerificationCooldown(emailToCheck) - if (result.status === 'success' && result.data?.data) { - setCooldownSeconds(result.data.data.cooldown_seconds || 0) - } else if (result.status === 'success' && result.data?.cooldown_seconds !== undefined) { - setCooldownSeconds(result.data.cooldown_seconds || 0) + const result = await authService.getEmailVerificationCooldown(emailToCheck); + if (result.status === "success" && result.data?.data) { + setCooldownSeconds(result.data.data.cooldown_seconds || 0); + } else if (result.status === "success" && result.data?.cooldown_seconds !== undefined) { + setCooldownSeconds(result.data.cooldown_seconds || 0); } } catch (error) { - debugError('Failed to fetch cooldown:', error) + debugError("Failed to fetch cooldown:", error); } - } + }; // Update email from location state and fetch cooldown when needed - const cooldownFetchedRef = useRef(false) + const cooldownFetchedRef = useRef(false); useEffect(() => { - const stateEmail = location.state?.email + const stateEmail = location.state?.email; if (stateEmail) { - setEmail(prev => prev !== stateEmail ? stateEmail : prev) - if ((verificationStatus === 'error' || verificationStatus === 'pending') && !cooldownFetchedRef.current) { - cooldownFetchedRef.current = true - fetchCooldownRef.current?.(stateEmail) + setEmail((prev) => (prev !== stateEmail ? stateEmail : prev)); + if ( + (verificationStatus === "error" || verificationStatus === "pending") && + !cooldownFetchedRef.current + ) { + cooldownFetchedRef.current = true; + fetchCooldownRef.current?.(stateEmail); } } else { - cooldownFetchedRef.current = false + cooldownFetchedRef.current = false; } - }, [location.state?.email, verificationStatus]) + }, [location.state?.email, verificationStatus]); // Update cooldown countdown timer useEffect(() => { @@ -226,91 +243,96 @@ export const VerifyEmailForm = ({ className, token }) => { setCooldownSeconds((prev) => { if (prev <= 1) { if (intervalRef.current) { - clearInterval(intervalRef.current) - intervalRef.current = null + clearInterval(intervalRef.current); + intervalRef.current = null; } - return 0 + return 0; } - return prev - 1 - }) - }, 1000) + return prev - 1; + }); + }, 1000); } else { if (intervalRef.current) { - clearInterval(intervalRef.current) - intervalRef.current = null + clearInterval(intervalRef.current); + intervalRef.current = null; } } return () => { if (intervalRef.current) { - clearInterval(intervalRef.current) - intervalRef.current = null + clearInterval(intervalRef.current); + intervalRef.current = null; } - } - }, [cooldownSeconds]) + }; + }, [cooldownSeconds]); // Resend verification email const handleResend = useCallback(async () => { if (!email || cooldownSeconds > 0 || isResending) { - return + return; } - setIsResending(true) + setIsResending(true); try { - await authService.resendVerification(email, { showErrorToast: true, showSuccessToast: true }) - await fetchCooldownRef.current?.(email) + await authService.resendVerification(email, { showErrorToast: true, showSuccessToast: true }); + await fetchCooldownRef.current?.(email); } catch (error) { - debugError('Failed to resend verification email:', error) - await fetchCooldownRef.current?.(email) + debugError("Failed to resend verification email:", error); + await fetchCooldownRef.current?.(email); } finally { - setIsResending(false) + setIsResending(false); } - }, [email, cooldownSeconds, isResending]) + }, [email, cooldownSeconds, isResending]); - const navigateRef = useRef(navigate) + const navigateRef = useRef(navigate); useEffect(() => { - navigateRef.current = navigate - }, [navigate]) + navigateRef.current = navigate; + }, [navigate]); if (isVerifying) { return ( -
+
- ) + ); } - if (verificationStatus === 'success') { + if (verificationStatus === "success") { return ( -
+

- {t('pages.auth.verifyEmail.messages.success', { defaultValue: 'Email verified successfully' })} + {t("pages.auth.verifyEmail.messages.success", { + defaultValue: "Email verified successfully", + })}

- ) + ); } - if (verificationStatus === 'pending' && email) { + if (verificationStatus === "pending" && email) { return ( -
+

- {t('pages.auth.verifyEmail.messages.pending', { defaultValue: 'A verification email has been sent to your email address. Please check your inbox and click the verification link.' })} + {t("pages.auth.verifyEmail.messages.pending", { + defaultValue: + "A verification email has been sent to your email address. Please check your inbox and click the verification link.", + })}

- +
- +
- ) + ); } - if (verificationStatus === 'error') { + if (verificationStatus === "error") { return ( -
+

- {errorMessage || t('pages.auth.verifyEmail.messages.invalidToken', { defaultValue: 'The verification link has expired or is invalid. Please request a new one.' })} + {errorMessage || + t("pages.auth.verifyEmail.messages.invalidToken", { + defaultValue: + "The verification link has expired or is invalid. Please request a new one.", + })}

- + {email && (
)} - +
- ) + ); } - return null -} + return null; +}; -export default VerifyEmailForm \ No newline at end of file +export default VerifyEmailForm; diff --git a/frontend/src/components/core/dock.jsx b/frontend/src/components/core/dock.jsx index 33d94f0..1b5c4bd 100644 --- a/frontend/src/components/core/dock.jsx +++ b/frontend/src/components/core/dock.jsx @@ -16,7 +16,7 @@ export function Dock({ className }) { const [dockState, setDockState] = useState("default"); const [isInitialRender, setIsInitialRender] = useState(true); const { theme, setTheme, themes } = useTheme(); - const { i18n: i18nInstance, t } = useTranslation(); + const { i18n: i18nInstance } = useTranslation(); const dockRef = useRef(null); useEffect(() => { @@ -172,12 +172,7 @@ export function Dock({ className }) { exit="exit" layout > - + @@ -241,12 +231,7 @@ export function Dock({ className }) { exit="exit" layout > - + ); @@ -298,4 +276,4 @@ export function Dock({ className }) { )} ); -} \ No newline at end of file +} diff --git a/frontend/src/components/core/layout.jsx b/frontend/src/components/core/layout.jsx index ed7f2bf..5d5e63a 100644 --- a/frontend/src/components/core/layout.jsx +++ b/frontend/src/components/core/layout.jsx @@ -1,12 +1,12 @@ -import { motion } from 'motion/react'; -import { Dock } from '@/components/core/dock'; -import { cn, debugError } from '@/lib/utils'; -import { useAuth } from '@/hooks/useAuth'; -import { useState, useEffect, useRef } from 'react'; -import { useTranslation } from 'react-i18next'; -import { useIsMobile } from '@/hooks/useMobile'; -import { useNavigate, useLocation } from 'react-router-dom'; -import { Button } from '@/components/ui/button'; +import { motion } from "motion/react"; +import { Dock } from "@/components/core/dock"; +import { cn, debugError } from "@/lib/utils"; +import { useAuth } from "@/hooks/useAuth"; +import { useState, useEffect, useRef } from "react"; +import { useTranslation } from "react-i18next"; +import { useIsMobile } from "@/hooks/useMobile"; +import { useNavigate, useLocation } from "react-router-dom"; +import { Button } from "@/components/ui/button"; import { AlertDialog, AlertDialogAction, @@ -16,15 +16,11 @@ import { AlertDialogFooter, AlertDialogHeader, AlertDialogTitle, -} from '@/components/ui/alert-dialog'; -import { SidebarProvider, SidebarInset, SidebarTrigger } from '@/components/ui/sidebar'; -import { AppSidebar } from '@/components/sidebar/app-sidebar'; +} from "@/components/ui/alert-dialog"; +import { SidebarProvider, SidebarInset, SidebarTrigger } from "@/components/ui/sidebar"; +import { AppSidebar } from "@/components/sidebar/app-sidebar"; -export const Layout = ({ - children, - showDock = true, - dockPosition = 'top-right' -}) => { +export const Layout = ({ children, showDock = true, dockPosition = "top-right" }) => { const { user, isAuthenticated, logout } = useAuth(); const { t } = useTranslation(); const isMobile = useIsMobile(); @@ -33,16 +29,16 @@ export const Layout = ({ const [showLogoutDialog, setShowLogoutDialog] = useState(false); const [shouldDelayLoginButton, setShouldDelayLoginButton] = useState(false); const prevIsAuthenticatedRef = useRef(isAuthenticated); - + // Check if current path is any auth-related page - const isAuthPage = location.pathname.startsWith('/auth'); + const isAuthPage = location.pathname.startsWith("/auth"); const prevIsAuthPageRef = useRef(isAuthPage); const dockPositionClasses = { - 'bottom-right': 'bottom-4 right-4', - 'bottom-left': 'bottom-4 left-4', - 'top-right': 'top-4 right-4', - 'top-left': 'top-4 left-4', + "bottom-right": "bottom-4 right-4", + "bottom-left": "bottom-4 left-4", + "top-right": "top-4 right-4", + "top-left": "top-4 left-4", }; const handleLogout = async () => { @@ -50,36 +46,37 @@ export const Layout = ({ setShowLogoutDialog(false); await logout(); } catch (error) { - debugError('Logout failed:', error); + debugError("Logout failed:", error); setShowLogoutDialog(false); } }; const handleLoginClick = () => { - navigate('/auth/login'); + navigate("/auth/login"); }; - const userName = user ? `${user.first_name || ''} ${user.last_name || ''}`.trim() || user.email || 'User' : 'User'; - const userEmail = user?.email || ''; + const userName = user + ? `${user.first_name || ""} ${user.last_name || ""}`.trim() || user.email || "User" + : "User"; + const userEmail = user?.email || ""; useEffect(() => { let timer = null; - + if (prevIsAuthenticatedRef.current === true && isAuthenticated === false) { setShouldDelayLoginButton(true); timer = setTimeout(() => { setShouldDelayLoginButton(false); }, 600); - } - else if (prevIsAuthPageRef.current === true && isAuthPage === false && !isAuthenticated) { + } else if (prevIsAuthPageRef.current === true && isAuthPage === false && !isAuthenticated) { setShouldDelayLoginButton(true); timer = setTimeout(() => { setShouldDelayLoginButton(false); }, 600); } - + prevIsAuthenticatedRef.current = isAuthenticated; prevIsAuthPageRef.current = isAuthPage; - + return () => { if (timer) { clearTimeout(timer); @@ -91,51 +88,44 @@ export const Layout = ({ setShowLogoutDialog(true); }; - const sidebarUser = user ? { - name: userName, - email: userEmail, - avatar: user.avatar, - first_name: user.first_name, - last_name: user.last_name, - } : null; + const sidebarUser = user + ? { + name: userName, + email: userEmail, + avatar: user.avatar, + first_name: user.first_name, + last_name: user.last_name, + } + : null; if (isAuthenticated && !isAuthPage) { return ( -
+
-
- {children} -
+
{children}
- + {t("pages.auth.logout.title")} {t("pages.auth.logout.confirmMessage")} - + {t("common.actions.cancel")} - -
- {children} -
- {(showDock && (dockPosition === 'top-right' || dockPosition === 'bottom-right')) || !isAuthPage ? ( -
{children}
+ {(showDock && (dockPosition === "top-right" || dockPosition === "bottom-right")) || + !isAuthPage ? ( +
- {showDock && (dockPosition === 'top-right' || dockPosition === 'bottom-right') && ( + {showDock && (dockPosition === "top-right" || dockPosition === "bottom-right") && ( )} {!isAuthPage && ( @@ -189,7 +177,9 @@ export const Layout = ({ className="h-12 px-5 rounded-xl hover:bg-accent hover:text-accent-foreground" onClick={handleLoginClick} > - {t("pages.auth.login.title", { defaultValue: "Sign in" })} + + {t("pages.auth.login.title", { defaultValue: "Sign in" })} + ) : ( @@ -199,20 +189,17 @@ export const Layout = ({ className="h-12 px-5 rounded-xl hover:bg-accent hover:text-accent-foreground" onClick={handleLoginClick} > - {t("pages.auth.login.title", { defaultValue: "Sign in" })} + + {t("pages.auth.login.title", { defaultValue: "Sign in" })} + )} )}
) : null} - {showDock && (dockPosition === 'top-left' || dockPosition === 'bottom-left') && ( -
+ {showDock && (dockPosition === "top-left" || dockPosition === "bottom-left") && ( +
)} @@ -220,4 +207,4 @@ export const Layout = ({ ); }; -export default Layout; \ No newline at end of file +export default Layout; diff --git a/frontend/src/components/core/protected-route.jsx b/frontend/src/components/core/protected-route.jsx index f125c5a..5bb8335 100644 --- a/frontend/src/components/core/protected-route.jsx +++ b/frontend/src/components/core/protected-route.jsx @@ -1,20 +1,28 @@ -import React, { useRef, useEffect } from 'react'; -import { Navigate, useLocation, useNavigate } from 'react-router-dom'; -import { useAuth } from '@/hooks/useAuth'; -import { Spinner } from '@/components/ui/spinner'; -import Error from '@/pages/Error'; -import { debugError } from '@/lib/utils'; +import React, { useRef, useEffect } from "react"; +import { Navigate, useLocation, useNavigate } from "react-router-dom"; +import { useAuth } from "@/hooks/useAuth"; +import { Spinner } from "@/components/ui/spinner"; +import Error from "@/pages/Error"; +import { debugError } from "@/lib/utils"; /** * ProtectedRoute - Protects routes by checking authentication and permissions */ export const ProtectedRoute = ({ children, requireAuth = true, permissions = null }) => { - const { isAuthenticated, isLoading, isLoadingPermissions, checkPermissions, logout, isResettingPasswordRef, user } = useAuth(); + const { + isAuthenticated, + isLoading, + isLoadingPermissions, + checkPermissions, + logout, + isResettingPasswordRef, + user, + } = useAuth(); const location = useLocation(); const navigate = useNavigate(); const isInitialLoadRef = useRef(true); const wasAuthenticatedRef = useRef(isAuthenticated); - + // Track initial load state useEffect(() => { if (!isLoading && isInitialLoadRef.current) { @@ -22,11 +30,11 @@ export const ProtectedRoute = ({ children, requireAuth = true, permissions = nul } }, [isLoading]); - const isResetPasswordPage = location.pathname === '/auth/reset-password'; - const isVerificationPage = location.pathname === '/auth/verify-email'; - const isForgotPasswordPage = location.pathname === '/auth/forgot-password'; + const isResetPasswordPage = location.pathname === "/auth/reset-password"; + const isVerificationPage = location.pathname === "/auth/verify-email"; + const isForgotPasswordPage = location.pathname === "/auth/forgot-password"; const isAuthFlowPage = isResetPasswordPage || isVerificationPage || isForgotPasswordPage; - + // Auto logout authenticated users visiting reset password page (except during password reset) useEffect(() => { if (isResetPasswordPage && isAuthenticated && !isLoading && !isResettingPasswordRef?.current) { @@ -36,8 +44,13 @@ export const ProtectedRoute = ({ children, requireAuth = true, permissions = nul // Redirect to login when user becomes unauthenticated useEffect(() => { - if (wasAuthenticatedRef.current && !isAuthenticated && location.pathname !== '/auth/login' && !isAuthFlowPage) { - navigate('/auth/login', { state: { from: location }, replace: true }); + if ( + wasAuthenticatedRef.current && + !isAuthenticated && + location.pathname !== "/auth/login" && + !isAuthFlowPage + ) { + navigate("/auth/login", { state: { from: location }, replace: true }); } wasAuthenticatedRef.current = isAuthenticated; }, [isAuthenticated, location, navigate, isAuthFlowPage]); @@ -47,15 +60,15 @@ export const ProtectedRoute = ({ children, requireAuth = true, permissions = nul if (!permissions || permissions.length === 0) { return true; } - + if (!isAuthenticated || !checkPermissions) { return false; } - + try { return checkPermissions(permissions); } catch (error) { - debugError('Failed to check permissions:', error); + debugError("Failed to check permissions:", error); return false; } }, [permissions, isAuthenticated, checkPermissions]); @@ -104,7 +117,7 @@ export const ProtectedRoute = ({ children, requireAuth = true, permissions = nul return React.Children.map(children, (child) => { if (React.isValidElement(child)) { return React.cloneElement(child, { - children: + children: , }); } return child; @@ -114,4 +127,4 @@ export const ProtectedRoute = ({ children, requireAuth = true, permissions = nul return children; }; -export default ProtectedRoute; \ No newline at end of file +export default ProtectedRoute; diff --git a/frontend/src/components/data-grid/data-grid-column-filter.jsx b/frontend/src/components/data-grid/data-grid-column-filter.jsx index 973c289..49726bf 100644 --- a/frontend/src/components/data-grid/data-grid-column-filter.jsx +++ b/frontend/src/components/data-grid/data-grid-column-filter.jsx @@ -1,11 +1,11 @@ -import * as React from 'react'; -import { useTranslation } from 'react-i18next'; -import { cn } from '@/lib/utils'; -import { Badge } from '@/components/ui/badge'; -import { Button } from '@/components/ui/button'; -import { Popover, PopoverContent, PopoverTrigger } from '@/components/ui/popover'; -import { Separator } from '@/components/ui/separator'; -import { Check, CirclePlus } from 'lucide-react'; +import * as React from "react"; +import { useTranslation } from "react-i18next"; +import { cn } from "@/lib/utils"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover"; +import { Separator } from "@/components/ui/separator"; +import { Check, CirclePlus } from "lucide-react"; import { Command, CommandEmpty, @@ -14,19 +14,17 @@ import { CommandItem, CommandList, CommandSeparator, -} from '@/components/ui/command'; +} from "@/components/ui/command"; // Column filter component with multi-select dropdown -function DataGridColumnFilter( - { - column, - title, - options, - className, - buttonClassName, - popoverClassName - } -) { +function DataGridColumnFilter({ + column, + title, + options, + className, + buttonClassName, + popoverClassName, +}) { const { t } = useTranslation(); const facets = column?.getFacetedUniqueValues(); const selectedValues = new Set(column?.getFilterValue()); @@ -37,8 +35,10 @@ function DataGridColumnFilter( column.columnDef.meta = {}; } // Only update if not already set or if options have changed - if (!column.columnDef.meta.filterOptions || - JSON.stringify(column.columnDef.meta.filterOptions) !== JSON.stringify(options)) { + if ( + !column.columnDef.meta.filterOptions || + JSON.stringify(column.columnDef.meta.filterOptions) !== JSON.stringify(options) + ) { column.columnDef.meta.filterOptions = options; } } @@ -62,13 +62,20 @@ function DataGridColumnFilter( {selectedValues?.size > 0 && ( <> - + {selectedValues.size}
{selectedValues.size > 2 ? ( - - {selectedValues.size} {t('components.dataGrid.columnFilter.selected', 'selected')} + + {selectedValues.size}{" "} + {t("components.dataGrid.columnFilter.selected", "selected")} ) : ( options @@ -77,7 +84,8 @@ function DataGridColumnFilter( + className="rounded-sm px-2 font-semibold bg-primary/15 text-primary border-primary/50" + > {option.label} )) @@ -91,7 +99,9 @@ function DataGridColumnFilter( - {t('components.dataGrid.columnFilter.noResults', 'No results found')} + + {t("components.dataGrid.columnFilter.noResults", "No results found")} + {options.map((option) => { const isSelected = selectedValues.has(option.value); @@ -106,18 +116,19 @@ function DataGridColumnFilter( } const filterValues = Array.from(selectedValues); column?.setFilterValue(filterValues.length ? filterValues : undefined); - }}> + }} + >
- + "me-2 flex h-4 w-4 items-center justify-center rounded border border-muted-foreground", + isSelected ? "bg-primary border-primary" : "opacity-50 [&_svg]:invisible" + )} + > +
- {option.label} + {option.label} {facets?.get(option.value) && ( - + {facets.get(option.value)} )} @@ -131,8 +142,9 @@ function DataGridColumnFilter( column?.setFilterValue(undefined)} - className="justify-center text-center"> - {t('components.dataGrid.columnFilter.clearFilters', 'Clear filters')} + className="justify-center text-center" + > + {t("components.dataGrid.columnFilter.clearFilters", "Clear filters")} @@ -145,4 +157,3 @@ function DataGridColumnFilter( } export { DataGridColumnFilter }; - diff --git a/frontend/src/components/data-grid/data-grid-column-header.jsx b/frontend/src/components/data-grid/data-grid-column-header.jsx index 95998ac..93dc574 100644 --- a/frontend/src/components/data-grid/data-grid-column-header.jsx +++ b/frontend/src/components/data-grid/data-grid-column-header.jsx @@ -1,8 +1,8 @@ -import * as React from 'react'; -import { useTranslation } from 'react-i18next'; -import { cn } from '@/lib/utils'; -import { useDataGrid } from '@/components/data-grid/data-grid'; -import { Tooltip, TooltipTrigger, TooltipContent } from '@/components/ui/tooltip'; +import * as React from "react"; +import { useTranslation } from "react-i18next"; +import { cn } from "@/lib/utils"; +import { useDataGrid } from "@/components/data-grid/data-grid"; +import { Tooltip, TooltipTrigger, TooltipContent } from "@/components/ui/tooltip"; import { DropdownMenu, DropdownMenuCheckboxItem, @@ -15,7 +15,7 @@ import { DropdownMenuSubContent, DropdownMenuSubTrigger, DropdownMenuTrigger, -} from '@/components/ui/dropdown-menu'; +} from "@/components/ui/dropdown-menu"; import { ArrowDown, ArrowLeft, @@ -27,21 +27,12 @@ import { ChevronsUpDown, Settings2, ChevronDown, - Filter -} from 'lucide-react'; - -function DataGridColumnHeader( - { - column, - title = '', - icon, - className, - filter, - visibility = false - } -) { + Filter, +} from "lucide-react"; + +function DataGridColumnHeader({ column, title = "", icon, className, filter, visibility = false }) { const { t } = useTranslation(); - const { isLoading, table, props, recordCount } = useDataGrid(); + const { table, props } = useDataGrid(); if (title && (!column.columnDef.meta || column.columnDef.meta.headerTitle !== title)) { if (!column.columnDef.meta) { @@ -54,14 +45,14 @@ function DataGridColumnHeader( const currentOrder = [...table.getState().columnOrder]; const currentIndex = currentOrder.indexOf(column.id); - if (direction === 'left' && currentIndex > 0) { + if (direction === "left" && currentIndex > 0) { const newOrder = [...currentOrder]; const [movedColumn] = newOrder.splice(currentIndex, 1); newOrder.splice(currentIndex - 1, 0, movedColumn); table.setColumnOrder(newOrder); } - if (direction === 'right' && currentIndex < currentOrder.length - 1) { + if (direction === "right" && currentIndex < currentOrder.length - 1) { const newOrder = [...currentOrder]; const [movedColumn] = newOrder.splice(currentIndex, 1); newOrder.splice(currentIndex + 1, 0, movedColumn); @@ -69,10 +60,10 @@ function DataGridColumnHeader( } }; - const canMove = direction => { + const canMove = (direction) => { const currentOrder = table.getState().columnOrder; const currentIndex = currentOrder.indexOf(column.id); - if (direction === 'left') { + if (direction === "left") { return currentIndex > 0; } else { return currentIndex < currentOrder.length - 1; @@ -84,74 +75,76 @@ function DataGridColumnHeader( const filterValue = column.getFilterValue(); if (filterValue === undefined || filterValue === null) return false; if (Array.isArray(filterValue)) return filterValue.length > 0; - if (typeof filterValue === 'string') return filterValue.trim().length > 0; + if (typeof filterValue === "string") return filterValue.trim().length > 0; return Boolean(filterValue); }, [column, columnFilters]); const getFilterTooltipContent = React.useMemo(() => { if (!hasFilter) return null; - + const filterValue = column.getFilterValue(); let filterOptions = column.columnDef.meta?.filterOptions; - + if (!filterOptions || !Array.isArray(filterOptions)) { filterOptions = null; } - + if (!filterValue) return null; - + const normalizeValue = (val) => { - if (val === true) return 'true'; - if (val === false) return 'false'; - if (val === 'true' || val === 'True' || val === 'TRUE') return 'true'; - if (val === 'false' || val === 'False' || val === 'FALSE') return 'false'; - if (val === 1 || val === '1') return 'true'; - if (val === 0 || val === '0') return 'false'; + if (val === true) return "true"; + if (val === false) return "false"; + if (val === "true" || val === "True" || val === "TRUE") return "true"; + if (val === "false" || val === "False" || val === "FALSE") return "false"; + if (val === 1 || val === "1") return "true"; + if (val === 0 || val === "0") return "false"; return String(val); }; - + const findOptionLabel = (value) => { if (!filterOptions || !Array.isArray(filterOptions) || filterOptions.length === 0) { return null; } - + const normalizedValue = normalizeValue(value); - + for (const opt of filterOptions) { - if (!opt || typeof opt !== 'object') continue; - + if (!opt || typeof opt !== "object") continue; + const normalizedOptValue = normalizeValue(opt.value); if (normalizedOptValue === normalizedValue && opt.label) { return opt.label; } } - + return null; }; - + if (Array.isArray(filterValue)) { if (filterOptions && Array.isArray(filterOptions) && filterOptions.length > 0) { const selectedLabels = filterValue - .map(value => findOptionLabel(value)) - .filter(label => label !== null && label !== undefined); - + .map((value) => findOptionLabel(value)) + .filter((label) => label !== null && label !== undefined); + if (selectedLabels.length > 0) { - return selectedLabels.join(', '); + return selectedLabels.join(", "); } } - return filterValue.map(v => { - const label = findOptionLabel(v); - return label || String(v); - }).join(', '); + return filterValue + .map((v) => { + const label = findOptionLabel(v); + return label || String(v); + }) + .join(", "); } - + if (filterOptions && Array.isArray(filterOptions) && filterOptions.length > 0) { const label = findOptionLabel(filterValue); if (label) { return label; } } - + return String(filterValue); }, [hasFilter, column, columnFilters]); @@ -159,9 +152,10 @@ function DataGridColumnHeader( return (
+ )} + > {icon && icon} {title}
@@ -172,15 +166,16 @@ function DataGridColumnHeader( return (
+ )} + > {icon && {icon}} {title} {column.getCanSort() && - (column.getIsSorted() === 'desc' ? ( + (column.getIsSorted() === "desc" ? ( - ) : column.getIsSorted() === 'asc' ? ( + ) : column.getIsSorted() === "asc" ? ( ) : ( @@ -200,17 +195,25 @@ function DataGridColumnHeader( getFilterTooltipContent ? ( - + - -
-
{t('components.dataGrid.columnFilter.activeFilter', 'Active filter')}
-
{getFilterTooltipContent}
-
-
+ +
+
+ {t("components.dataGrid.columnFilter.activeFilter", "Active filter")} +
+
{getFilterTooltipContent}
+
+
) : ( - + ) ) : ( @@ -220,56 +223,79 @@ function DataGridColumnHeader( {filter && {filter}} - {filter && (column.getCanSort() || column.getCanPin() || visibility) && } + {filter && (column.getCanSort() || column.getCanPin() || visibility) && ( + + )} {column.getCanSort() && ( <> { - if (column.getIsSorted() === 'asc') { + if (column.getIsSorted() === "asc") { column.clearSorting(); } else { column.toggleSorting(false); } }} - disabled={!column.getCanSort()}> + disabled={!column.getCanSort()} + > - {t('components.dataGrid.columnHeader.asc', 'Ascending')} - {column.getIsSorted() === 'asc' && } + + {t("components.dataGrid.columnHeader.asc", "Ascending")} + + {column.getIsSorted() === "asc" && ( + + )} { - if (column.getIsSorted() === 'desc') { + if (column.getIsSorted() === "desc") { column.clearSorting(); } else { column.toggleSorting(true); } }} - disabled={!column.getCanSort()}> + disabled={!column.getCanSort()} + > - {t('components.dataGrid.columnHeader.desc', 'Descending')} - {column.getIsSorted() === 'desc' && } + + {t("components.dataGrid.columnHeader.desc", "Descending")} + + {column.getIsSorted() === "desc" && ( + + )} )} - {(filter || column.getCanSort()) && (column.getCanSort() || column.getCanPin() || visibility) && ( - - )} + {(filter || column.getCanSort()) && + (column.getCanSort() || column.getCanPin() || visibility) && ( + + )} {props.tableLayout?.columnsPinnable && column.getCanPin() && ( <> column.pin(column.getIsPinned() === 'left' ? false : 'left')}> + onClick={() => column.pin(column.getIsPinned() === "left" ? false : "left")} + > column.pin(column.getIsPinned() === 'right' ? false : 'right')}> + onClick={() => column.pin(column.getIsPinned() === "right" ? false : "right")} + > )} @@ -278,16 +304,18 @@ function DataGridColumnHeader( <> moveColumn('left')} - disabled={!canMove('left') || column.getIsPinned() !== false}> + onClick={() => moveColumn("left")} + disabled={!canMove("left") || column.getIsPinned() !== false} + > moveColumn('right')} - disabled={!canMove('right') || column.getIsPinned() !== false}> + onClick={() => moveColumn("right")} + disabled={!canMove("right") || column.getIsPinned() !== false} + > )} @@ -300,13 +328,13 @@ function DataGridColumnHeader( - {t('components.dataGrid.columnHeader.columns', 'Columns')} + {t("components.dataGrid.columnHeader.columns", "Columns")} {table .getAllColumns() - .filter((col) => typeof col.accessorFn !== 'undefined' && col.getCanHide()) + .filter((col) => typeof col.accessorFn !== "undefined" && col.getCanHide()) .map((col) => { return ( event.preventDefault()} onCheckedChange={(value) => col.toggleVisibility(!!value)} - className="capitalize"> + className="capitalize" + > {col.columnDef.meta?.headerTitle || col.id} ); @@ -348,17 +377,25 @@ function DataGridColumnHeader( getFilterTooltipContent ? ( - +
-
{t('components.dataGrid.columnFilter.activeFilter', 'Active filter')}
+
+ {t("components.dataGrid.columnFilter.activeFilter", "Active filter")} +
{getFilterTooltipContent}
) : ( - + ) ) : ( @@ -370,29 +407,35 @@ function DataGridColumnHeader( <> { - if (column.getIsSorted() === 'asc') { + if (column.getIsSorted() === "asc") { column.clearSorting(); } else { column.toggleSorting(false); } }} - disabled={!column.getCanSort()}> + disabled={!column.getCanSort()} + > - {t('components.dataGrid.columnHeader.asc')} - {column.getIsSorted() === 'asc' && } + {t("components.dataGrid.columnHeader.asc")} + {column.getIsSorted() === "asc" && ( + + )} { - if (column.getIsSorted() === 'desc') { + if (column.getIsSorted() === "desc") { column.clearSorting(); } else { column.toggleSorting(true); } }} - disabled={!column.getCanSort()}> + disabled={!column.getCanSort()} + > - {t('components.dataGrid.columnHeader.desc')} - {column.getIsSorted() === 'desc' && } + {t("components.dataGrid.columnHeader.desc")} + {column.getIsSorted() === "desc" && ( + + )} )} @@ -405,4 +448,3 @@ function DataGridColumnHeader( } export { DataGridColumnHeader }; - diff --git a/frontend/src/components/data-grid/data-grid-column-visibility.jsx b/frontend/src/components/data-grid/data-grid-column-visibility.jsx index 24945c6..5063273 100644 --- a/frontend/src/components/data-grid/data-grid-column-visibility.jsx +++ b/frontend/src/components/data-grid/data-grid-column-visibility.jsx @@ -1,19 +1,14 @@ -import { useTranslation } from 'react-i18next'; +import { useTranslation } from "react-i18next"; import { DropdownMenu, DropdownMenuCheckboxItem, DropdownMenuContent, DropdownMenuLabel, DropdownMenuTrigger, -} from '@/components/ui/dropdown-menu'; +} from "@/components/ui/dropdown-menu"; // Column visibility toggle dropdown -function DataGridColumnVisibility( - { - table, - trigger - } -) { +function DataGridColumnVisibility({ table, trigger }) { const { t } = useTranslation(); // Get display title for column @@ -22,7 +17,7 @@ function DataGridColumnVisibility( return column.columnDef.meta.headerTitle; } - if (typeof column.columnDef.header === 'string') { + if (typeof column.columnDef.header === "string") { return column.columnDef.header; } @@ -33,12 +28,12 @@ function DataGridColumnVisibility( {trigger} - {t('components.dataGrid.columnVisibility.toggleColumns', 'Toggle columns')} + + {t("components.dataGrid.columnVisibility.toggleColumns", "Toggle columns")} + {table .getAllColumns() - .filter( - (column) => typeof column.accessorFn !== 'undefined' && column.getCanHide() - ) + .filter((column) => typeof column.accessorFn !== "undefined" && column.getCanHide()) .map((column) => { const title = getColumnTitle(column); return ( @@ -47,7 +42,8 @@ function DataGridColumnVisibility( className="capitalize" checked={column.getIsVisible()} onSelect={(event) => event.preventDefault()} - onCheckedChange={(value) => column.toggleVisibility(!!value)}> + onCheckedChange={(value) => column.toggleVisibility(!!value)} + > {title} ); @@ -58,4 +54,3 @@ function DataGridColumnVisibility( } export { DataGridColumnVisibility }; - diff --git a/frontend/src/components/data-grid/data-grid-pagination.jsx b/frontend/src/components/data-grid/data-grid-pagination.jsx index f1f82e6..c0e162c 100644 --- a/frontend/src/components/data-grid/data-grid-pagination.jsx +++ b/frontend/src/components/data-grid/data-grid-pagination.jsx @@ -1,8 +1,14 @@ -import { useTranslation } from 'react-i18next'; -import { useDataGrid } from '@/components/data-grid/data-grid'; -import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; -import { Skeleton } from '@/components/ui/skeleton'; -import { Separator } from '@/components/ui/separator'; +import { useTranslation } from "react-i18next"; +import { useDataGrid } from "@/components/data-grid/data-grid"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { Skeleton } from "@/components/ui/skeleton"; +import { Separator } from "@/components/ui/separator"; import { Pagination, PaginationContent, @@ -11,8 +17,8 @@ import { PaginationLink, PaginationNext, PaginationPrevious, -} from '@/components/ui/pagination'; -import { cn } from '@/lib/utils'; +} from "@/components/ui/pagination"; +import { cn } from "@/lib/utils"; // Pagination component with page navigation and rows per page selector function DataGridPagination(props) { @@ -21,17 +27,17 @@ function DataGridPagination(props) { const defaultProps = { sizes: [5, 10, 25, 50, 100], - sizesLabel: t('components.pagination.show', 'Show'), - sizesDescription: t('components.pagination.perPage', 'per page'), + sizesLabel: t("components.pagination.show", "Show"), + sizesDescription: t("components.pagination.perPage", "per page"), sizesSkeleton: , moreLimit: 5, more: false, - info: 'components.pagination.info', + info: "components.pagination.info", infoSkeleton: , - rowsPerPageLabel: t('components.pagination.rowsPerPage', 'Rows per page'), - previousPageLabel: t('components.pagination.previousPage', 'Previous page'), - nextPageLabel: t('components.pagination.nextPage', 'Next page'), - ellipsisText: '...', + rowsPerPageLabel: t("components.pagination.rowsPerPage", "Rows per page"), + previousPageLabel: t("components.pagination.previousPage", "Previous page"), + nextPageLabel: t("components.pagination.nextPage", "Next page"), + ellipsisText: "...", }; const mergedProps = { ...defaultProps, ...props }; @@ -44,10 +50,10 @@ function DataGridPagination(props) { // Generate page numbers with ellipsis for pagination const generatePageNumbers = () => { if (pageCount <= 1) return []; - + const pages = []; const totalPages = pageCount; - + // If total pages <= 7, show all pages if (totalPages <= 7) { for (let i = 1; i <= totalPages; i++) { @@ -55,35 +61,35 @@ function DataGridPagination(props) { } return pages; } - + // Always show first page pages.push(1); - + // Current page is near the beginning (pages 1-4) if (currentPage <= 4) { for (let i = 2; i <= 5; i++) { pages.push(i); } - pages.push('ellipsis-end'); + pages.push("ellipsis-end"); pages.push(totalPages); } // Current page is near the end (last 4 pages) else if (currentPage >= totalPages - 3) { - pages.push('ellipsis-start'); + pages.push("ellipsis-start"); for (let i = totalPages - 4; i <= totalPages; i++) { pages.push(i); } } // Current page is in the middle else { - pages.push('ellipsis-start'); + pages.push("ellipsis-start"); pages.push(currentPage - 1); pages.push(currentPage); pages.push(currentPage + 1); - pages.push('ellipsis-end'); + pages.push("ellipsis-end"); pages.push(totalPages); } - + return pages; }; @@ -92,14 +98,14 @@ function DataGridPagination(props) { // Render pagination buttons from page numbers const renderPageButtons = () => { return pageNumbers.map((page, index) => { - if (page === 'ellipsis-start' || page === 'ellipsis-end') { + if (page === "ellipsis-start" || page === "ellipsis-end") { return ( ); } - + const pageIndexValue = page - 1; return ( @@ -111,7 +117,11 @@ function DataGridPagination(props) { table.setPageIndex(pageIndexValue); } }} - className={cn('h-8 w-8 p-0 text-sm rounded-xs', pageIndex === pageIndexValue ? '' : 'text-muted-foreground')}> + className={cn( + "h-8 w-8 p-0 text-sm rounded-xs", + pageIndex === pageIndexValue ? "" : "text-muted-foreground" + )} + > {page} @@ -123,21 +133,26 @@ function DataGridPagination(props) {
-
+ )} + > +
{isLoading ? ( mergedProps?.sizesSkeleton ) : ( <> {recordCount !== undefined && recordCount !== null && (
- {t('components.pagination.totalAccounts', 'Total {{count}} accounts', { count: recordCount.toLocaleString() })} + {t("components.pagination.totalAccounts", "Total {{count}} accounts", { + count: recordCount.toLocaleString(), + })}
)} - +
{mergedProps.rowsPerPageLabel}
+ aria-label={t("common.actions.reset", "Reset")} + > )}
-
@@ -140,33 +141,27 @@ function DataGridToolbar({ const shouldShow = hasSelectedRows && onDeleteSelected; return shouldShow ? true : null; })() && ( - )} {showAdd && onAdd && ( )} {showSelect && onSelectionToggle && ( )} @@ -176,7 +171,7 @@ function DataGridToolbar({ trigger={ } /> @@ -187,4 +182,3 @@ function DataGridToolbar({ } export { DataGridToolbar }; - diff --git a/frontend/src/components/data-grid/data-grid.jsx b/frontend/src/components/data-grid/data-grid.jsx index 05110fa..0bf1b3b 100644 --- a/frontend/src/components/data-grid/data-grid.jsx +++ b/frontend/src/components/data-grid/data-grid.jsx @@ -1,6 +1,6 @@ -import * as React from 'react'; -import { createContext, useContext } from 'react'; -import { cn } from '@/lib/utils'; +import * as React from "react"; +import { createContext, useContext } from "react"; +import { cn } from "@/lib/utils"; const DataGridContext = createContext(undefined); @@ -8,19 +8,13 @@ const DataGridContext = createContext(undefined); function useDataGrid() { const context = useContext(DataGridContext); if (!context) { - throw new Error('useDataGrid must be used within a DataGridProvider'); + throw new Error("useDataGrid must be used within a DataGridProvider"); } return context; } // Context provider for DataGrid state -function DataGridProvider( - { - children, - table, - ...props - } -) { +function DataGridProvider({ children, table, ...props }) { const [hasVerticalScrollbar, setHasVerticalScrollbar] = React.useState(false); const [delayedLoading, setDelayedLoading] = React.useState(false); const rawIsLoading = props.isLoading || false; @@ -58,22 +52,17 @@ function DataGridProvider( rawIsLoading, hasVerticalScrollbar, setHasVerticalScrollbar, - }}> + }} + > {children} ); } // Main DataGrid component with default props -function DataGrid( - { - children, - table, - ...props - } -) { +function DataGrid({ children, table, ...props }) { const defaultProps = { - loadingMode: 'skeleton', + loadingMode: "skeleton", loadingDelayMs: 0, tableLayout: { dense: false, @@ -84,7 +73,7 @@ function DataGrid( headerSticky: false, headerBackground: true, headerBorder: true, - width: 'fixed', + width: "fixed", columnsVisibility: false, columnsResizable: false, columnsPinnable: false, @@ -93,14 +82,14 @@ function DataGrid( rowsDraggable: false, }, tableClassNames: { - base: '', - header: '', - headerRow: '', - headerSticky: 'sticky top-0 z-10 bg-background/90 backdrop-blur-xs', - body: '', - bodyRow: '', - footer: '', - edgeCell: '', + base: "", + header: "", + headerRow: "", + headerSticky: "sticky top-0 z-10 bg-background/90 backdrop-blur-xs", + body: "", + bodyRow: "", + footer: "", + edgeCell: "", }, }; @@ -129,19 +118,19 @@ function DataGrid( } // Container wrapper for DataGrid with border styling -function DataGridContainer({ - children, - className, - border = true -}) { +function DataGridContainer({ children, className, border = true }) { return (
+ className={cn( + "grid w-full", + border && "border border-border rounded-xl overflow-hidden shadow-2xs", + className + )} + > {children}
); } export { useDataGrid, DataGridProvider, DataGrid, DataGridContainer }; - diff --git a/frontend/src/components/profile/change-password-form.jsx b/frontend/src/components/profile/change-password-form.jsx index 975a94f..6fe2411 100644 --- a/frontend/src/components/profile/change-password-form.jsx +++ b/frontend/src/components/profile/change-password-form.jsx @@ -1,12 +1,12 @@ -import { useState, useCallback, useMemo, useEffect } from 'react'; -import { useForm } from 'react-hook-form'; -import { zodResolver } from '@hookform/resolvers/zod'; -import { z } from 'zod'; -import { useTranslation } from 'react-i18next'; -import { toast } from 'sonner'; -import { Check } from 'lucide-react'; -import { cn, debugError } from '@/lib/utils'; -import { Button } from '@/components/ui/button'; +import { useState, useCallback, useMemo, useEffect } from "react"; +import { useForm } from "react-hook-form"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { z } from "zod"; +import { useTranslation } from "react-i18next"; +import { toast } from "sonner"; +import { Check } from "lucide-react"; +import { cn, debugError } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; import { Form, FormControl, @@ -14,17 +14,15 @@ import { FormItem, FormLabel, FormMessage, -} from '@/components/ui/form'; -import { Input } from '@/components/ui/input'; -import { Spinner } from '@/components/ui/spinner'; -import { Progress } from '@/components/ui/progress'; -import { useAuth } from '@/hooks/useAuth'; -import { useIsMobile } from '@/hooks/useMobile'; +} from "@/components/ui/form"; +import { Input } from "@/components/ui/input"; +import { Spinner } from "@/components/ui/spinner"; +import { Progress } from "@/components/ui/progress"; +import { useAuth } from "@/hooks/useAuth"; export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { const { t } = useTranslation(); const { changePassword, isLoading } = useAuth(); - const isMobile = useIsMobile(); const [isSubmitting, setIsSubmitting] = useState(false); const [passwordStrength, setPasswordStrength] = useState(0); @@ -35,29 +33,38 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { }, [isSubmitting, isLoading, onSubmittingChange]); const formSchema = useMemo(() => { - return z.object({ - current_password: z - .string() - .min(1, { message: t('pages.profile.security.fields.currentPassword.validation.required') }) - .min(6, { message: t('pages.profile.security.fields.currentPassword.validation.minLength') }), - new_password: z - .string() - .min(1, { message: t('pages.profile.security.fields.newPassword.validation.required') }) - .min(6, { message: t('pages.profile.security.fields.newPassword.validation.minLength') }), - confirm_new_password: z - .string() - .min(1, { message: t('pages.profile.security.fields.confirmNewPassword.validation.required') }), - }).refine((data) => data.new_password === data.confirm_new_password, { - message: t('pages.profile.security.fields.confirmNewPassword.validation.notMatch'), - path: ['confirm_new_password'], - }); + return z + .object({ + current_password: z + .string() + .min(1, { + message: t("pages.profile.security.fields.currentPassword.validation.required"), + }) + .min(6, { + message: t("pages.profile.security.fields.currentPassword.validation.minLength"), + }), + new_password: z + .string() + .min(1, { message: t("pages.profile.security.fields.newPassword.validation.required") }) + .min(6, { message: t("pages.profile.security.fields.newPassword.validation.minLength") }), + confirm_new_password: z.string().min(1, { + message: t("pages.profile.security.fields.confirmNewPassword.validation.required"), + }), + }) + .refine((data) => data.new_password === data.confirm_new_password, { + message: t("pages.profile.security.fields.confirmNewPassword.validation.notMatch"), + path: ["confirm_new_password"], + }); }, [t]); - const defaultValues = useMemo(() => ({ - current_password: '', - new_password: '', - confirm_new_password: '', - }), []); + const defaultValues = useMemo( + () => ({ + current_password: "", + new_password: "", + confirm_new_password: "", + }), + [] + ); const checkPasswordStrength = useCallback((password) => { let strength = 0; @@ -71,29 +78,29 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { const strengthConfig = useMemo(() => { if (passwordStrength < 50) { return { - label: t('pages.profile.security.fields.newPassword.strength.weak'), - textColor: 'text-destructive', - barColor: '[&>div]:bg-destructive dark:[&>div]:bg-destructive', + label: t("pages.profile.security.fields.newPassword.strength.weak"), + textColor: "text-destructive", + barColor: "[&>div]:bg-destructive dark:[&>div]:bg-destructive", }; } if (passwordStrength < 75) { return { - label: t('pages.profile.security.fields.newPassword.strength.medium'), - textColor: 'text-warning', - barColor: '[&>div]:bg-warning dark:[&>div]:bg-warning', + label: t("pages.profile.security.fields.newPassword.strength.medium"), + textColor: "text-warning", + barColor: "[&>div]:bg-warning dark:[&>div]:bg-warning", }; } if (passwordStrength < 100) { return { - label: t('pages.profile.security.fields.newPassword.strength.strong'), - textColor: 'text-success', - barColor: '[&>div]:bg-success dark:[&>div]:bg-success', + label: t("pages.profile.security.fields.newPassword.strength.strong"), + textColor: "text-success", + barColor: "[&>div]:bg-success dark:[&>div]:bg-success", }; } return { - label: t('pages.profile.security.fields.newPassword.strength.veryStrong'), - textColor: 'text-success', - barColor: '[&>div]:bg-success dark:[&>div]:bg-success', + label: t("pages.profile.security.fields.newPassword.strength.veryStrong"), + textColor: "text-success", + barColor: "[&>div]:bg-success dark:[&>div]:bg-success", }; }, [passwordStrength, t]); @@ -102,83 +109,94 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { defaultValues, }); - const submitPasswordChange = useCallback(async (formValues) => { - setIsSubmitting(true); - - try { - const result = await changePassword({ - current_password: formValues.current_password, - new_password: formValues.new_password, - logout_other_devices: true, - }); - - if (result.success) { - form.reset(defaultValues); - setPasswordStrength(0); - if (onSuccess) { - await onSuccess(); + const submitPasswordChange = useCallback( + async (formValues) => { + setIsSubmitting(true); + + try { + const result = await changePassword({ + current_password: formValues.current_password, + new_password: formValues.new_password, + logout_other_devices: true, + }); + + if (result.success) { + form.reset(defaultValues); + setPasswordStrength(0); + if (onSuccess) { + await onSuccess(); + } + } else { + // Check if it's a password error (401 with password error message) + const error = result.error; + const errorMessage = typeof error === "string" ? error : error?.message || ""; + const isPasswordError = + errorMessage.includes("Current password is incorrect") || + errorMessage.includes("password is incorrect") || + errorMessage.toLowerCase().includes("incorrect password"); + + if (isPasswordError) { + // Set error on current_password field + form.setError("current_password", { + type: "manual", + message: t("pages.profile.security.fields.currentPassword.validation.incorrect"), + }); + } else { + toast.error(errorMessage || t("pages.profile.security.messages.incorrect")); + } } - } else { - // Check if it's a password error (401 with password error message) - const error = result.error; - const errorMessage = typeof error === 'string' ? error : error?.message || ''; - const isPasswordError = errorMessage.includes('Current password is incorrect') || - errorMessage.includes('password is incorrect') || - errorMessage.toLowerCase().includes('incorrect password'); - - if (isPasswordError) { - // Set error on current_password field - form.setError('current_password', { - type: 'manual', - message: t('pages.profile.security.fields.currentPassword.validation.incorrect'), + } catch (error) { + debugError("Change password error:", error); + + if (error.isPasswordError) { + const errorMessage = t( + "pages.profile.security.fields.currentPassword.validation.incorrect" + ); + const fieldName = error.passwordErrorField || "current_password"; + + form.setError(fieldName, { + type: "manual", + message: errorMessage, }); - } else { - toast.error(errorMessage || t('pages.profile.security.messages.incorrect')); } + } finally { + setIsSubmitting(false); } - } catch (error) { - debugError('Change password error:', error); - - if (error.isPasswordError) { - const errorMessage = t('pages.profile.security.fields.currentPassword.validation.incorrect'); - const fieldName = error.passwordErrorField || 'current_password'; - - form.setError(fieldName, { - type: 'manual', - message: errorMessage, - }); - } - } finally { - setIsSubmitting(false); - } - }, [changePassword, form, defaultValues, onSuccess, t]); + }, + [changePassword, form, defaultValues, onSuccess, t] + ); - const handleSubmit = useCallback(async (formValues) => { - await submitPasswordChange(formValues); - }, [submitPasswordChange]); + const handleSubmit = useCallback( + async (formValues) => { + await submitPasswordChange(formValues); + }, + [submitPasswordChange] + ); - const onSubmitHandler = useCallback((e) => { - e?.preventDefault?.(); - if (isSubmitting || isLoading) { - return; - } - form.handleSubmit(handleSubmit)(); - }, [form, handleSubmit, isSubmitting, isLoading]); + const onSubmitHandler = useCallback( + (e) => { + e?.preventDefault?.(); + if (isSubmitting || isLoading) { + return; + } + form.handleSubmit(handleSubmit)(); + }, + [form, handleSubmit, isSubmitting, isLoading] + ); - const handleKeyDown = useCallback((e) => { - if (e.key === 'Enter' && !isSubmitting && !isLoading) { - e.preventDefault(); - onSubmitHandler(e); - } - }, [onSubmitHandler, isSubmitting, isLoading]); + const handleKeyDown = useCallback( + (e) => { + if (e.key === "Enter" && !isSubmitting && !isLoading) { + e.preventDefault(); + onSubmitHandler(e); + } + }, + [onSubmitHandler, isSubmitting, isLoading] + ); return (
- e.preventDefault()} - onKeyDown={handleKeyDown} - className="space-y-6" - > + e.preventDefault()} onKeyDown={handleKeyDown} className="space-y-6">
- {t('pages.profile.security.fields.currentPassword.label')} + {t("pages.profile.security.fields.currentPassword.label")} * @@ -206,7 +224,7 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { disabled={isSubmitting || isLoading} className={cn( form.formState.errors.current_password && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} /> @@ -223,11 +241,11 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { - {t('pages.profile.security.fields.newPassword.label')} + {t("pages.profile.security.fields.newPassword.label")} * @@ -240,7 +258,7 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { disabled={isSubmitting || isLoading} className={cn( form.formState.errors.new_password && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} onChange={(e) => { field.onChange(e); @@ -261,11 +279,11 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { - {t('pages.profile.security.fields.confirmNewPassword.label')} + {t("pages.profile.security.fields.confirmNewPassword.label")} * @@ -278,7 +296,7 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { disabled={isSubmitting || isLoading} className={cn( form.formState.errors.confirm_new_password && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} /> @@ -286,64 +304,54 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { )} /> -
- - {strengthConfig.label} - + {strengthConfig.label}
  • - = 8 && 'text-success' - )}> - {t('pages.profile.security.fields.newPassword.validation.minLength')} + = 8 && "text-success")}> + {t("pages.profile.security.fields.newPassword.validation.minLength")} - {form.watch('new_password')?.length >= 8 && ( + {form.watch("new_password")?.length >= 8 && ( )}
  • - - {t('pages.profile.security.fields.newPassword.strength.uppercase')} + + {t("pages.profile.security.fields.newPassword.strength.uppercase")} - {form.watch('new_password')?.match(/[A-Z]/) && ( + {form.watch("new_password")?.match(/[A-Z]/) && ( )}
  • - - {t('pages.profile.security.fields.newPassword.strength.number')} + + {t("pages.profile.security.fields.newPassword.strength.number")} - {form.watch('new_password')?.match(/[0-9]/) && ( + {form.watch("new_password")?.match(/[0-9]/) && ( )}
  • - - {t('pages.profile.security.fields.newPassword.strength.special')} + + {t("pages.profile.security.fields.newPassword.strength.special")} - {form.watch('new_password')?.match(/[^A-Za-z0-9]/) && ( + {form.watch("new_password")?.match(/[^A-Za-z0-9]/) && ( )}
  • @@ -360,7 +368,7 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { disabled={isSubmitting || isLoading} className="flex items-center gap-1.5 md:gap-2 text-sm" > - {t('common.actions.cancel')} + {t("common.actions.cancel")} )}
@@ -383,4 +391,4 @@ export function ChangePasswordForm({ onSuccess, onClose, onSubmittingChange }) { ); } -export default ChangePasswordForm; \ No newline at end of file +export default ChangePasswordForm; diff --git a/frontend/src/components/profile/profile-info.jsx b/frontend/src/components/profile/profile-info.jsx index 74d625a..f0c11c4 100644 --- a/frontend/src/components/profile/profile-info.jsx +++ b/frontend/src/components/profile/profile-info.jsx @@ -1,6 +1,6 @@ -import { useTranslation } from 'react-i18next'; -import { Label } from '@/components/ui/label'; -import { Input } from '@/components/ui/input'; +import { useTranslation } from "react-i18next"; +import { Label } from "@/components/ui/label"; +import { Input } from "@/components/ui/input"; export function ProfileInfo({ user }) { const { t } = useTranslation(); @@ -13,23 +13,19 @@ export function ProfileInfo({ user }) {
- +
- + @@ -37,27 +33,23 @@ export function ProfileInfo({ user }) {
- +
- + @@ -67,4 +59,4 @@ export function ProfileInfo({ user }) { ); } -export default ProfileInfo; \ No newline at end of file +export default ProfileInfo; diff --git a/frontend/src/components/profile/update-profile-form.jsx b/frontend/src/components/profile/update-profile-form.jsx index e9f7bc2..5869d41 100644 --- a/frontend/src/components/profile/update-profile-form.jsx +++ b/frontend/src/components/profile/update-profile-form.jsx @@ -1,10 +1,10 @@ -import { useState, useCallback, useMemo, useEffect } from 'react'; -import { useForm } from 'react-hook-form'; -import { zodResolver } from '@hookform/resolvers/zod'; -import { z } from 'zod'; -import { useTranslation } from 'react-i18next'; -import { cn, debugError } from '@/lib/utils'; -import { Button } from '@/components/ui/button'; +import { useState, useCallback, useMemo, useEffect } from "react"; +import { useForm } from "react-hook-form"; +import { zodResolver } from "@hookform/resolvers/zod"; +import { z } from "zod"; +import { useTranslation } from "react-i18next"; +import { cn, debugError } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; import { Form, FormControl, @@ -12,10 +12,10 @@ import { FormItem, FormLabel, FormMessage, -} from '@/components/ui/form'; -import { Input } from '@/components/ui/input'; -import { Spinner } from '@/components/ui/spinner'; -import { useAuth } from '@/hooks/useAuth'; +} from "@/components/ui/form"; +import { Input } from "@/components/ui/input"; +import { Spinner } from "@/components/ui/spinner"; +import { useAuth } from "@/hooks/useAuth"; export function UpdateProfileForm({ user, @@ -38,32 +38,34 @@ export function UpdateProfileForm({ return z.object({ first_name: z .string() - .min(1, t('pages.profile.profile.fields.firstName.validation.required')), - last_name: z - .string() - .min(1, t('pages.profile.profile.fields.lastName.validation.required')), + .min(1, t("pages.profile.profile.fields.firstName.validation.required")), + last_name: z.string().min(1, t("pages.profile.profile.fields.lastName.validation.required")), email: z .string() - .min(1, t('pages.profile.profile.fields.email.validation.required')) + .min(1, t("pages.profile.profile.fields.email.validation.required")) .transform((val) => val.trim().toLowerCase()) - .refine((val) => { - const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/; - return emailRegex.test(val); - }, { - message: t('pages.profile.profile.fields.email.validation.invalid'), - }), - phone: z - .string() - .min(1, t('pages.profile.profile.fields.phone.validation.required')), + .refine( + (val) => { + const emailRegex = /^[^\s@]+@[^\s@]+\.[^\s@]+$/; + return emailRegex.test(val); + }, + { + message: t("pages.profile.profile.fields.email.validation.invalid"), + } + ), + phone: z.string().min(1, t("pages.profile.profile.fields.phone.validation.required")), }); }, [t]); - const defaultValues = useMemo(() => ({ - first_name: user?.first_name || '', - last_name: user?.last_name || '', - email: user?.email || '', - phone: user?.phone || '', - }), [user]); + const defaultValues = useMemo( + () => ({ + first_name: user?.first_name || "", + last_name: user?.last_name || "", + email: user?.email || "", + phone: user?.phone || "", + }), + [user] + ); const form = useForm({ resolver: zodResolver(formSchema), @@ -73,43 +75,49 @@ export function UpdateProfileForm({ useEffect(() => { if (user) { form.reset({ - first_name: user.first_name || '', - last_name: user.last_name || '', - email: user.email || '', - phone: user.phone || '', + first_name: user.first_name || "", + last_name: user.last_name || "", + email: user.email || "", + phone: user.phone || "", }); } }, [user, form]); - const handleSubmit = useCallback(async (formValues) => { - setIsSubmitting(true); + const handleSubmit = useCallback( + async (formValues) => { + setIsSubmitting(true); - try { - const result = await updateUserProfile({ - first_name: formValues.first_name, - last_name: formValues.last_name, - email: formValues.email, - phone: formValues.phone, - }, {returnStatus: true}); + try { + const result = await updateUserProfile( + { + first_name: formValues.first_name, + last_name: formValues.last_name, + email: formValues.email, + phone: formValues.phone, + }, + { returnStatus: true } + ); - if (result?.requiresEmailVerification) { - if (onRequiresEmailVerification) { - await onRequiresEmailVerification({ - email: result.email || result.data?.pending_email || formValues.email, - profile: result.data, - }); - } - } else if (result?.status === 'success') { - if (onSuccess) { - await onSuccess(result.data); + if (result?.requiresEmailVerification) { + if (onRequiresEmailVerification) { + await onRequiresEmailVerification({ + email: result.email || result.data?.pending_email || formValues.email, + profile: result.data, + }); + } + } else if (result?.status === "success") { + if (onSuccess) { + await onSuccess(result.data); + } } + } catch (error) { + debugError("Update profile error:", error); + } finally { + setIsSubmitting(false); } - } catch (error) { - debugError('Update profile error:', error); - } finally { - setIsSubmitting(false); - } - }, [updateUserProfile, onSuccess, onRequiresEmailVerification]); + }, + [updateUserProfile, onSuccess, onRequiresEmailVerification] + ); return ( @@ -123,11 +131,11 @@ export function UpdateProfileForm({ - {t('pages.profile.profile.fields.firstName.label')} + {t("pages.profile.profile.fields.firstName.label")} * @@ -138,7 +146,7 @@ export function UpdateProfileForm({ disabled={isSubmitting || isLoading} className={cn( form.formState.errors.first_name && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} /> @@ -154,11 +162,11 @@ export function UpdateProfileForm({ - {t('pages.profile.profile.fields.lastName.label')} + {t("pages.profile.profile.fields.lastName.label")} * @@ -169,7 +177,7 @@ export function UpdateProfileForm({ disabled={isSubmitting || isLoading} className={cn( form.formState.errors.last_name && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} /> @@ -187,11 +195,11 @@ export function UpdateProfileForm({ - {t('pages.profile.profile.fields.email.label')} + {t("pages.profile.profile.fields.email.label")} * @@ -203,7 +211,7 @@ export function UpdateProfileForm({ disabled={isSubmitting || isLoading} className={cn( form.formState.errors.email && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} /> @@ -219,11 +227,11 @@ export function UpdateProfileForm({ - {t('pages.profile.profile.fields.phone.label')} + {t("pages.profile.profile.fields.phone.label")} * @@ -235,7 +243,7 @@ export function UpdateProfileForm({ disabled={isSubmitting || isLoading} className={cn( form.formState.errors.phone && - 'ring-2 ring-destructive focus-visible:ring-destructive' + "ring-2 ring-destructive focus-visible:ring-destructive" )} /> @@ -253,7 +261,7 @@ export function UpdateProfileForm({ disabled={isSubmitting || isLoading} className="flex items-center gap-1.5 md:gap-2 text-sm" > - {t('common.actions.cancel')} + {t("common.actions.cancel")} )}
@@ -275,4 +283,4 @@ export function UpdateProfileForm({ ); } -export default UpdateProfileForm; \ No newline at end of file +export default UpdateProfileForm; diff --git a/frontend/src/components/roles/delete-role-dialog.jsx b/frontend/src/components/roles/delete-role-dialog.jsx index 42d76bf..0c88119 100644 --- a/frontend/src/components/roles/delete-role-dialog.jsx +++ b/frontend/src/components/roles/delete-role-dialog.jsx @@ -11,17 +11,11 @@ import { AlertDialogTitle, } from "@/components/ui/alert-dialog"; -export function DeleteRoleDialog({ - roleName = "", - isSubmitting = false, - onConfirm, -}) { +export function DeleteRoleDialog({ roleName = "", isSubmitting = false, onConfirm }) { const { t } = useTranslation(); return ( - e.preventDefault()} - > + e.preventDefault()}> {t("pages.rolesManagement.dialog.deleteTitle", "Delete role")} @@ -51,4 +45,4 @@ export function DeleteRoleDialog({ ); } -export default DeleteRoleDialog; \ No newline at end of file +export default DeleteRoleDialog; diff --git a/frontend/src/components/roles/role-form-dialog.jsx b/frontend/src/components/roles/role-form-dialog.jsx index fb695de..d77717e 100644 --- a/frontend/src/components/roles/role-form-dialog.jsx +++ b/frontend/src/components/roles/role-form-dialog.jsx @@ -1,16 +1,16 @@ -import * as React from "react" -import { useTranslation } from "react-i18next" -import { Button } from "@/components/ui/button" -import { Input } from "@/components/ui/input" -import { Spinner } from "@/components/ui/spinner" -import { - Dialog, - DialogContent, - DialogDescription, - DialogFooter, - DialogHeader, - DialogTitle -} from "@/components/ui/dialog" +import * as React from "react"; +import { useTranslation } from "react-i18next"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Spinner } from "@/components/ui/spinner"; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; export function RoleFormDialog({ open, @@ -63,10 +63,7 @@ export function RoleFormDialog({ const level = Number(formData.level); if (!Number.isInteger(level) || level < 1) { setLevelError( - t( - "pages.rolesManagement.fields.level.validation.min", - "Level must be at least 1" - ) + t("pages.rolesManagement.fields.level.validation.min", "Level must be at least 1") ); return; } @@ -106,24 +103,21 @@ export function RoleFormDialog({ [formData, isSubmitting, validateAndSubmit] ); - const canSubmit = - !isSubmitting && - !!formData?.name?.trim() && - maxAssignableLevel >= 1; + const canSubmit = !isSubmitting && !!formData?.name?.trim() && maxAssignableLevel >= 1; return ( - {isCreate - ? t("pages.rolesManagement.dialog.createTitle", "Create role") - : t("pages.rolesManagement.dialog.editTitle", "Edit role")} + {isCreate + ? t("pages.rolesManagement.dialog.createTitle", "Create role") + : t("pages.rolesManagement.dialog.editTitle", "Edit role")} - {isCreate - ? t("pages.rolesManagement.dialog.createDescription", "Create a new role") - : t("pages.rolesManagement.dialog.editDescription", "Update role info")} + {isCreate + ? t("pages.rolesManagement.dialog.createDescription", "Create a new role") + : t("pages.rolesManagement.dialog.editDescription", "Update role info")} @@ -164,7 +158,10 @@ export function RoleFormDialog({ description: e.target.value, })) } - placeholder={t("pages.rolesManagement.fields.description.placeholder", "Enter description")} + placeholder={t( + "pages.rolesManagement.fields.description.placeholder", + "Enter description" + )} />
@@ -188,10 +185,7 @@ export function RoleFormDialog({ level: e.target.value === "" ? "" : Number(e.target.value), })); }} - placeholder={t( - "pages.rolesManagement.fields.level.placeholder", - "Enter level" - )} + placeholder={t("pages.rolesManagement.fields.level.placeholder", "Enter level")} />

{t( @@ -200,9 +194,7 @@ export function RoleFormDialog({ { max: maxAssignableLevel } )}

- {levelError ? ( -

{levelError}

- ) : null} + {levelError ?

{levelError}

: null}
@@ -210,13 +202,19 @@ export function RoleFormDialog({ - diff --git a/frontend/src/components/roles/role-permissions-desktop.jsx b/frontend/src/components/roles/role-permissions-desktop.jsx index 6206c73..29e40fc 100644 --- a/frontend/src/components/roles/role-permissions-desktop.jsx +++ b/frontend/src/components/roles/role-permissions-desktop.jsx @@ -1,15 +1,15 @@ -import * as React from "react" -import { cn } from "@/lib/utils" -import { useTranslation } from "react-i18next" -import { Button } from "@/components/ui/button" -import { Checkbox } from "@/components/ui/checkbox" -import { Spinner } from "@/components/ui/spinner" -import { Skeleton } from "@/components/ui/skeleton" -import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible" -import { Scroller } from "@/components/ui/scroller" -import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card" -import { Badge } from "@/components/ui/badge" -import { Shield, CheckCircle2, Circle, ChevronRight } from "lucide-react" +import * as React from "react"; +import { cn } from "@/lib/utils"; +import { useTranslation } from "react-i18next"; +import { Button } from "@/components/ui/button"; +import { Checkbox } from "@/components/ui/checkbox"; +import { Spinner } from "@/components/ui/spinner"; +import { Skeleton } from "@/components/ui/skeleton"; +import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; +import { Scroller } from "@/components/ui/scroller"; +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { Badge } from "@/components/ui/badge"; +import { Shield, CheckCircle2, Circle, ChevronRight } from "lucide-react"; const toCamelKey = (input) => { if (!input || typeof input !== "string") return ""; @@ -37,7 +37,8 @@ const normalizeGroupsForDisplay = (groups, attributesFlat = {}) => { .filter((a) => a?.name) .map((a) => ({ name: a.name, - value: typeof attributesFlat?.[a.name] === "boolean" ? attributesFlat[a.name] : !!a.value, + value: + typeof attributesFlat?.[a.name] === "boolean" ? attributesFlat[a.name] : !!a.value, })); return { name: categoryName, attributes: attrs }; @@ -176,7 +177,8 @@ export function RolePermissionsDesktop({ const handleSelectAllGroup = React.useCallback( (groupObj) => { const allAttrNames = getGroupAttributeNames(groupObj); - const allSelected = allAttrNames.length > 0 && allAttrNames.every((name) => !!attributes?.[name]); + const allSelected = + allAttrNames.length > 0 && allAttrNames.every((name) => !!attributes?.[name]); allAttrNames.forEach((name) => { if (allSelected) { if (attributes?.[name]) onAttributeToggle?.(name); @@ -209,7 +211,9 @@ export function RolePermissionsDesktop({ const baseXPaddingPx = 20; const scrollerPaddingRightPx = Math.max(0, baseXPaddingPx - scrollbarWidth - 2); - const roleTitle = selectedRole?.name || t("pages.rolesManagement.selectRoleFromLeft", "Select a role from the left"); + const roleTitle = + selectedRole?.name || + t("pages.rolesManagement.selectRoleFromLeft", "Select a role from the left"); return ( <> @@ -218,7 +222,9 @@ export function RolePermissionsDesktop({ {selectedRole - ? t("pages.rolesManagement.rolePermissionsFor", "Permissions for {{role}}", { role: roleTitle }) + ? t("pages.rolesManagement.rolePermissionsFor", "Permissions for {{role}}", { + role: roleTitle, + }) : t("pages.rolesManagement.selectRoleFromLeft", "Select a role from the left")} {selectedRole?.level != null && ( @@ -232,11 +238,7 @@ export function RolePermissionsDesktop({ {hasChanges && canManageRoles && (
-
) : ( <> @@ -276,7 +284,10 @@ export function RolePermissionsDesktop({ organizedGroups.map((g) => { const groupAttrNames = getGroupAttributeNames(g); const groupTotal = groupAttrNames.length; - const groupSelectedCount = groupAttrNames.reduce((acc, name) => acc + (attributes?.[name] ? 1 : 0), 0); + const groupSelectedCount = groupAttrNames.reduce( + (acc, name) => acc + (attributes?.[name] ? 1 : 0), + 0 + ); const groupCheckedState = groupTotal > 0 && groupSelectedCount === groupTotal ? true @@ -291,21 +302,33 @@ export function RolePermissionsDesktop({ key={g.group} open={isOpen} onOpenChange={(nextOpen) => { - setOpenGroups((prev) => ({ ...(prev || {}), [g.group]: nextOpen })); + setOpenGroups((prev) => ({ + ...(prev || {}), + [g.group]: nextOpen, + })); }} className="border rounded-lg overflow-hidden bg-card" >
-

{getGroupLabel(g.group)}

- +

+ {getGroupLabel(g.group)} +

+
{groupTotal > 0 && (
- ) + ); } return ( <>
-
- {showEmptyState ? ( -
- -

- {t("pages.rolesManagement.dialog.noAttributes", "No available permissions")} -

-
- ) : ( - <> - {organizedGroups.length > 0 && - organizedGroups.map((g) => { - const groupAttrNames = getGroupAttributeNames(g) - const groupTotal = groupAttrNames.length - const groupSelectedCount = groupAttrNames.reduce((acc, name) => acc + (attributes?.[name] ? 1 : 0), 0) - const groupAllSelected = isGroupAllSelected(g) - const groupCheckedState = - groupTotal > 0 && groupSelectedCount === groupTotal - ? true - : groupTotal > 0 && groupSelectedCount > 0 - ? "indeterminate" - : false - const groupCheckboxId = `role-perms-mobile-${selectedRole?.id ?? "new"}-${toCamelKey(g.group) || "group"}` + {showEmptyState ? ( +
+ +

+ {t("pages.rolesManagement.dialog.noAttributes", "No available permissions")} +

+
+ ) : ( + <> + {organizedGroups.length > 0 && + organizedGroups.map((g) => { + const groupAttrNames = getGroupAttributeNames(g); + const groupTotal = groupAttrNames.length; + const groupSelectedCount = groupAttrNames.reduce( + (acc, name) => acc + (attributes?.[name] ? 1 : 0), + 0 + ); + const groupAllSelected = isGroupAllSelected(g); + const groupCheckedState = + groupTotal > 0 && groupSelectedCount === groupTotal + ? true + : groupTotal > 0 && groupSelectedCount > 0 + ? "indeterminate" + : false; + const groupCheckboxId = `role-perms-mobile-${selectedRole?.id ?? "new"}-${toCamelKey(g.group) || "group"}`; - return ( -
-
-
-
- {getGroupLabel(g.group)} + return ( +
+
+
+
+ {getGroupLabel(g.group)} +
+
- -
-
- {g.categories.map((cat, catIdx) => { - const hasAttrs = (cat?.attributes || []).length > 0 - return ( -
-
- {getCategoryLabel(cat.name)} -
- +
+ {g.categories.map((cat, catIdx) => { + const hasAttrs = (cat?.attributes || []).length > 0; + return ( +
+
+ {getCategoryLabel(cat.name)} +
+ - {hasAttrs ? ( - cat.attributes.map((attr, idx) => { - const isEnabled = !!attributes?.[attr.name] - const isLastRow = catIdx === g.categories.length - 1 && idx === cat.attributes.length - 1 - const key = attr?.name || `${cat.name}-${idx}` - return ( - - {renderSettingRow({ - label: getAttributeLabel(attr.name), - checked: isEnabled, - disabled: !canManageRoles || isSubmitting || showSkeletonView, - onToggle: () => onAttributeToggle?.(attr.name), - showDivider: !isLastRow, - isSkeleton: showSkeletonView, - })} - - ) - }) - ) : ( -
-
- )} -
- ) - })} + {hasAttrs ? ( + cat.attributes.map((attr, idx) => { + const isEnabled = !!attributes?.[attr.name]; + const isLastRow = + catIdx === g.categories.length - 1 && + idx === cat.attributes.length - 1; + const key = attr?.name || `${cat.name}-${idx}`; + return ( + + {renderSettingRow({ + label: getAttributeLabel(attr.name), + checked: isEnabled, + disabled: + !canManageRoles || isSubmitting || showSkeletonView, + onToggle: () => onAttributeToggle?.(attr.name), + showDivider: !isLastRow, + isSkeleton: showSkeletonView, + })} + + ); + }) + ) : ( +
-
+ )} +
+ ); + })} +
-
- ) - })} + ); + })} - {organizedGroups.length === 0 && - legacyAttributes.length > 0 && (() => { - const allSelected = legacyAttributes.every((a) => !!attributes?.[a.name]) - const total = legacyAttributes.length - const selectedCount = legacyAttributes.reduce((acc, a) => acc + (attributes?.[a.name] ? 1 : 0), 0) - const checkedState = - total > 0 && selectedCount === total ? true : total > 0 && selectedCount > 0 ? "indeterminate" : false - const legacyCheckboxId = `role-perms-mobile-${selectedRole?.id ?? "new"}-legacy` - return ( -
-
-
-
- {t("pages.rolesManagement.categories.other", "Other")} + {organizedGroups.length === 0 && + legacyAttributes.length > 0 && + (() => { + const allSelected = legacyAttributes.every((a) => !!attributes?.[a.name]); + const total = legacyAttributes.length; + const selectedCount = legacyAttributes.reduce( + (acc, a) => acc + (attributes?.[a.name] ? 1 : 0), + 0 + ); + const checkedState = + total > 0 && selectedCount === total + ? true + : total > 0 && selectedCount > 0 + ? "indeterminate" + : false; + const legacyCheckboxId = `role-perms-mobile-${selectedRole?.id ?? "new"}-legacy`; + return ( +
+
+
+
+ {t("pages.rolesManagement.categories.other", "Other")} +
-
- -
+ : t("pages.rolesManagement.actions.selectAll", "Select all")} + + + + )} + + {selectedCount}/{total} + + +
-
- {legacyAttributes.map((attr, idx) => { - const isEnabled = !!attributes?.[attr.name] - const isLastRow = idx === legacyAttributes.length - 1 - const key = attr?.name || `legacy-${idx}` - return ( - - {renderSettingRow({ - label: getAttributeLabel(attr.name), - checked: isEnabled, - disabled: !canManageRoles || isSubmitting || showSkeletonView, - onToggle: () => onAttributeToggle?.(attr.name), - showDivider: !isLastRow, - isSkeleton: showSkeletonView, - })} - - ) - })} +
+ {legacyAttributes.map((attr, idx) => { + const isEnabled = !!attributes?.[attr.name]; + const isLastRow = idx === legacyAttributes.length - 1; + const key = attr?.name || `legacy-${idx}`; + return ( + + {renderSettingRow({ + label: getAttributeLabel(attr.name), + checked: isEnabled, + disabled: !canManageRoles || isSubmitting || showSkeletonView, + onToggle: () => onAttributeToggle?.(attr.name), + showDivider: !isLastRow, + isSkeleton: showSkeletonView, + })} + + ); + })} +
-
- ) - })()} - - )} + ); + })()} + + )}
@@ -431,10 +469,10 @@ export function RolePermissionsMobile({ { - if (!open) onResetAttributes?.() + if (!open) onResetAttributes?.(); }} onEscapeKeyDown={() => { - onResetAttributes?.() + onResetAttributes?.(); }} className="w-auto max-w-xs" > @@ -446,8 +484,8 @@ export function RolePermissionsMobile({ { - e.preventDefault() - onResetAttributes?.() + e.preventDefault(); + onResetAttributes?.(); }} disabled={isSubmitting} className="size-10 p-0 flex items-center justify-center" @@ -457,8 +495,8 @@ export function RolePermissionsMobile({ { - e.preventDefault() - onSaveAttributes?.() + e.preventDefault(); + onSaveAttributes?.(); }} disabled={isSubmitting} className={cn( @@ -473,7 +511,7 @@ export function RolePermissionsMobile({ )} - ) + ); } -export default RolePermissionsMobile \ No newline at end of file +export default RolePermissionsMobile; diff --git a/frontend/src/components/roles/role-selector.jsx b/frontend/src/components/roles/role-selector.jsx index 704d0b9..6dcfd66 100644 --- a/frontend/src/components/roles/role-selector.jsx +++ b/frontend/src/components/roles/role-selector.jsx @@ -1,19 +1,19 @@ -import * as React from "react" -import { cn } from "@/lib/utils" -import { useTranslation } from "react-i18next" -import { Button } from "@/components/ui/button" -import { Spinner } from "@/components/ui/spinner" -import { Badge } from "@/components/ui/badge" -import { Check, ChevronsUpDown, Edit, Trash2 } from "lucide-react" -import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover" -import { - Command, - CommandEmpty, - CommandGroup, - CommandInput, - CommandItem, +import * as React from "react"; +import { cn } from "@/lib/utils"; +import { useTranslation } from "react-i18next"; +import { Button } from "@/components/ui/button"; +import { Spinner } from "@/components/ui/spinner"; +import { Badge } from "@/components/ui/badge"; +import { Check, ChevronsUpDown, Edit, Trash2 } from "lucide-react"; +import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover"; +import { + Command, + CommandEmpty, + CommandGroup, + CommandInput, + CommandItem, CommandList, -} from "@/components/ui/command" +} from "@/components/ui/command"; export function RoleSelector({ filteredRoles = [], @@ -29,19 +29,20 @@ export function RoleSelector({ onEditClick, onDeleteClick, }) { - const { t } = useTranslation() - const [open, setOpen] = React.useState(false) + const { t } = useTranslation(); + const [open, setOpen] = React.useState(false); - const selectedLabel = selectedRole?.name || t("pages.rolesManagement.selectRole", "Select a role") - const rawIsLoading = !!isLoading - const hasKeyword = !!searchKeyword?.trim() + const selectedLabel = + selectedRole?.name || t("pages.rolesManagement.selectRole", "Select a role"); + const rawIsLoading = !!isLoading; + const hasKeyword = !!searchKeyword?.trim(); const canEditOrDelete = !!selectedRole && canManageRoles && canEditSelected && !rawIsLoading && !disabled && - !isSubmitting + !isSubmitting; return (
@@ -66,7 +67,7 @@ export function RoleSelector({ align="start" onOpenAutoFocus={(e) => { // Mobile UX: do not auto-focus the search input when opening. - e.preventDefault() + e.preventDefault(); }} > @@ -91,45 +92,50 @@ export function RoleSelector({ : t("pages.rolesManagement.noRoles", "No roles yet")} - {filteredRoles.map((role) => ( + {filteredRoles.map((role) => (() => { - const isSelected = selectedRole?.id === role.id + const isSelected = selectedRole?.id === role.id; return ( - { - onRoleSelect?.(role) - onSearchChange?.("") - setOpen(false) - }} - className={cn( - "flex items-center justify-between gap-2 py-2 px-3 rounded-md", - "data-[selected=true]:bg-transparent aria-selected:bg-transparent", - isSelected && "!bg-primary/15 border border-primary/20" - )} - > -
-
-
{role.name}
- {role.level != null && ( - - Lv.{role.level} - + { + onRoleSelect?.(role); + onSearchChange?.(""); + setOpen(false); + }} + className={cn( + "flex items-center justify-between gap-2 py-2 px-3 rounded-md", + "data-[selected=true]:bg-transparent aria-selected:bg-transparent", + isSelected && "!bg-primary/15 border border-primary/20" )} -
-
{role.description || "-"}
-
- -
- ) + > +
+
+
{role.name}
+ {role.level != null && ( + + Lv.{role.level} + + )} +
+
+ {role.description || "-"} +
+
+ + + ); })() - ))} + )}
)} @@ -171,7 +177,7 @@ export function RoleSelector({ {selectedRole?.description || ""}
- ) + ); } -export default RoleSelector \ No newline at end of file +export default RoleSelector; diff --git a/frontend/src/components/roles/roles-list.jsx b/frontend/src/components/roles/roles-list.jsx index a5140b0..f3048a1 100644 --- a/frontend/src/components/roles/roles-list.jsx +++ b/frontend/src/components/roles/roles-list.jsx @@ -1,18 +1,19 @@ -import * as React from "react" -import { useTranslation } from "react-i18next" -import { cn } from "@/lib/utils" -import { Button } from "@/components/ui/button" -import { Input } from "@/components/ui/input" -import { Spinner } from "@/components/ui/spinner" -import { Badge } from "@/components/ui/badge" -import { Card, CardContent } from "@/components/ui/card" -import { Separator } from "@/components/ui/separator" +import * as React from "react"; +import { useTranslation } from "react-i18next"; +import { cn } from "@/lib/utils"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Spinner } from "@/components/ui/spinner"; +import { Badge } from "@/components/ui/badge"; +import { Card, CardContent } from "@/components/ui/card"; +import { Separator } from "@/components/ui/separator"; +import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover"; import { - Popover, - PopoverContent, - PopoverTrigger, -} from "@/components/ui/popover" -import { DropdownMenu, DropdownMenuContent, DropdownMenuItem, DropdownMenuTrigger } from "@/components/ui/dropdown-menu" + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; import { ArrowDown10, ArrowDownAZ, @@ -26,11 +27,11 @@ import { Shield, Trash2, X, -} from "lucide-react" -import { Scroller } from "@/components/ui/scroller" +} from "lucide-react"; +import { Scroller } from "@/components/ui/scroller"; -export function RolesList({ - filteredRoles = [], +export function RolesList({ + filteredRoles = [], selectedRole, isLoading = false, loadingDelayMs = 0, @@ -225,10 +226,10 @@ export function RolesList({ className={cn("w-full h-10 rounded-md", searchKeyword?.trim() && "pr-10")} /> {!!searchKeyword?.trim() && ( -
{canManageRoles ? ( -
) : (
- {filteredRoles.map((role) => ( + {filteredRoles.map((role) => (() => { const isSelected = selectedRole?.id === role.id; return ( @@ -307,58 +308,60 @@ export function RolesList({ )}
-
{role.description || "-"}
+
+ {role.description || "-"} +
{canEditOrDelete && Number(role.level ?? 0) <= Number(actorLevel ?? 0) && role.id !== actorRoleId && ( -
e.stopPropagation()} - > - - - - - - onEditClick?.(role)} - className="justify-between gap-2" - disabled={isSubmitting} - > - {t("common.actions.edit", "Edit")} - - - onDeleteClick?.(role)} - className="justify-between gap-2 text-destructive focus:text-destructive hover:!bg-destructive/10" - disabled={isSubmitting} - > - {t("common.actions.delete", "Delete")} - - - - -
- )} +
e.stopPropagation()} + > + + + + + + onEditClick?.(role)} + className="justify-between gap-2" + disabled={isSubmitting} + > + {t("common.actions.edit", "Edit")} + + + onDeleteClick?.(role)} + className="justify-between gap-2 text-destructive focus:text-destructive hover:!bg-destructive/10" + disabled={isSubmitting} + > + {t("common.actions.delete", "Delete")} + + + + +
+ )}
); })() - ))} + )}
)} diff --git a/frontend/src/components/roles/roles-management-panel.jsx b/frontend/src/components/roles/roles-management-panel.jsx index b8ba4a3..b4802f7 100644 --- a/frontend/src/components/roles/roles-management-panel.jsx +++ b/frontend/src/components/roles/roles-management-panel.jsx @@ -9,7 +9,10 @@ import { RolePermissionsMobile } from "./role-permissions-mobile"; import { AlertDialog } from "@/components/ui/alert-dialog"; import { DeleteRoleDialog } from "@/components/roles/delete-role-dialog"; -export const RolesManagementPanel = React.forwardRef(function RolesManagementPanel({ canManageRoles = false }, ref) { +export const RolesManagementPanel = React.forwardRef(function RolesManagementPanel( + { canManageRoles = false }, + ref +) { const isMobile = useIsMobile(); const [roles, setRoles] = React.useState([]); const [filteredRoles, setFilteredRoles] = React.useState([]); @@ -155,7 +158,7 @@ export const RolesManagementPanel = React.forwardRef(function RolesManagementPan sensitivity: "base", }) * direction; if (nameDiff !== 0) return nameDiff; - return (Number(b.level ?? 0) - Number(a.level ?? 0)); + return Number(b.level ?? 0) - Number(a.level ?? 0); }); setFilteredRoles(nextRoles); @@ -257,7 +260,8 @@ export const RolesManagementPanel = React.forwardRef(function RolesManagementPan ...prev, [attributeName]: !prev[attributeName], }; - const hasChanged = JSON.stringify(newAttributes) !== JSON.stringify(initialAttributesRef.current); + const hasChanged = + JSON.stringify(newAttributes) !== JSON.stringify(initialAttributesRef.current); setHasChanges(hasChanged); return newAttributes; }); @@ -267,7 +271,11 @@ export const RolesManagementPanel = React.forwardRef(function RolesManagementPan const handleSaveAttributes = React.useCallback(async () => { if (!selectedRole?.id) return; setIsSubmitting(true); - const result = await rolesService.updateRoleAttributes(selectedRole.id, { attributes }, { returnStatus: true }); + const result = await rolesService.updateRoleAttributes( + selectedRole.id, + { attributes }, + { returnStatus: true } + ); if (result.status === "success") { initialAttributesRef.current = { ...attributes }; setHasChanges(false); @@ -291,7 +299,9 @@ export const RolesManagementPanel = React.forwardRef(function RolesManagementPan try { let createdOrUpdatedRoleId = roleId || null; if (createdOrUpdatedRoleId) { - const updateResult = await rolesService.updateRole(createdOrUpdatedRoleId, roleData, { returnStatus: true }); + const updateResult = await rolesService.updateRole(createdOrUpdatedRoleId, roleData, { + returnStatus: true, + }); if (updateResult.status !== "success") { return false; } @@ -365,12 +375,9 @@ export const RolesManagementPanel = React.forwardRef(function RolesManagementPan } }, [selectedRole]); - const handleRoleSelect = React.useCallback( - (role) => { - setSelectedRole(role); - }, - [] - ); + const handleRoleSelect = React.useCallback((role) => { + setSelectedRole(role); + }, []); if (isMobile) { return ( @@ -588,4 +595,4 @@ export const RolesManagementPanel = React.forwardRef(function RolesManagementPan ); }); -export default RolesManagementPanel; \ No newline at end of file +export default RolesManagementPanel; diff --git a/frontend/src/components/sidebar/app-sidebar.jsx b/frontend/src/components/sidebar/app-sidebar.jsx index 66ba2c0..be1bfd6 100644 --- a/frontend/src/components/sidebar/app-sidebar.jsx +++ b/frontend/src/components/sidebar/app-sidebar.jsx @@ -1,53 +1,49 @@ -import * as React from "react" -import { useTranslation } from "react-i18next" -import { useAuth } from "@/hooks/useAuth" -import { useSidebarRoutes } from "@/lib/sidebar-routes" -import { NavMain } from "./nav-main" -import { NavProjects } from "./nav-projects" -import { NavUser } from "./nav-user" -import { TeamSwitcher } from "./team-switcher" +import * as React from "react"; +import { useTranslation } from "react-i18next"; +import { useAuth } from "@/hooks/useAuth"; +import { useSidebarRoutes } from "@/lib/sidebar-routes"; +import { NavMain } from "./nav-main"; +import { NavProjects } from "./nav-projects"; +import { NavUser } from "./nav-user"; +import { TeamSwitcher } from "./team-switcher"; import { Sidebar, SidebarContent, SidebarFooter, SidebarHeader, SidebarSeparator, -} from "@/components/ui/sidebar" +} from "@/components/ui/sidebar"; -export function AppSidebar({ - user, - onLogout, - ...props -}) { - const { isAuthenticated } = useAuth() - const { t } = useTranslation() - const { navMain, projects, groups } = useSidebarRoutes(isAuthenticated) +export function AppSidebar({ user: _user, onLogout: _onLogout, ...props }) { + const { isAuthenticated } = useAuth(); + const { t } = useTranslation(); + const { navMain, projects, groups } = useSidebarRoutes(isAuthenticated); // Group navMain items by their group const groupedNavMain = React.useMemo(() => { - const grouped = {} + const grouped = {}; navMain.forEach((item) => { - const group = item.group || "Other" + const group = item.group || "Other"; if (!grouped[group]) { - grouped[group] = [] + grouped[group] = []; } - grouped[group].push(item) - }) - return grouped - }, [navMain]) + grouped[group].push(item); + }); + return grouped; + }, [navMain]); // Group projects items by their group const groupedProjects = React.useMemo(() => { - const grouped = {} + const grouped = {}; projects.forEach((item) => { - const group = item.group || "Projects" + const group = item.group || "Projects"; if (!grouped[group]) { - grouped[group] = [] + grouped[group] = []; } - grouped[group].push(item) - }) - return grouped - }, [projects]) + grouped[group].push(item); + }); + return grouped; + }, [projects]); return ( @@ -59,30 +55,30 @@ export function AppSidebar({ {groups .map((group) => { if (group === "Projects") { - const groupProjects = groupedProjects[group] || [] - if (groupProjects.length === 0) return null + const groupProjects = groupedProjects[group] || []; + if (groupProjects.length === 0) return null; return ( - - ) + ); } else { - const groupNavMain = groupedNavMain[group] || [] - if (groupNavMain.length === 0) return null + const groupNavMain = groupedNavMain[group] || []; + if (groupNavMain.length === 0) return null; return ( - - ) + ); } }) .filter(Boolean) .map((component, index, array) => { - const isLastGroup = index === array.length - 1 + const isLastGroup = index === array.length - 1; return ( {component} @@ -90,7 +86,7 @@ export function AppSidebar({ )} - ) + ); })} diff --git a/frontend/src/components/sidebar/nav-main.jsx b/frontend/src/components/sidebar/nav-main.jsx index dd18ea3..8ba9b6d 100644 --- a/frontend/src/components/sidebar/nav-main.jsx +++ b/frontend/src/components/sidebar/nav-main.jsx @@ -2,11 +2,7 @@ import * as React from "react"; import { ChevronRight } from "lucide-react"; import { Link } from "react-router-dom"; import { Icon } from "@/components/ui/icon"; -import { - Collapsible, - CollapsibleContent, - CollapsibleTrigger, -} from "@/components/ui/collapsible" +import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; import { DropdownMenu, DropdownMenuContent, @@ -14,7 +10,7 @@ import { DropdownMenuLabel, DropdownMenuSeparator, DropdownMenuTrigger, -} from "@/components/ui/dropdown-menu" +} from "@/components/ui/dropdown-menu"; import { SidebarGroup, SidebarGroupLabel, @@ -25,13 +21,13 @@ import { SidebarMenuSubButton, SidebarMenuSubItem, useSidebar, -} from "@/components/ui/sidebar" -import { cn } from "@/lib/utils" +} from "@/components/ui/sidebar"; +import { cn } from "@/lib/utils"; -function CollapsedMenuItemWithDropdown({ item }) { +function CollapsedMenuItemWithDropdown({ item, onItemClick }) { const [dropdownOpen, setDropdownOpen] = React.useState(false); const [showTooltip, setShowTooltip] = React.useState(true); - + React.useEffect(() => { if (dropdownOpen) { setShowTooltip(false); @@ -42,14 +38,12 @@ function CollapsedMenuItemWithDropdown({ item }) { return () => clearTimeout(timer); } }, [dropdownOpen]); - + return ( - + - + item.isActive && "hover:!bg-primary/15 hover:!text-primary" + )} + > {item.iconName && ( - svg]:!size-7 transition-[width,height] duration-300 ease-in-out", - "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7", + "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7" )} /> )} {item.title} - - {item.title} - + + + {item.title} + + {item.items.map((subItem) => ( - + subItem.isActive && + "bg-primary/15 text-primary font-semibold hover:!bg-primary/15 hover:!text-primary" + )} + > {subItem.iconName && ( - svg]:!size-6", + "text-sidebar-foreground transition-[width,height] duration-300 ease-in-out [&>svg]:!size-6" )} /> )} @@ -104,20 +108,17 @@ function CollapsedMenuItemWithDropdown({ item }) { ); } -export function NavMain({ - items, - groupLabel -}) { - const { state, isMobile, toggleSidebar } = useSidebar() - const isCollapsed = state === "collapsed" - const shouldUseDropdown = isCollapsed && !isMobile - +export function NavMain({ items, groupLabel }) { + const { state, isMobile, toggleSidebar } = useSidebar(); + const isCollapsed = state === "collapsed"; + const shouldUseDropdown = isCollapsed && !isMobile; + const handleItemClick = React.useCallback(() => { if (isMobile) { - toggleSidebar() + toggleSidebar(); } - }, [isMobile, toggleSidebar]) - + }, [isMobile, toggleSidebar]); + if (!items || items.length === 0) { return null; } @@ -128,23 +129,28 @@ export function NavMain({ {items.map((item) => { const hasItems = item.items && item.items.length > 0; - + if (hasItems) { if (shouldUseDropdown) { return ( - + ); } - + return ( + className={cn("group/collapsible")} + > - + item.isActive && "hover:!bg-primary/15 hover:!text-primary" + )} + > {item.iconName && ( - svg]:!size-7 transition-[width,height] duration-300 ease-in-out", - "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7", + "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7" )} /> )} @@ -170,27 +177,29 @@ export function NavMain({ className={cn( "ml-auto transition-transform duration-300 ease-in-out", "group-data-[state=open]/collapsible:rotate-90" - )} /> + )} + /> {item.items.map((subItem) => ( - + subItem.isActive && "hover:bg-primary/15 hover:text-primary" + )} + > {subItem.iconName && ( - svg]:!size-7 transition-[width,height] duration-300 ease-in-out", - "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7", + "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7" )} /> )} @@ -207,7 +216,7 @@ export function NavMain({ } else { return ( - + item.isActive && "hover:bg-primary/15 hover:text-primary" + )} + > {item.iconName && ( - svg]:!size-7 transition-[width,height] duration-300 ease-in-out", - "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7", + "group-data-[collapsible=icon]:mx-0.5 group-data-[collapsible=icon]:!size-7" )} /> )} @@ -240,4 +250,4 @@ export function NavMain({ ); -} \ No newline at end of file +} diff --git a/frontend/src/components/sidebar/nav-projects.jsx b/frontend/src/components/sidebar/nav-projects.jsx index 558416d..22b98be 100644 --- a/frontend/src/components/sidebar/nav-projects.jsx +++ b/frontend/src/components/sidebar/nav-projects.jsx @@ -8,7 +8,7 @@ import { DropdownMenuItem, DropdownMenuSeparator, DropdownMenuTrigger, -} from "@/components/ui/dropdown-menu" +} from "@/components/ui/dropdown-menu"; import { SidebarGroup, SidebarGroupLabel, @@ -17,20 +17,17 @@ import { SidebarMenuButton, SidebarMenuItem, useSidebar, -} from "@/components/ui/sidebar" -import { cn } from "@/lib/utils" +} from "@/components/ui/sidebar"; +import { cn } from "@/lib/utils"; + +export function NavProjects({ projects, groupLabel }) { + const { isMobile, toggleSidebar } = useSidebar(); -export function NavProjects({ - projects, - groupLabel -}) { - const { isMobile, toggleSidebar } = useSidebar() - const handleItemClick = React.useCallback(() => { if (isMobile) { - toggleSidebar() + toggleSidebar(); } - }, [isMobile, toggleSidebar]) + }, [isMobile, toggleSidebar]); if (!projects || projects.length === 0) { return null; @@ -44,9 +41,7 @@ export function NavProjects({ - {item.iconName && ( - - )} + {item.iconName && } {item.title || item.name} @@ -60,7 +55,8 @@ export function NavProjects({ + align={isMobile ? "end" : "start"} + > View Project @@ -87,4 +83,4 @@ export function NavProjects({ ); -} \ No newline at end of file +} diff --git a/frontend/src/components/sidebar/nav-user.jsx b/frontend/src/components/sidebar/nav-user.jsx index 4dca1a5..4bd0f91 100644 --- a/frontend/src/components/sidebar/nav-user.jsx +++ b/frontend/src/components/sidebar/nav-user.jsx @@ -1,24 +1,12 @@ -import * as React from "react" -import { Link } from "react-router-dom" -import { useTranslation } from "react-i18next" -import { useTheme } from "@/contexts/themeContext" -import { useAuth } from "@/hooks/useAuth" -import { useIsMobile } from "@/hooks/useMobile" -import { cn } from "@/lib/utils" -import { - LogOut, - Languages, - Moon, - Sun, - Check, - CircleUser, - EllipsisVertical, -} from "lucide-react" -import { - Avatar, - AvatarFallback, - AvatarImage, -} from "@/components/ui/avatar" +import * as React from "react"; +import { Link } from "react-router-dom"; +import { useTranslation } from "react-i18next"; +import { useTheme } from "@/contexts/themeContext"; +import { useAuth } from "@/hooks/useAuth"; +import { useIsMobile } from "@/hooks/useMobile"; +import { cn } from "@/lib/utils"; +import { LogOut, Languages, Moon, Sun, Check, CircleUser, EllipsisVertical } from "lucide-react"; +import { Avatar, AvatarFallback, AvatarImage } from "@/components/ui/avatar"; import { DropdownMenu, DropdownMenuContent, @@ -30,13 +18,13 @@ import { DropdownMenuSubContent, DropdownMenuSubTrigger, DropdownMenuTrigger, -} from "@/components/ui/dropdown-menu" +} from "@/components/ui/dropdown-menu"; import { SidebarMenu, SidebarMenuButton, SidebarMenuItem, useSidebar, -} from "@/components/ui/sidebar" +} from "@/components/ui/sidebar"; import { AlertDialog, AlertDialogContent, @@ -46,59 +34,59 @@ import { AlertDialogFooter, AlertDialogCancel, AlertDialogAction, -} from "@/components/ui/alert-dialog" +} from "@/components/ui/alert-dialog"; export function NavUser() { - const { user, logout } = useAuth() - const { t, i18n: i18nInstance } = useTranslation() - const { theme, setTheme, themes } = useTheme() - const isMobile = useIsMobile() - const { state, toggleSidebar } = useSidebar() + const { user, logout } = useAuth(); + const { t, i18n: i18nInstance } = useTranslation(); + const { theme, setTheme, themes } = useTheme(); + const isMobile = useIsMobile(); + const { state, toggleSidebar } = useSidebar(); - const [alertOpen, setAlertOpen] = React.useState(false) + const [alertOpen, setAlertOpen] = React.useState(false); const themeOptions = [ { value: themes.LIGHT, label: t("components.sidebar.setting.themeOptions.light"), icon: Sun }, { value: themes.DARK, label: t("components.sidebar.setting.themeOptions.dark"), icon: Moon }, - ] + ]; const languageOptions = [ { value: "zh-TW", label: "繁體中文" }, { value: "en", label: "English" }, - ] + ]; - const currentTheme = themeOptions.find(opt => opt.value === theme)?.label - const currentLanguage = languageOptions.find(opt => opt.value === i18nInstance.language)?.label + const currentTheme = themeOptions.find((opt) => opt.value === theme)?.label; + const currentLanguage = languageOptions.find((opt) => opt.value === i18nInstance.language)?.label; const username = React.useMemo(() => { - if (!user) return ''; - return `${user.first_name || ''} ${user.last_name || ''}`.trim(); + if (!user) return ""; + return `${user.first_name || ""} ${user.last_name || ""}`.trim(); }, [user]); const email = React.useMemo(() => { - return user?.email || ''; + return user?.email || ""; }, [user]); - const displayName = username || t("components.sidebar.setting.guest") - const displayEmail = email + const displayName = username || t("components.sidebar.setting.guest"); + const displayEmail = email; const userInitials = React.useMemo(() => { - if (!user) return 'G'; - const first = user.first_name?.[0]?.toUpperCase() || ''; - const last = user.last_name?.[0]?.toUpperCase() || ''; - return (first + last) || user.email?.[0]?.toUpperCase() || 'G'; + if (!user) return "G"; + const first = user.first_name?.[0]?.toUpperCase() || ""; + const last = user.last_name?.[0]?.toUpperCase() || ""; + return first + last || user.email?.[0]?.toUpperCase() || "G"; }, [user]); const handleItemClick = () => { if (isMobile) { - toggleSidebar() + toggleSidebar(); } - } + }; const handleLanguageChange = (languageValue) => { i18nInstance.changeLanguage(languageValue); localStorage.setItem("app-language", languageValue); - } + }; return ( <> @@ -116,42 +104,45 @@ export function NavUser() { "group-data-[collapsible=icon]:p-1!", "group-data-[collapsible=icon]:mx-0.5!", "group-data-[collapsible=icon]:rounded-lg!", - "flex items-center gap-2 p-2 h-13 rounded-md transition-[padding,margin,width] duration-300 ease-in-out will-change-transform", - )}> - + + )} + > {userInitials} -
+
{displayName} {displayEmail}
- + sideOffset={isMobile ? 8 : 4} + > @@ -162,17 +153,22 @@ export function NavUser() {
- {theme === themes.DARK ? : } + {theme === themes.DARK ? ( + + ) : ( + + )} {t("components.sidebar.setting.theme")}
{currentTheme}
- + alignOffset={isMobile ? -200 : 0} + > {themeOptions.map((option) => ( + )} + >
{option.label}
- {currentLanguage} - + alignOffset={isMobile ? -200 : 0} + > {languageOptions.map((option) => ( + )} + > {option.label} {i18nInstance.language === option.value && ( @@ -227,7 +226,10 @@ export function NavUser() {
- setAlertOpen(true)}> + setAlertOpen(true)} + > {t("components.sidebar.setting.logout")} @@ -238,10 +240,7 @@ export function NavUser() { - + {t("components.sidebar.setting.logoutConfirm")} {t("components.sidebar.setting.logoutConfirmDescription")} @@ -251,9 +250,10 @@ export function NavUser() { {t("common.actions.cancel")} { - setAlertOpen(false) - logout() - }}> + setAlertOpen(false); + logout(); + }} + > {t("common.actions.confirm")} @@ -261,4 +261,4 @@ export function NavUser() { ); -} \ No newline at end of file +} diff --git a/frontend/src/components/sidebar/team-switcher.jsx b/frontend/src/components/sidebar/team-switcher.jsx index 5693c19..d04db6a 100644 --- a/frontend/src/components/sidebar/team-switcher.jsx +++ b/frontend/src/components/sidebar/team-switcher.jsx @@ -1,7 +1,7 @@ -import * as React from "react" -import { ChevronsUpDown, Plus } from "lucide-react" -import { useNavigate } from "react-router-dom" -import { useTranslation } from "react-i18next" +import * as React from "react"; +import { ChevronsUpDown, Plus } from "lucide-react"; +import { useNavigate } from "react-router-dom"; +import { useTranslation } from "react-i18next"; import { DropdownMenu, DropdownMenuContent, @@ -10,40 +10,38 @@ import { DropdownMenuSeparator, DropdownMenuShortcut, DropdownMenuTrigger, -} from "@/components/ui/dropdown-menu" +} from "@/components/ui/dropdown-menu"; import { SidebarMenu, SidebarMenuButton, SidebarMenuItem, useSidebar, -} from "@/components/ui/sidebar" -import { cn } from "@/lib/utils" +} from "@/components/ui/sidebar"; +import { cn } from "@/lib/utils"; -export function TeamSwitcher({ - teams = [] -}) { - const { isMobile, state, toggleSidebar } = useSidebar() - const navigate = useNavigate() - const { t } = useTranslation() - const [activeTeam, setActiveTeam] = React.useState(teams?.[0]) - - const shouldShowDropdown = teams && teams.length > 1 - const isCollapsed = state === "collapsed" +export function TeamSwitcher({ teams = [] }) { + const { isMobile, state, toggleSidebar } = useSidebar(); + const navigate = useNavigate(); + const { t } = useTranslation(); + const [activeTeam, setActiveTeam] = React.useState(teams?.[0]); + + const shouldShowDropdown = teams && teams.length > 1; + const isCollapsed = state === "collapsed"; const handleClick = () => { if (!shouldShowDropdown) { - navigate("/") + navigate("/"); if (isMobile) { - toggleSidebar() + toggleSidebar(); } } - } + }; if (!shouldShowDropdown) { - const appName = t("components.app.name") - const displayTeam = activeTeam || { logo: null, name: appName, plan: "" } - const LogoComponent = displayTeam.logo - + const appName = t("components.app.name"); + const displayTeam = activeTeam || { logo: null, name: appName, plan: "" }; + const LogoComponent = displayTeam.logo; + return ( @@ -56,30 +54,33 @@ export function TeamSwitcher({ "group-data-[collapsible=icon]:size-13!", "group-data-[collapsible=icon]:p-1!", "group-data-[collapsible=icon]:mx-0.5!", - "group-data-[collapsible=icon]:rounded-lg!", - )}> + "group-data-[collapsible=icon]:rounded-lg!" + )} + > {LogoComponent ? ( - ) : ( - {appName} )} -
+
{displayTeam.name} {displayTeam.plan && ( {displayTeam.plan} @@ -88,11 +89,11 @@ export function TeamSwitcher({ - ) + ); } if (!activeTeam) { - return null + return null; } return ( @@ -105,12 +106,14 @@ export function TeamSwitcher({ className={cn( "data-[state=open]:bg-sidebar-accent", "data-[state=open]:text-sidebar-accent-foreground" - )}> + )} + >
+ )} + >
@@ -121,24 +124,21 @@ export function TeamSwitcher({ + sideOffset={4} + > Teams {teams.map((team, index) => ( - setActiveTeam(team)} - className={cn("gap-2 p-2")}> -
+ setActiveTeam(team)} + className={cn("gap-2 p-2")} + > +
{team.name} @@ -150,7 +150,8 @@ export function TeamSwitcher({
+ )} + >
Add team
@@ -160,4 +161,4 @@ export function TeamSwitcher({ ); -} \ No newline at end of file +} diff --git a/frontend/src/components/ui/action-bar.jsx b/frontend/src/components/ui/action-bar.jsx index 5b4af4e..388a682 100644 --- a/frontend/src/components/ui/action-bar.jsx +++ b/frontend/src/components/ui/action-bar.jsx @@ -18,10 +18,7 @@ const ITEM_SELECT = "actionbar.itemSelect"; const ENTRY_FOCUS = "actionbarFocusGroup.onEntryFocus"; const EVENT_OPTIONS = { bubbles: false, cancelable: true }; -function focusFirst( - candidates, - preventScroll = false, -) { +function focusFirst(candidates, preventScroll = false) { const PREVIOUSLY_FOCUSED_ELEMENT = document.activeElement; for (const candidateRef of candidates) { const candidate = candidateRef.current; @@ -38,11 +35,7 @@ function wrapArray(array, startIndex) { function getDirectionAwareKey(key, dir) { if (dir !== "rtl") return key; - return key === "ArrowLeft" - ? "ArrowRight" - : key === "ArrowRight" - ? "ArrowLeft" - : key; + return key === "ArrowLeft" ? "ArrowRight" : key === "ArrowRight" ? "ArrowLeft" : key; } const ActionBarContext = React.createContext(null); @@ -119,15 +112,17 @@ function ActionBar(props) { return () => ownerDocument.removeEventListener("keydown", onKeyDown); }, [open, propsRef]); - const contextValue = React.useMemo(() => ({ - onOpenChange, - dir, - orientation, - loop, - }), [onOpenChange, dir, orientation, loop]); + const contextValue = React.useMemo( + () => ({ + onOpenChange, + dir, + orientation, + loop, + }), + [onOpenChange, dir, orientation, loop] + ); - const portalContainer = - portalContainerProp ?? (mounted ? globalThis.document?.body : null); + const portalContainer = portalContainerProp ?? (mounted ? globalThis.document?.body : null); if (!portalContainer || !open) return null; @@ -135,65 +130,69 @@ function ActionBar(props) { return ( - {ReactDOM.createPortal(, portalContainer)} + {ReactDOM.createPortal( + , + portalContainer + )} ); } @@ -211,7 +210,8 @@ function ActionBarSelection(props) { className={cn( "bg-input flex items-center gap-1 rounded-sm border px-3 py-1 font-medium text-sm tabular-nums shrink-0", className - )} /> + )} + /> ); } @@ -279,64 +279,72 @@ function ActionBarGroup(props) { }); }, []); - const onBlur = React.useCallback((event) => { - onBlurProp?.(event); - if (event.defaultPrevented) return; - - setIsTabbingBackOut(false); - }, [onBlurProp]); - - const onFocus = React.useCallback((event) => { - onFocusProp?.(event); - if (event.defaultPrevented) return; - - const isKeyboardFocus = !isClickFocusRef.current; - if ( - event.target === event.currentTarget && - isKeyboardFocus && - !isTabbingBackOut - ) { - const entryFocusEvent = new CustomEvent(ENTRY_FOCUS, EVENT_OPTIONS); - event.currentTarget.dispatchEvent(entryFocusEvent); - - if (!entryFocusEvent.defaultPrevented) { - const items = Array.from(itemsRef.current.values()).filter((item) => !item.disabled); - const currentItem = items.find((item) => item.id === tabStopId); - - const candidateItems = [currentItem, ...items].filter(Boolean); - const candidateRefs = candidateItems.map((item) => item.ref); - focusFirst(candidateRefs, false); + const onBlur = React.useCallback( + (event) => { + onBlurProp?.(event); + if (event.defaultPrevented) return; + + setIsTabbingBackOut(false); + }, + [onBlurProp] + ); + + const onFocus = React.useCallback( + (event) => { + onFocusProp?.(event); + if (event.defaultPrevented) return; + + const isKeyboardFocus = !isClickFocusRef.current; + if (event.target === event.currentTarget && isKeyboardFocus && !isTabbingBackOut) { + const entryFocusEvent = new CustomEvent(ENTRY_FOCUS, EVENT_OPTIONS); + event.currentTarget.dispatchEvent(entryFocusEvent); + + if (!entryFocusEvent.defaultPrevented) { + const items = Array.from(itemsRef.current.values()).filter((item) => !item.disabled); + const currentItem = items.find((item) => item.id === tabStopId); + + const candidateItems = [currentItem, ...items].filter(Boolean); + const candidateRefs = candidateItems.map((item) => item.ref); + focusFirst(candidateRefs, false); + } } - } - isClickFocusRef.current = false; - }, [onFocusProp, isTabbingBackOut, tabStopId]); - - const onMouseDown = React.useCallback((event) => { - onMouseDownProp?.(event); - if (event.defaultPrevented) return; - - isClickFocusRef.current = true; - }, [onMouseDownProp]); - - const focusContextValue = React.useMemo(() => ({ - tabStopId, - onItemFocus, - onItemShiftTab, - onFocusableItemAdd, - onFocusableItemRemove, - onItemRegister, - onItemUnregister, - getItems, - }), [ - tabStopId, - onItemFocus, - onItemShiftTab, - onFocusableItemAdd, - onFocusableItemRemove, - onItemRegister, - onItemUnregister, - getItems, - ]); + isClickFocusRef.current = false; + }, + [onFocusProp, isTabbingBackOut, tabStopId] + ); + + const onMouseDown = React.useCallback( + (event) => { + onMouseDownProp?.(event); + if (event.defaultPrevented) return; + + isClickFocusRef.current = true; + }, + [onMouseDownProp] + ); + + const focusContextValue = React.useMemo( + () => ({ + tabStopId, + onItemFocus, + onItemShiftTab, + onFocusableItemAdd, + onFocusableItemRemove, + onItemRegister, + onItemUnregister, + getItems, + }), + [ + tabStopId, + onItemFocus, + onItemShiftTab, + onFocusableItemAdd, + onFocusableItemRemove, + onItemRegister, + onItemUnregister, + getItems, + ] + ); const GroupPrimitive = asChild ? Slot : motion.div; @@ -356,14 +364,17 @@ function ActionBarGroup(props) { type: "spring", stiffness: 500, damping: 40, - } + }, }} - className={cn("flex gap-2 outline-none shrink-0", orientation === "horizontal" - ? "items-center" - : "w-full flex-col items-start", className)} + className={cn( + "flex gap-2 outline-none shrink-0", + orientation === "horizontal" ? "items-center" : "w-full flex-col items-start", + className + )} onBlur={onBlur} onFocus={onFocus} - onMouseDown={onMouseDown} /> + onMouseDown={onMouseDown} + /> ); } @@ -385,8 +396,7 @@ function ActionBarItem(props) { const composedRef = useComposedRefs(ref, itemRef); const isMouseClickRef = React.useRef(false); - const { onOpenChange, dir, orientation, loop } = - useActionBarContext(ITEM_NAME); + const { onOpenChange, dir, orientation, loop } = useActionBarContext(ITEM_NAME); const focusContext = useFocusContext(ITEM_NAME); const itemId = React.useId(); @@ -411,97 +421,110 @@ function ActionBarItem(props) { }; }, [focusContext, itemId, disabled]); - const onClick = React.useCallback((event) => { - onClickProp?.(event); - if (event.defaultPrevented) return; - - const item = itemRef.current; - if (!item) return; + const onClick = React.useCallback( + (event) => { + onClickProp?.(event); + if (event.defaultPrevented) return; - const itemSelectEvent = new CustomEvent(ITEM_SELECT, { - bubbles: true, - cancelable: true, - }); - - item.addEventListener(ITEM_SELECT, (event) => onSelect?.(event), { - once: true, - }); + const item = itemRef.current; + if (!item) return; - item.dispatchEvent(itemSelectEvent); + const itemSelectEvent = new CustomEvent(ITEM_SELECT, { + bubbles: true, + cancelable: true, + }); - if (!itemSelectEvent.defaultPrevented) { - onOpenChange?.(false); - } - }, [onClickProp, onOpenChange, onSelect]); + item.addEventListener(ITEM_SELECT, (event) => onSelect?.(event), { + once: true, + }); - const onFocus = React.useCallback((event) => { - onFocusProp?.(event); - if (event.defaultPrevented) return; + item.dispatchEvent(itemSelectEvent); - focusContext.onItemFocus(itemId); - isMouseClickRef.current = false; - }, [onFocusProp, focusContext, itemId]); + if (!itemSelectEvent.defaultPrevented) { + onOpenChange?.(false); + } + }, + [onClickProp, onOpenChange, onSelect] + ); - const onKeyDown = React.useCallback((event) => { - onKeyDownProp?.(event); - if (event.defaultPrevented) return; + const onFocus = React.useCallback( + (event) => { + onFocusProp?.(event); + if (event.defaultPrevented) return; - if (event.key === "Tab" && event.shiftKey) { - focusContext.onItemShiftTab(); - return; - } + focusContext.onItemFocus(itemId); + isMouseClickRef.current = false; + }, + [onFocusProp, focusContext, itemId] + ); - if (event.target !== event.currentTarget) return; - - const key = getDirectionAwareKey(event.key, dir); - let focusIntent; - - if (orientation === "horizontal") { - if (key === "ArrowLeft") focusIntent = "prev"; - else if (key === "ArrowRight") focusIntent = "next"; - else if (key === "Home") focusIntent = "first"; - else if (key === "End") focusIntent = "last"; - } else { - if (key === "ArrowUp") focusIntent = "prev"; - else if (key === "ArrowDown") focusIntent = "next"; - else if (key === "Home") focusIntent = "first"; - else if (key === "End") focusIntent = "last"; - } + const onKeyDown = React.useCallback( + (event) => { + onKeyDownProp?.(event); + if (event.defaultPrevented) return; - if (focusIntent !== undefined) { - if (event.metaKey || event.ctrlKey || event.altKey || event.shiftKey) + if (event.key === "Tab" && event.shiftKey) { + focusContext.onItemShiftTab(); return; - event.preventDefault(); - - const items = focusContext.getItems().filter((item) => !item.disabled); - let candidateRefs = items.map((item) => item.ref); - - if (focusIntent === "last") { - candidateRefs.reverse(); - } else if (focusIntent === "prev" || focusIntent === "next") { - if (focusIntent === "prev") candidateRefs.reverse(); - const currentIndex = candidateRefs.findIndex((ref) => ref.current === event.currentTarget); - candidateRefs = loop - ? wrapArray(candidateRefs, currentIndex + 1) - : candidateRefs.slice(currentIndex + 1); } - queueMicrotask(() => focusFirst(candidateRefs)); - } - }, [onKeyDownProp, focusContext, dir, orientation, loop]); + if (event.target !== event.currentTarget) return; + + const key = getDirectionAwareKey(event.key, dir); + let focusIntent; + + if (orientation === "horizontal") { + if (key === "ArrowLeft") focusIntent = "prev"; + else if (key === "ArrowRight") focusIntent = "next"; + else if (key === "Home") focusIntent = "first"; + else if (key === "End") focusIntent = "last"; + } else { + if (key === "ArrowUp") focusIntent = "prev"; + else if (key === "ArrowDown") focusIntent = "next"; + else if (key === "Home") focusIntent = "first"; + else if (key === "End") focusIntent = "last"; + } + + if (focusIntent !== undefined) { + if (event.metaKey || event.ctrlKey || event.altKey || event.shiftKey) return; + event.preventDefault(); + + const items = focusContext.getItems().filter((item) => !item.disabled); + let candidateRefs = items.map((item) => item.ref); + + if (focusIntent === "last") { + candidateRefs.reverse(); + } else if (focusIntent === "prev" || focusIntent === "next") { + if (focusIntent === "prev") candidateRefs.reverse(); + const currentIndex = candidateRefs.findIndex( + (ref) => ref.current === event.currentTarget + ); + candidateRefs = loop + ? wrapArray(candidateRefs, currentIndex + 1) + : candidateRefs.slice(currentIndex + 1); + } - const onMouseDown = React.useCallback((event) => { - onMouseDownProp?.(event); - if (event.defaultPrevented) return; + queueMicrotask(() => focusFirst(candidateRefs)); + } + }, + [onKeyDownProp, focusContext, dir, orientation, loop] + ); - isMouseClickRef.current = true; + const onMouseDown = React.useCallback( + (event) => { + onMouseDownProp?.(event); + if (event.defaultPrevented) return; - if (disabled) { - event.preventDefault(); - } else { - focusContext.onItemFocus(itemId); - } - }, [onMouseDownProp, focusContext, itemId, disabled]); + isMouseClickRef.current = true; + + if (disabled) { + event.preventDefault(); + } else { + focusContext.onItemFocus(itemId); + } + }, + [onMouseDownProp, focusContext, itemId, disabled] + ); return ( )} @@ -216,35 +220,21 @@ function Alert({ ); } -function AlertTitle({ - className, - ...props -}) { +function AlertTitle({ className, ...props }) { return ( -
+
); } -function AlertIcon({ - children, - className, - ...props -}) { +function AlertIcon({ children, className, ...props }) { return ( -
+
{children}
); } -function AlertToolbar({ - children, - className, - ...props -}) { +function AlertToolbar({ children, className, ...props }) { return (
{children} @@ -252,27 +242,23 @@ function AlertToolbar({ ); } -function AlertDescription({ - className, - ...props -}) { +function AlertDescription({ className, ...props }) { return (
+ className={cn("text-sm [&_p]:leading-relaxed [&_p]:mb-2", className)} + {...props} + /> ); } -function AlertContent({ - className, - ...props -}) { +function AlertContent({ className, ...props }) { return (
+ className={cn("space-y-2 [&_[data-slot=alert-title]]:font-semibold", className)} + {...props} + /> ); } diff --git a/frontend/src/components/ui/avatar.jsx b/frontend/src/components/ui/avatar.jsx index 4e7eb9a..ce93e55 100644 --- a/frontend/src/components/ui/avatar.jsx +++ b/frontend/src/components/ui/avatar.jsx @@ -1,29 +1,26 @@ -import * as React from "react" -import { cn } from "@/lib/utils" +import * as React from "react"; +import { cn } from "@/lib/utils"; const Avatar = React.forwardRef(({ className, ...props }, ref) => (
-)) -Avatar.displayName = "Avatar" +)); +Avatar.displayName = "Avatar"; const AvatarImage = React.forwardRef(({ className, src, onError, ...props }, ref) => { - const [imageError, setImageError] = React.useState(false) + const [imageError, setImageError] = React.useState(false); const handleError = (e) => { - setImageError(true) - if (onError) onError(e) - } + setImageError(true); + if (onError) onError(e); + }; // If no src or image error, don't render the image if (!src || imageError) { - return null + return null; } return ( @@ -34,20 +31,18 @@ const AvatarImage = React.forwardRef(({ className, src, onError, ...props }, ref onError={handleError} {...props} /> - ) -}) -AvatarImage.displayName = "AvatarImage" + ); +}); +AvatarImage.displayName = "AvatarImage"; const AvatarFallback = React.forwardRef(({ className, children, ...props }, ref) => { // If children is provided, use it; otherwise show default fallback - const displayText = children || "U" - + const displayText = children || "U"; + // If it's a single character or initials (2 chars), show as is // Otherwise, show first character of the text - const fallbackText = displayText.length <= 2 - ? displayText - : displayText.charAt(0).toUpperCase() - + const fallbackText = displayText.length <= 2 ? displayText : displayText.charAt(0).toUpperCase(); + return (
{fallbackText}
- ) -}) -AvatarFallback.displayName = "AvatarFallback" + ); +}); +AvatarFallback.displayName = "AvatarFallback"; -export { Avatar, AvatarImage, AvatarFallback } \ No newline at end of file +export { Avatar, AvatarImage, AvatarFallback }; diff --git a/frontend/src/components/ui/badge.jsx b/frontend/src/components/ui/badge.jsx index abda76c..de80b7a 100644 --- a/frontend/src/components/ui/badge.jsx +++ b/frontend/src/components/ui/badge.jsx @@ -1,44 +1,34 @@ -import * as React from "react" -import { Slot } from "@radix-ui/react-slot" +import * as React from "react"; +import { Slot } from "@radix-ui/react-slot"; import { cva } from "class-variance-authority"; -import { cn } from "@/lib/utils" +import { cn } from "@/lib/utils"; const badgeVariants = cva( "inline-flex items-center justify-center rounded-xs border px-2 py-0.5 text-xs font-medium w-fit whitespace-nowrap shrink-0 [&>svg]:size-3 gap-1 [&>svg]:pointer-events-none focus-visible:border-ring focus-visible:ring-ring/50 focus-visible:ring-[3px] aria-invalid:ring-destructive/20 dark:aria-invalid:ring-destructive/40 aria-invalid:border-destructive transition-[color,box-shadow] overflow-hidden", { variants: { variant: { - default: - "border-transparent bg-primary text-primary-foreground [a&]:hover:bg-primary/90", + default: "border-transparent bg-primary text-primary-foreground [a&]:hover:bg-primary/90", secondary: "border-transparent bg-secondary text-secondary-foreground [a&]:hover:bg-secondary/90", destructive: "border-transparent bg-destructive text-white [a&]:hover:bg-destructive/90 focus-visible:ring-destructive/20 dark:focus-visible:ring-destructive/40 dark:bg-destructive/60", - outline: - "text-foreground [a&]:hover:bg-accent [a&]:hover:text-accent-foreground", + outline: "text-foreground [a&]:hover:bg-accent [a&]:hover:text-accent-foreground", }, }, defaultVariants: { variant: "default", }, } -) +); -function Badge({ - className, - variant, - asChild = false, - ...props -}) { - const Comp = asChild ? Slot : "span" +function Badge({ className, variant, asChild = false, ...props }) { + const Comp = asChild ? Slot : "span"; return ( - + ); } -export { Badge, badgeVariants } +export { Badge, badgeVariants }; diff --git a/frontend/src/components/ui/breadcrumb.jsx b/frontend/src/components/ui/breadcrumb.jsx index 96e6f9e..0b65717 100644 --- a/frontend/src/components/ui/breadcrumb.jsx +++ b/frontend/src/components/ui/breadcrumb.jsx @@ -1,19 +1,14 @@ -import * as React from "react" -import { Slot } from "@radix-ui/react-slot" -import { ChevronRight, MoreHorizontal } from "lucide-react" +import * as React from "react"; +import { Slot } from "@radix-ui/react-slot"; +import { ChevronRight, MoreHorizontal } from "lucide-react"; -import { cn } from "@/lib/utils" +import { cn } from "@/lib/utils"; -function Breadcrumb({ - ...props -}) { +function Breadcrumb({ ...props }) { return