SameerMesiah97 commented on code in PR #74364:
URL: https://github.com/apache/airflow/pull/74364#discussion_r4199198081


##########
airflow-core/src/airflow/models/revoked_token.py:
##########
@@ -71,9 +80,26 @@ def _maybe_cleanup_expired(cls, session: Session) -> None:
         """
         now = time.monotonic()
         cleanup_interval = conf.getint("api_auth", "jwt_expiration_time", 
fallback=3600) * 2
-        if now - cls._last_cleanup_time >= cleanup_interval:
+        if now - cls._last_cleanup_time < cleanup_interval:
+            return
+        if not cls._cleanup_lock.acquire(blocking=False):
+            return
+        try:
+            # Set before the delete so a failing database is not retried on 
every request.
             cls._last_cleanup_time = now
-            try:
-                session.execute(delete(cls).where(cls.exp < 
datetime.now(tz=timezone.utc)))
-            except Exception:
-                log.exception("Failed to clean up expired revoked tokens")
+            expired_jtis = session.scalars(

Review Comment:
   I think there is a slight issue here. For example, consider the following 
sequence of events:
   
   1. Thread A acquires the lock and completes cleanup, updating 
`_last_cleanup_time`.
   2. Thread B was paused after its interval check.
   3. B resumes after A releases the lock, acquires it and runs cleanup again 
(even though cleanup is no longer due).
   
   I believe you should check for the cleanup time right after the line 87 
(just above `cls._last_cleanup_time = now`) like this:
   
   ```
   now = time.monotonic()
       if now - cls._last_cleanup_time < cleanup_interval:
           return
   ```
   
   



##########
airflow-core/tests/unit/models/test_revoked_token.py:
##########
@@ -97,7 +104,117 @@ def test_cleanup_skips_when_interval_not_passed(self):
             ):
                 RevokedToken.is_revoked("test-jti", session=mock_session)
 
-            # session.execute should NOT be called
+            mock_session.scalars.assert_not_called()
+            mock_session.execute.assert_not_called()
+        finally:
+            RevokedToken._last_cleanup_time = original_last_cleanup
+
+    def test_cleanup_skipped_while_another_thread_is_cleaning(self):
+        """The interval bookkeeping is not thread safe, so only one pass may 
run at a time."""
+        mock_session = MagicMock()
+        mock_session.scalar.return_value = False
+
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        RevokedToken._cleanup_lock.acquire()
+        try:
+            RevokedToken._last_cleanup_time = 0.0
+            with (
+                patch("airflow.models.revoked_token.time.monotonic", 
return_value=8000.0),
+                patch("airflow.models.revoked_token.conf.getint", 
return_value=3600),
+            ):
+                assert RevokedToken.is_revoked("test-jti", 
session=mock_session) is False
+
+            mock_session.scalars.assert_not_called()
             mock_session.execute.assert_not_called()
+            # a skipped pass must not claim the interval either
+            assert RevokedToken._last_cleanup_time == 0.0
         finally:
+            RevokedToken._cleanup_lock.release()
             RevokedToken._last_cleanup_time = original_last_cleanup
+
+    def 
test_failed_cleanup_rolls_back_so_the_revocation_read_still_works(self):

Review Comment:
   I think this only covers rollback for a the failed select query. Not the 
delete. I would either parameterize this test (if possible) to cover the delete 
rollback or add a new test specifically for that codepath.



##########
airflow-core/tests/unit/models/test_revoked_token.py:
##########
@@ -97,7 +104,117 @@ def test_cleanup_skips_when_interval_not_passed(self):
             ):
                 RevokedToken.is_revoked("test-jti", session=mock_session)
 
-            # session.execute should NOT be called
+            mock_session.scalars.assert_not_called()
+            mock_session.execute.assert_not_called()
+        finally:
+            RevokedToken._last_cleanup_time = original_last_cleanup
+
+    def test_cleanup_skipped_while_another_thread_is_cleaning(self):
+        """The interval bookkeeping is not thread safe, so only one pass may 
run at a time."""
+        mock_session = MagicMock()
+        mock_session.scalar.return_value = False
+
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        RevokedToken._cleanup_lock.acquire()
+        try:
+            RevokedToken._last_cleanup_time = 0.0
+            with (
+                patch("airflow.models.revoked_token.time.monotonic", 
return_value=8000.0),
+                patch("airflow.models.revoked_token.conf.getint", 
return_value=3600),
+            ):
+                assert RevokedToken.is_revoked("test-jti", 
session=mock_session) is False
+
+            mock_session.scalars.assert_not_called()
             mock_session.execute.assert_not_called()
+            # a skipped pass must not claim the interval either
+            assert RevokedToken._last_cleanup_time == 0.0
         finally:
+            RevokedToken._cleanup_lock.release()
             RevokedToken._last_cleanup_time = original_last_cleanup
+
+    def 
test_failed_cleanup_rolls_back_so_the_revocation_read_still_works(self):
+        """A failed statement aborts the transaction on PostgreSQL; the read 
after it must not inherit that."""
+        mock_session = MagicMock()
+        mock_session.scalar.return_value = False
+        mock_session.scalars.side_effect = RuntimeError("database is on fire")
+
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        try:
+            RevokedToken._last_cleanup_time = 0.0
+            with (
+                patch("airflow.models.revoked_token.time.monotonic", 
return_value=8000.0),
+                patch("airflow.models.revoked_token.conf.getint", 
return_value=3600),
+            ):
+                assert RevokedToken.is_revoked("test-jti", 
session=mock_session) is False
+
+            mock_session.rollback.assert_called_once()
+            # the lock must not stay held after a failure
+            assert RevokedToken._cleanup_lock.acquire(blocking=False)
+            RevokedToken._cleanup_lock.release()
+        finally:
+            RevokedToken._last_cleanup_time = original_last_cleanup
+
+
[email protected]_test
+class TestRevokedTokenCleanupIsBounded:
+    """Cleanup runs on the request path, so a single pass must not issue an 
unbounded DELETE."""
+
+    @pytest.fixture(autouse=True)
+    def reset_cleanup_state(self):
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        with create_session() as session:
+            session.execute(delete(RevokedToken))
+        yield
+        RevokedToken._last_cleanup_time = original_last_cleanup
+        with create_session() as session:
+            session.execute(delete(RevokedToken))
+
+    @staticmethod
+    def _add_tokens(expired: int, live: int) -> None:
+        now = datetime.now(tz=timezone.utc)
+        with create_session() as session:
+            for i in range(expired):
+                session.add(RevokedToken(jti=f"expired-{i}", exp=now - 
timedelta(hours=1)))
+            for i in range(live):
+                session.add(RevokedToken(jti=f"live-{i}", exp=now + 
timedelta(hours=1)))
+
+    @staticmethod
+    def _remaining() -> int:
+        with create_session() as session:
+            return 
session.scalars(select(func.count()).select_from(RevokedToken)).one()
+
+    @conf_vars({("api_auth", "jwt_expiration_time"): "3600"})
+    def test_cleanup_deletes_at_most_one_batch(self):
+        self._add_tokens(expired=7, live=2)
+        RevokedToken._last_cleanup_time = 0.0
+
+        with (
+            patch("airflow.models.revoked_token._CLEANUP_BATCH_SIZE", 3),
+            patch("airflow.models.revoked_token.time.monotonic", 
return_value=100_000.0),
+        ):
+            RevokedToken.is_revoked("live-0")
+
+        # 3 of the 7 expired rows gone, both unexpired rows untouched
+        assert self._remaining() == 6

Review Comment:
   This assertion is too vague. What I would do is query the live tokesn and 
make sure they are still there and assert like this:
   
   ```
   assert "live-0" in remaining_jtis
   assert "live-1" in remaining_jtis
   ```



##########
airflow-core/tests/unit/models/test_revoked_token.py:
##########
@@ -97,7 +104,117 @@ def test_cleanup_skips_when_interval_not_passed(self):
             ):
                 RevokedToken.is_revoked("test-jti", session=mock_session)
 
-            # session.execute should NOT be called
+            mock_session.scalars.assert_not_called()
+            mock_session.execute.assert_not_called()
+        finally:
+            RevokedToken._last_cleanup_time = original_last_cleanup
+
+    def test_cleanup_skipped_while_another_thread_is_cleaning(self):
+        """The interval bookkeeping is not thread safe, so only one pass may 
run at a time."""
+        mock_session = MagicMock()
+        mock_session.scalar.return_value = False
+
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        RevokedToken._cleanup_lock.acquire()
+        try:
+            RevokedToken._last_cleanup_time = 0.0
+            with (
+                patch("airflow.models.revoked_token.time.monotonic", 
return_value=8000.0),
+                patch("airflow.models.revoked_token.conf.getint", 
return_value=3600),
+            ):
+                assert RevokedToken.is_revoked("test-jti", 
session=mock_session) is False
+
+            mock_session.scalars.assert_not_called()
             mock_session.execute.assert_not_called()
+            # a skipped pass must not claim the interval either
+            assert RevokedToken._last_cleanup_time == 0.0
         finally:
+            RevokedToken._cleanup_lock.release()
             RevokedToken._last_cleanup_time = original_last_cleanup
+
+    def 
test_failed_cleanup_rolls_back_so_the_revocation_read_still_works(self):
+        """A failed statement aborts the transaction on PostgreSQL; the read 
after it must not inherit that."""
+        mock_session = MagicMock()
+        mock_session.scalar.return_value = False
+        mock_session.scalars.side_effect = RuntimeError("database is on fire")
+
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        try:
+            RevokedToken._last_cleanup_time = 0.0
+            with (
+                patch("airflow.models.revoked_token.time.monotonic", 
return_value=8000.0),
+                patch("airflow.models.revoked_token.conf.getint", 
return_value=3600),
+            ):
+                assert RevokedToken.is_revoked("test-jti", 
session=mock_session) is False
+
+            mock_session.rollback.assert_called_once()
+            # the lock must not stay held after a failure
+            assert RevokedToken._cleanup_lock.acquire(blocking=False)
+            RevokedToken._cleanup_lock.release()
+        finally:
+            RevokedToken._last_cleanup_time = original_last_cleanup
+
+
[email protected]_test
+class TestRevokedTokenCleanupIsBounded:
+    """Cleanup runs on the request path, so a single pass must not issue an 
unbounded DELETE."""
+
+    @pytest.fixture(autouse=True)
+    def reset_cleanup_state(self):
+        original_last_cleanup = RevokedToken._last_cleanup_time
+        with create_session() as session:
+            session.execute(delete(RevokedToken))
+        yield
+        RevokedToken._last_cleanup_time = original_last_cleanup
+        with create_session() as session:
+            session.execute(delete(RevokedToken))
+
+    @staticmethod
+    def _add_tokens(expired: int, live: int) -> None:
+        now = datetime.now(tz=timezone.utc)
+        with create_session() as session:
+            for i in range(expired):
+                session.add(RevokedToken(jti=f"expired-{i}", exp=now - 
timedelta(hours=1)))
+            for i in range(live):
+                session.add(RevokedToken(jti=f"live-{i}", exp=now + 
timedelta(hours=1)))
+
+    @staticmethod
+    def _remaining() -> int:
+        with create_session() as session:
+            return 
session.scalars(select(func.count()).select_from(RevokedToken)).one()
+
+    @conf_vars({("api_auth", "jwt_expiration_time"): "3600"})
+    def test_cleanup_deletes_at_most_one_batch(self):
+        self._add_tokens(expired=7, live=2)
+        RevokedToken._last_cleanup_time = 0.0
+
+        with (
+            patch("airflow.models.revoked_token._CLEANUP_BATCH_SIZE", 3),
+            patch("airflow.models.revoked_token.time.monotonic", 
return_value=100_000.0),
+        ):
+            RevokedToken.is_revoked("live-0")
+
+        # 3 of the 7 expired rows gone, both unexpired rows untouched
+        assert self._remaining() == 6
+
+    @conf_vars({("api_auth", "jwt_expiration_time"): "3600"})
+    def test_full_batch_lets_the_next_check_resume_draining(self):
+        self._add_tokens(expired=7, live=0)
+        RevokedToken._last_cleanup_time = 0.0
+
+        with (
+            patch("airflow.models.revoked_token._CLEANUP_BATCH_SIZE", 3),
+            patch("airflow.models.revoked_token.time.monotonic", 
return_value=100_000.0),
+        ):
+            # Each full batch rewinds the interval, so the passes chain on a 
frozen clock.
+            RevokedToken.is_revoked("expired-0")
+            assert self._remaining() == 4
+            RevokedToken.is_revoked("expired-0")
+            assert self._remaining() == 1
+            RevokedToken.is_revoked("expired-0")
+            assert self._remaining() == 0
+
+            # The last pass did not fill the batch, so the interval applies 
again
+            self._add_tokens(expired=2, live=0)
+            RevokedToken.is_revoked("expired-0")
+            assert self._remaining() == 2

Review Comment:
   Add a test where another thread updates `_last_cleanup_time `between the 
initial interval check and acquiring the lock. This is the coverage for the gap 
pointed out in the implementation above.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to