فهرست منبع

test: count only the racing test's own archive folders (#2957)

Other xdist workers archive into the same printer directory, sometimes
within the same second, and the race test counted their folders too.
maziggy 3 روز پیش
والد
کامیت
10ed899130
3فایلهای تغییر یافته به همراه99 افزوده شده و 31 حذف شده
  1. 26 30
      backend/app/api/routes/mfa.py
  2. 65 0
      backend/tests/integration/test_mfa_api.py
  3. 8 1
      backend/tests/unit/test_fallback_archive_recovery_2957.py

+ 26 - 30
backend/app/api/routes/mfa.py

@@ -421,27 +421,17 @@ async def record_email_otp_send(db: AsyncSession, username: str) -> None:
 # ---------------------------------------------------------------------------
 # ---------------------------------------------------------------------------
 # TOTP replay-protection helper
 # TOTP replay-protection helper
 # ---------------------------------------------------------------------------
 # ---------------------------------------------------------------------------
-def _assert_totp_not_replayed(totp_obj: pyotp.TOTP, totp_record: UserTOTP, code: str) -> None:
-    """Raise HTTP 400 if this TOTP code was already accepted in its time window.
+def _matched_totp_counter(totp_obj: pyotp.TOTP, code: str) -> int | None:
+    """The time step whose code matches, one step either side of now, or None.
 
 
-    M3 fix: store the counter of the *accepted* code rather than the current
-    wall-clock counter.  With valid_window=1, pyotp accepts codes from the
-    previous 30-second step.  Using timecode(now) would store the wrong counter
-    when the previous-window code is accepted, allowing immediate replay.
+    The caller records the returned step with ``accept_counter``. Matching and
+    naming the step share one clock reading.
     """
     """
-    # Determine which time-step the accepted code belongs to.
-    now = datetime.now(timezone.utc)
-    accepted_counter: int | None = None
-    for offset in (0, -1):  # current window first, then previous
-        candidate_time = now.timestamp() + offset * totp_obj.interval
-        candidate_counter = totp_obj.timecode(datetime.fromtimestamp(candidate_time, tz=timezone.utc))
-        if totp_obj.at(candidate_counter) == code:
-            accepted_counter = candidate_counter
-            break
-    if accepted_counter is None:
-        accepted_counter = totp_obj.timecode(now)  # fallback (should not happen after verify())
-
-    totp_record.accept_counter(accepted_counter)
+    current = totp_obj.timecode(datetime.now(timezone.utc))
+    for counter in (current, current - 1, current + 1):
+        if pyotp.utils.strings_equal(str(code), totp_obj.generate_otp(counter)):
+            return counter
+    return None
 
 
 
 
 # ---------------------------------------------------------------------------
 # ---------------------------------------------------------------------------
@@ -679,7 +669,7 @@ async def setup_totp(
         # S4: narrow the RuntimeError catch to ONLY the property access — that
         # S4: narrow the RuntimeError catch to ONLY the property access — that
         # is the single line that raises on key-loss. The previous wide try
         # is the single line that raises on key-loss. The previous wide try
         # block also covered record_failed_attempt, clear_failed_attempts,
         # block also covered record_failed_attempt, clear_failed_attempts,
-        # and _assert_totp_not_replayed, so a future RuntimeError from any
+        # and the replay guard, so a future RuntimeError from any
         # of those would have been misreported as "TOTP secret unavailable".
         # of those would have been misreported as "TOTP secret unavailable".
         try:
         try:
             secret_plain = existing.secret
             secret_plain = existing.secret
@@ -689,14 +679,15 @@ async def setup_totp(
                 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
                 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
                 detail="TOTP secret unavailable",
                 detail="TOTP secret unavailable",
             )
             )
