| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184 |
- 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
|