test_password_reset.py 7.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184
  1. import hashlib
  2. import uuid
  3. from datetime import UTC, datetime, timedelta
  4. import pytest
  5. from httpx import AsyncClient
  6. from sqlalchemy import func, select
  7. from sqlalchemy.orm import Session
  8. from app.core.security import get_password_hash, verify_password
  9. from app.core.session import session as session_maker
  10. from app.models import PasswordResetToken, User
  11. pytestmark = pytest.mark.asyncio
  12. @pytest.fixture
  13. def existing_user() -> User:
  14. with session_maker() as db:
  15. user = User(
  16. email="perceval@test.com",
  17. hashed_password=get_password_hash("OldPassword123"),
  18. name="Perceval",
  19. )
  20. db.add(user)
  21. db.commit()
  22. db.refresh(user)
  23. return user
  24. def make_token(
  25. session: Session, user: User, *, purpose="reset", expires_in=timedelta(hours=1), used=False
  26. ):
  27. raw = "raw-token-" + uuid.uuid4().hex
  28. token_hash = hashlib.sha256(raw.encode()).hexdigest()
  29. record = PasswordResetToken(
  30. user_id=user.id,
  31. token_hash=token_hash,
  32. expires_at=datetime.now(UTC) + expires_in,
  33. purpose=purpose,
  34. used_at=datetime.now(UTC) if used else None,
  35. )
  36. session.add(record)
  37. session.commit()
  38. return raw
  39. class TestForgotPassword:
  40. async def test_existing_email_returns_204_and_creates_token(
  41. self, client: AsyncClient, existing_user: User, session: Session
  42. ):
  43. response = await client.post("/users/forgot-password", json={"email": existing_user.email})
  44. assert response.status_code == 204
  45. token = session.scalar(
  46. select(PasswordResetToken).where(PasswordResetToken.user_id == existing_user.id)
  47. )
  48. assert token is not None
  49. assert token.purpose == "reset"
  50. async def test_unknown_email_also_returns_204_no_enumeration(self, client: AsyncClient):
  51. response = await client.post("/users/forgot-password", json={"email": "nobody@test.com"})
  52. assert response.status_code == 204
  53. async def test_unknown_email_creates_no_token(self, client: AsyncClient, session: Session):
  54. await client.post("/users/forgot-password", json={"email": "nobody@test.com"})
  55. assert session.scalar(select(PasswordResetToken)) is None
  56. async def test_response_body_identical_for_both_cases(
  57. self, client: AsyncClient, existing_user: User
  58. ):
  59. # guards against a future refactor leaking a distinguishable signal
  60. r1 = await client.post("/users/forgot-password", json={"email": existing_user.email})
  61. r2 = await client.post("/users/forgot-password", json={"email": "nobody@test.com"})
  62. assert r1.status_code == r2.status_code == 204
  63. assert r1.content == r2.content == b""
  64. class TestSetPasswordWithToken:
  65. async def test_valid_token_sets_password(
  66. self, client: AsyncClient, existing_user: User, session: Session
  67. ):
  68. raw_token = make_token(session, existing_user)
  69. response = await client.post(
  70. "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"}
  71. )
  72. assert response.status_code == 200
  73. refreshed = session.get(User, existing_user.id)
  74. assert verify_password("NewSecurePass456", refreshed.hashed_password)
  75. async def test_valid_token_clears_must_change_password(
  76. self, client: AsyncClient, existing_user: User, session: Session
  77. ):
  78. existing_user.must_change_password = True
  79. session.merge(existing_user)
  80. session.commit()
  81. raw_token = make_token(session, existing_user)
  82. response = await client.post(
  83. "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"}
  84. )
  85. assert response.status_code == 200
  86. assert session.get(User, existing_user.id).must_change_password is False
  87. async def test_token_marked_used_after_success(
  88. self, client: AsyncClient, existing_user: User, session: Session
  89. ):
  90. raw_token = make_token(session, existing_user)
  91. await client.post(
  92. "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"}
  93. )
  94. token_hash = hashlib.sha256(raw_token.encode()).hexdigest()
  95. record = session.scalar(
  96. select(PasswordResetToken).where(PasswordResetToken.token_hash == token_hash)
  97. )
  98. assert record.used_at is not None
  99. async def test_token_cannot_be_reused(
  100. self, client: AsyncClient, existing_user: User, session: Session
  101. ):
  102. raw_token = make_token(session, existing_user)
  103. await client.post(
  104. "/users/set-password", json={"token": raw_token, "password": "First12345"}
  105. )
  106. response = await client.post(
  107. "/users/set-password", json={"token": raw_token, "password": "Second67890"}
  108. )
  109. assert response.status_code == 400
  110. async def test_expired_token_rejected(
  111. self, client: AsyncClient, existing_user: User, session: Session
  112. ):
  113. raw_token = make_token(session, existing_user, expires_in=timedelta(hours=-1))
  114. response = await client.post(
  115. "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"}
  116. )
  117. assert response.status_code == 400
  118. async def test_already_used_token_rejected(
  119. self, client: AsyncClient, existing_user: User, session: Session
  120. ):
  121. raw_token = make_token(session, existing_user, used=True)
  122. response = await client.post(
  123. "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"}
  124. )
  125. assert response.status_code == 400
  126. async def test_unknown_token_rejected(self, client: AsyncClient):
  127. response = await client.post(
  128. "/users/set-password",
  129. json={"token": "not-a-real-token", "password": "NewSecurePass456"},
  130. )
  131. assert response.status_code == 400
  132. async def test_invite_purpose_token_also_works(
  133. self, client: AsyncClient, existing_user: User, session: Session
  134. ):
  135. raw_token = make_token(
  136. session, existing_user, purpose="invite", expires_in=timedelta(days=7)
  137. )
  138. response = await client.post(
  139. "/users/set-password", json={"token": raw_token, "password": "NewSecurePass456"}
  140. )
  141. assert response.status_code == 200
  142. class TestForgotPasswordRateLimit:
  143. async def test_exceeding_limit_stops_sending_but_still_204s(
  144. self, client: AsyncClient, existing_user: User, session: Session
  145. ):
  146. for _ in range(3):
  147. await client.post("/users/forgot-password", json={"email": existing_user.email})
  148. tokens_before = session.scalar(
  149. select(func.count())
  150. .select_from(PasswordResetToken)
  151. .where(PasswordResetToken.user_id == existing_user.id)
  152. )
  153. response = await client.post("/users/forgot-password", json={"email": existing_user.email})
  154. tokens_after = session.scalar(
  155. select(func.count())
  156. .select_from(PasswordResetToken)
  157. .where(PasswordResetToken.user_id == existing_user.id)
  158. )
  159. assert response.status_code == 204
  160. assert tokens_after == tokens_before