-        if not pyotp.TOTP(secret_plain).verify(supplied_code, valid_window=1):
+        counter = _matched_totp_counter(pyotp.TOTP(secret_plain), supplied_code)
+        if counter is None:
             await record_failed_attempt(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
             await record_failed_attempt(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
             raise HTTPException(
             raise HTTPException(
                 status_code=status.HTTP_400_BAD_REQUEST,
                 status_code=status.HTTP_400_BAD_REQUEST,
                 detail="Current TOTP code required to replace an active authenticator",
                 detail="Current TOTP code required to replace an active authenticator",
             )
             )
         await clear_failed_attempts(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
         await clear_failed_attempts(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
-        _assert_totp_not_replayed(pyotp.TOTP(secret_plain), existing, supplied_code)
+        existing.accept_counter(counter)
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
 
 
     secret = pyotp.random_base32()
     secret = pyotp.random_base32()
@@ -782,10 +773,12 @@ async def disable_totp(
     # code path so the user can still disable 2FA with their printed codes.
     # code path so the user can still disable 2FA with their printed codes.
     totp_obj: pyotp.TOTP | None = None
     totp_obj: pyotp.TOTP | None = None
     code_valid = False
     code_valid = False
+    counter: int | None = None
     decryption_failed = False
     decryption_failed = False
     try:
     try:
         totp_obj = pyotp.TOTP(totp_record.secret)
         totp_obj = pyotp.TOTP(totp_record.secret)
-        code_valid = totp_obj.verify(body.code, valid_window=1)
+        counter = _matched_totp_counter(totp_obj, body.code)
+        code_valid = counter is not None
     except RuntimeError:
     except RuntimeError:
         # S3: track that the failure was server-side so we don't penalise
         # S3: track that the failure was server-side so we don't penalise
         # the user with a fail-counter increment for a problem they can't fix.
         # the user with a fail-counter increment for a problem they can't fix.
@@ -795,8 +788,8 @@ async def disable_totp(
             totp_record.user_id,
             totp_record.user_id,
         )
         )
 
 
-    if code_valid and totp_obj is not None:
-        _assert_totp_not_replayed(totp_obj, totp_record, body.code)
+    if code_valid and counter is not None:
+        totp_record.accept_counter(counter)
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
     else:
     else:
         # Check backup codes — always iterate all entries (L-R9-A: no early break
         # Check backup codes — always iterate all entries (L-R9-A: no early break
@@ -844,10 +837,12 @@ async def regenerate_backup_codes(
     # rotate their codes with a printed backup code.
     # rotate their codes with a printed backup code.
     totp_obj: pyotp.TOTP | None = None
     totp_obj: pyotp.TOTP | None = None
     code_valid = False
     code_valid = False
+    counter: int | None = None
     decryption_failed = False
     decryption_failed = False
     try:
     try:
         totp_obj = pyotp.TOTP(totp_record.secret)
         totp_obj = pyotp.TOTP(totp_record.secret)
-        code_valid = totp_obj.verify(body.code, valid_window=1)
+        counter = _matched_totp_counter(totp_obj, body.code)
+        code_valid = counter is not None
     except RuntimeError:
     except RuntimeError:
         # S3: track server-side failure so we skip the fail-counter debit.
         # S3: track server-side failure so we skip the fail-counter debit.
         decryption_failed = True
         decryption_failed = True
@@ -856,8 +851,8 @@ async def regenerate_backup_codes(
             totp_record.user_id,
             totp_record.user_id,
         )
         )
 
 
-    if code_valid and totp_obj is not None:
-        _assert_totp_not_replayed(totp_obj, totp_record, body.code)
+    if code_valid and counter is not None:
+        totp_record.accept_counter(counter)
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
     else:
     else:
         # Accept a backup code as an alternative (M10)
         # Accept a backup code as an alternative (M10)
@@ -1170,10 +1165,11 @@ async def verify_2fa(
         except RuntimeError:
         except RuntimeError:
             logger.exception("TOTP decryption failed for user_id=%s", totp_record.user_id)
             logger.exception("TOTP decryption failed for user_id=%s", totp_record.user_id)
             raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="TOTP secret unavailable")
             raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="TOTP secret unavailable")
-        if not totp_obj.verify(body.code, valid_window=1):
+        counter = _matched_totp_counter(totp_obj, body.code)
+        if counter is None:
             await record_failed_attempt(db, username)
             await record_failed_attempt(db, username)
             raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid TOTP code")
             raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid TOTP code")
-        _assert_totp_not_replayed(totp_obj, totp_record, body.code)
+        totp_record.accept_counter(counter)
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
         await db.flush()  # L-3: persist last_totp_counter immediately to block replay
 
 
     elif method == "email":
     elif method == "email":

+ 65 - 0
backend/tests/integration/test_mfa_api.py

@@ -1332,6 +1332,71 @@ class TestTOTPReplay:
         )
         )
         assert second.status_code == 400
         assert second.status_code == 400
 
 
+    @staticmethod
+    def _freeze_mfa_clock(monkeypatch, at: float) -> type:
+        """Pin the clock the MFA routes read; set ``clock.at`` to move it."""
+        from backend.app.api.routes import mfa as mfa_module
+
+        class Clock(datetime):
+            @classmethod
+            def now(cls, tz=None):
+                return datetime.fromtimestamp(Clock.at, tz)
+
+        Clock.at = at
+        monkeypatch.setattr(mfa_module, "datetime", Clock)
+        return Clock
+
+    async def _verify(self, client: AsyncClient, username: str, password: str, code: str):
+        pre_auth = await _login_get_pre_auth_token(client, username, password)
+        return await client.post(
+            "/api/v1/auth/2fa/verify",
+            json={"pre_auth_token": pre_auth, "method": "totp", "code": code},
+        )
+
+    @pytest.mark.asyncio
+    @pytest.mark.integration
+    async def test_totp_replay_in_the_next_step_is_rejected(self, async_client: AsyncClient, monkeypatch):
+        """A code stays valid into the next 30-second step; its reuse there must still fail."""
+        _token, secret = await _setup_totp_user(async_client, "replaynext", "replaynext1")
+        totp = pyotp.TOTP(secret)
+        step_start = (int(time.time()) // 30 + 1) * 30
+        code = totp.at(step_start)
+        clock = self._freeze_mfa_clock(monkeypatch, step_start + 29)
+
+        first = await self._verify(async_client, "replaynext", "replaynext1", code)
+        assert first.status_code == 200
+
+        clock.at = step_start + 31
+        second = await self._verify(async_client, "replaynext", "replaynext1", code)
+        assert second.status_code == 400
+        assert second.json()["detail"] == "TOTP code already used"
+
+    @pytest.mark.asyncio
+    @pytest.mark.integration
+    async def test_a_previous_step_code_records_its_own_step(self, async_client: AsyncClient, monkeypatch):
+        """Used first in the next step, the code is recorded under its own step: the
+        current step's code is still accepted afterwards, the replay is not."""
+        _token, secret = await _setup_totp_user(async_client, "prevstep", "prevstep12")
+        totp = pyotp.TOTP(secret)
+        step_start = (int(time.time()) // 30 + 1) * 30
+        old_code, new_code = totp.at(step_start), totp.at(step_start + 30)
+        self._freeze_mfa_clock(monkeypatch, step_start + 35)
+
+        assert (await self._verify(async_client, "prevstep", "prevstep12", old_code)).status_code == 200
+        assert (await self._verify(async_client, "prevstep", "prevstep12", old_code)).status_code == 400
+        assert (await self._verify(async_client, "prevstep", "prevstep12", new_code)).status_code == 200
+
+    @pytest.mark.asyncio
+    @pytest.mark.integration
+    async def test_a_code_two_steps_old_is_rejected(self, async_client: AsyncClient, monkeypatch):
+        _token, secret = await _setup_totp_user(async_client, "stalecode", "stalecode12")
+        step_start = (int(time.time()) // 30 + 1) * 30
+        code = pyotp.TOTP(secret).at(step_start)
+        self._freeze_mfa_clock(monkeypatch, step_start + 61)
+
+        resp = await self._verify(async_client, "stalecode", "stalecode12", code)
+        assert resp.status_code == 401
+
     @pytest.mark.asyncio
     @pytest.mark.asyncio
     @pytest.mark.integration
     @pytest.mark.integration
     async def test_totp_replay_rejected_on_disable(self, async_client: AsyncClient):
     async def test_totp_replay_rejected_on_disable(self, async_client: AsyncClient):

+ 8 - 1
backend/tests/unit/test_fallback_archive_recovery_2957.py

@@ -469,7 +469,14 @@ class TestConcurrentRecoveryIsSerialised:
 
 
         # Exactly one caller did the work; the rest saw a recovered archive.
         # Exactly one caller did the work; the rest saw a recovered archive.
         assert results.count(True) == 1
         assert results.count(True) == 1
-        created = (set(printer_root.iterdir()) if printer_root.exists() else set()) - before
+        # Only this run's folders: tests in other xdist workers archive into the
+        # same printer directory, sometimes within the same second.
+        stem = unique.removesuffix(".gcode.3mf")
+        created = {
+            p
+            for p in (set(printer_root.iterdir()) if printer_root.exists() else set()) - before
+            if p.name.endswith(stem)
+        }
         assert len(created) == 1, f"expected one archive directory, got {sorted(p.name for p in created)}"
         assert len(created) == 1, f"expected one archive directory, got {sorted(p.name for p in created)}"
 
 
         async with maker() as db:
         async with maker() as db: