import hashlib import uuid from datetime import UTC, datetime, timedelta import pytest from httpx import AsyncClient from sqlalchemy import func, select from sqlalchemy.orm import Session from app.core.security import get_password_hash, verify_password from app.core.session import session as session_maker from app.models import PasswordResetToken, User pytestmark = pytest.mark.asyncio @pytest.fixture def existing_user() -> User: with session_maker() as db: user = User( email="perceval@test.com", hashed_password=get_password_hash("OldPassword123"), name="Perceval", ) db.add(user) db.commit() db.refresh(user) return user def make_token( session: Session, user: User, *, purpose="reset", expires_in=timedelta(hours=1), used=False ): raw = "raw-token-" + uuid.uuid4().hex token_hash = hashlib.sha256(raw.encode()).hexdigest() record = PasswordResetToken( user_id=user.id, token_hash=token_hash, expires_at=datetime.now(UTC) + expires_in, purpose=purpose, used_at=datetime.now(UTC) if used else None, ) session.add(record) session.commit() return raw class TestForgotPassword: async def test_existing_email_returns_204_and_creates_token( self, client: AsyncClient, existing_user: User, session: Session ): response = await client.post("/users/forgot-password", json={"email": existing_user.email}) assert response.status_code == 204 token = session.scalar( select(PasswordResetToken).where(PasswordResetToken.user_id == existing_user.id) ) assert token is not None assert token.purpose == "reset" async def test_unknown_email_also_returns_204_no_enumeration(self, client: AsyncClient): response = await client.post("/users/forgot-password", json={"email": "nobody@test.com"}) assert response.status_code == 204 async def test_unknown_email_creates_no_token(self, client: AsyncClient, session: Session): await client.post("/users/forgot-password", json={"email": "nobody@test.com"}) assert session.scalar(select(PasswordResetToken)) is None async def test_response_body_identical_for_both_cases( self, client: AsyncClient, existing_user: User ): # guards against a future refactor leaking a distinguishable signal r1 = await client.post("/users/forgot-password", json={"email": existing_user.email}) r2 = await client.post("/users/forgot-password", json={"email": "nobody@test.com"}) assert r1.status_code == r2.status_code == 204 assert r1.content == r2.content == b"" class TestSetPasswordWithToken: async def test_valid_token_sets_password( self, client: AsyncClient, existing_user: User, session: Session ): raw_token = make_token(session, existing_user) response = await client.post( "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"} ) assert response.status_code == 200 refreshed = session.get(User, existing_user.id) assert verify_password("NewSecurePass456", refreshed.hashed_password) async def test_valid_token_clears_must_change_password( self, client: AsyncClient, existing_user: User, session: Session ): existing_user.must_change_password = True session.merge(existing_user) session.commit() raw_token = make_token(session, existing_user) response = await client.post( "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"} ) assert response.status_code == 200 assert session.get(User, existing_user.id).must_change_password is False async def test_token_marked_used_after_success( self, client: AsyncClient, existing_user: User, session: Session ): raw_token = make_token(session, existing_user) await client.post( "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"} ) token_hash = hashlib.sha256(raw_token.encode()).hexdigest() record = session.scalar( select(PasswordResetToken).where(PasswordResetToken.token_hash == token_hash) ) assert record.used_at is not None async def test_token_cannot_be_reused( self, client: AsyncClient, existing_user: User, session: Session ): raw_token = make_token(session, existing_user) await client.post( "/users/set-password", json={"token": raw_token, "password": "First12345"} ) response = await client.post( "/users/set-password", json={"token": raw_token, "password": "Second67890"} ) assert response.status_code == 400 async def test_expired_token_rejected( self, client: AsyncClient, existing_user: User, session: Session ): raw_token = make_token(session, existing_user, expires_in=timedelta(hours=-1)) response = await client.post( "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"} ) assert response.status_code == 400 async def test_already_used_token_rejected( self, client: AsyncClient, existing_user: User, session: Session ): raw_token = make_token(session, existing_user, used=True) response = await client.post( "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"} ) assert response.status_code == 400 async def test_unknown_token_rejected(self, client: AsyncClient): response = await client.post( "/users/set-password", json={"token": "not-a-real-token", "password": "NewSecurePass456"}, ) assert response.status_code == 400 async def test_invite_purpose_token_also_works( self, client: AsyncClient, existing_user: User, session: Session ): raw_token = make_token( session, existing_user, purpose="invite", expires_in=timedelta(days=7) ) response = await client.post( "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"} ) assert response.status_code == 200 class TestForgotPasswordRateLimit: async def test_exceeding_limit_stops_sending_but_still_204s( self, client: AsyncClient, existing_user: User, session: Session ): for _ in range(3): await client.post("/users/forgot-password", json={"email": existing_user.email}) tokens_before = session.scalar( select(func.count()) .select_from(PasswordResetToken) .where(PasswordResetToken.user_id == existing_user.id) ) response = await client.post("/users/forgot-password", json={"email": existing_user.email}) tokens_after = session.scalar( select(func.count()) .select_from(PasswordResetToken) .where(PasswordResetToken.user_id == existing_user.id) ) assert response.status_code == 204 assert tokens_after == tokens_before