Spaces:
Running
Running
| from fastapi import APIRouter, Depends, HTTPException, Response, status, BackgroundTasks | |
| from sqlalchemy import or_, select | |
| from sqlalchemy.ext.asyncio import AsyncSession | |
| from app.auth import ( | |
| clear_auth_cookie, create_session_token, get_current_user, hash_password, | |
| set_auth_cookie, verify_password, create_password_reset_token, | |
| decode_password_reset_token | |
| ) | |
| from app.database import get_db | |
| from app.models.user import ( | |
| User, UserLoginRequest, UserResponse, UserSignupRequest, | |
| PasswordResetRequestBody, PasswordResetConfirmBody, PasswordResetResponse | |
| ) | |
| from app.services.email_service import EmailService | |
| from itsdangerous import BadSignature | |
| router = APIRouter(prefix="/api/auth", tags=["auth"]) | |
| email_service = EmailService() | |
| async def signup(payload: UserSignupRequest, response: Response, db: AsyncSession = Depends(get_db)): | |
| stmt = select(User).where(or_(User.username == payload.username, User.email == payload.email)) | |
| existing = (await db.execute(stmt)).scalars().first() | |
| if existing: | |
| raise HTTPException(status_code=409, detail="Username or email already in use") | |
| user = User( | |
| username=payload.username.strip(), | |
| email=payload.email.strip().lower(), | |
| password_hash=hash_password(payload.password), | |
| ) | |
| db.add(user) | |
| await db.commit() | |
| await db.refresh(user) | |
| token = create_session_token(user.id) | |
| set_auth_cookie(response, token) | |
| return user | |
| async def login(payload: UserLoginRequest, response: Response, db: AsyncSession = Depends(get_db)): | |
| identifier = payload.identifier.strip() | |
| stmt = select(User).where(or_(User.username == identifier, User.email == identifier.lower())) | |
| user = (await db.execute(stmt)).scalars().first() | |
| if not user or not verify_password(payload.password, user.password_hash): | |
| raise HTTPException(status_code=401, detail="Invalid credentials") | |
| token = create_session_token(user.id) | |
| set_auth_cookie(response, token) | |
| return user | |
| async def logout(response: Response): | |
| clear_auth_cookie(response) | |
| return None | |
| async def me(current_user: User = Depends(get_current_user)): | |
| return current_user | |
| async def request_password_reset( | |
| payload: PasswordResetRequestBody, | |
| background_tasks: BackgroundTasks, | |
| db: AsyncSession = Depends(get_db), | |
| ): | |
| """Request a password reset email. Always returns success for security.""" | |
| stmt = select(User).where(User.email == payload.email.lower()) | |
| user = (await db.execute(stmt)).scalars().first() | |
| if user: | |
| reset_token = create_password_reset_token(user.id) | |
| background_tasks.add_task(email_service.send_password_reset_email, user.email, reset_token) | |
| # Always return success to prevent email enumeration attacks | |
| return PasswordResetResponse( | |
| success=True, | |
| message="If an account exists with this email, a password reset link has been sent." | |
| ) | |
| async def confirm_password_reset( | |
| payload: PasswordResetConfirmBody, | |
| db: AsyncSession = Depends(get_db), | |
| ): | |
| """Confirm password reset with token and new password.""" | |
| try: | |
| user_id = decode_password_reset_token(payload.token) | |
| except BadSignature: | |
| raise HTTPException(status_code=400, detail="Invalid or expired reset token") | |
| stmt = select(User).where(User.id == user_id) | |
| user = (await db.execute(stmt)).scalars().first() | |
| if not user: | |
| raise HTTPException(status_code=404, detail="User not found") | |
| user.password_hash = hash_password(payload.new_password) | |
| from datetime import datetime, timezone | |
| user.password_reset_count = datetime.now(timezone.utc) | |
| await db.commit() | |
| return PasswordResetResponse( | |
| success=True, | |
| message="Password has been reset successfully. Please login with your new password." | |
| ) | |
| async def change_password( | |
| old_password: str, | |
| new_password: str, | |
| current_user: User = Depends(get_current_user), | |
| db: AsyncSession = Depends(get_db), | |
| ): | |
| """Change password for authenticated user.""" | |
| if not verify_password(old_password, current_user.password_hash): | |
| raise HTTPException(status_code=401, detail="Current password is incorrect") | |
| current_user.password_hash = hash_password(new_password) | |
| from datetime import datetime, timezone | |
| current_user.password_reset_count = datetime.now(timezone.utc) | |
| await db.commit() | |
| return {"success": True, "message": "Password changed successfully"} | |