mirror of
https://github.com/rennf93/roboco.git
synced 2026-08-03 07:23:24 +02:00
[F135] re-check activity before sweeper closes a session (TOCTOU)
sweep_timed_out_sessions read last_activity_at once at the candidate SELECT, then closed. A message landing in that window refreshed last_activity_at in the DB, but the sweeper closed on its stale in-memory value — closing a just-used session. Re-read last_activity_at fresh right before the close and skip if the session is no longer timed out.
This commit is contained in:
@@ -536,6 +536,37 @@ class MessagingService(BaseService):
|
|||||||
)
|
)
|
||||||
return result.scalar_one_or_none()
|
return result.scalar_one_or_none()
|
||||||
|
|
||||||
|
async def _session_still_timed_out(
|
||||||
|
self,
|
||||||
|
session_id: UUID,
|
||||||
|
*,
|
||||||
|
timeout_seconds: int,
|
||||||
|
max_time_window: timedelta | None,
|
||||||
|
) -> bool:
|
||||||
|
"""Re-check the timeout against a FRESH ``last_activity_at`` from the DB.
|
||||||
|
|
||||||
|
Closes the sweeper TOCTOU: a message may have refreshed
|
||||||
|
``last_activity_at`` between the candidate SELECT and the close, so the
|
||||||
|
stale in-memory value would close a just-used session. Returns False if
|
||||||
|
the session is no longer timed out / window-exceeded per the fresh row
|
||||||
|
(or was closed concurrently).
|
||||||
|
"""
|
||||||
|
now = datetime.now(UTC)
|
||||||
|
result = await self.session.execute(
|
||||||
|
select(
|
||||||
|
SessionTable.last_activity_at,
|
||||||
|
SessionTable.started_at,
|
||||||
|
SessionTable.status,
|
||||||
|
).where(SessionTable.id == session_id)
|
||||||
|
)
|
||||||
|
row = result.one_or_none()
|
||||||
|
if row is None or row.status != SessionStatus.ACTIVE:
|
||||||
|
return False
|
||||||
|
last = row.last_activity_at or row.started_at
|
||||||
|
if (now - last).total_seconds() >= timeout_seconds:
|
||||||
|
return True
|
||||||
|
return max_time_window is not None and (now - row.started_at) >= max_time_window
|
||||||
|
|
||||||
async def sweep_timed_out_sessions(self) -> int:
|
async def sweep_timed_out_sessions(self) -> int:
|
||||||
"""Close sessions whose inactivity exceeds `timeout_seconds`.
|
"""Close sessions whose inactivity exceeds `timeout_seconds`.
|
||||||
|
|
||||||
@@ -567,6 +598,16 @@ class MessagingService(BaseService):
|
|||||||
if not (timeout_exceeded or window_exceeded):
|
if not (timeout_exceeded or window_exceeded):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# TOCTOU re-check: a message may have refreshed last_activity_at
|
||||||
|
# between the candidate SELECT above and here. Re-read fresh and
|
||||||
|
# skip if the session is no longer timed out (was just used).
|
||||||
|
if not await self._session_still_timed_out(
|
||||||
|
cast("UUID", session.id),
|
||||||
|
timeout_seconds=session.timeout_seconds,
|
||||||
|
max_time_window=session.max_time_window,
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
reason = "Inactivity timeout" if timeout_exceeded else "Max time window"
|
reason = "Inactivity timeout" if timeout_exceeded else "Max time window"
|
||||||
await self.close_session(cast("UUID", session.id), reason)
|
await self.close_session(cast("UUID", session.id), reason)
|
||||||
closed += 1
|
closed += 1
|
||||||
|
|||||||
@@ -1020,6 +1020,38 @@ async def test_sweep_timed_out_sessions_closes_idle_session(
|
|||||||
assert closed >= 1
|
assert closed >= 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sweep_skips_session_refreshed_after_candidate_select(
|
||||||
|
msg_setup: dict, db_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""TOCTOU: a session stale at the candidate SELECT but refreshed (a message
|
||||||
|
landed) before the close must NOT be closed. The sweeper re-reads
|
||||||
|
last_activity_at fresh and skips a just-used session."""
|
||||||
|
svc = msg_setup["svc"]
|
||||||
|
ch = await svc.create_channel(_channel_req(uuid4().hex[:6]))
|
||||||
|
grp = await svc.create_group(GroupCreateRequest(name="g1", channel_id=ch.id))
|
||||||
|
sess = await svc.create_session(
|
||||||
|
SessionCreateRequest(group_id=grp.id, timeout_seconds=1)
|
||||||
|
)
|
||||||
|
sid = sess.id
|
||||||
|
# Stale at SELECT time -> candidate.
|
||||||
|
sess.last_activity_at = datetime.now(UTC) - timedelta(seconds=120)
|
||||||
|
await db_session.flush()
|
||||||
|
# Fresh re-read sees stale -> still timed out.
|
||||||
|
assert (
|
||||||
|
await svc._session_still_timed_out(sid, timeout_seconds=1, max_time_window=None)
|
||||||
|
is True
|
||||||
|
)
|
||||||
|
# A message lands, refreshing activity.
|
||||||
|
sess.last_activity_at = datetime.now(UTC)
|
||||||
|
await db_session.flush()
|
||||||
|
# Fresh re-read sees recent -> no longer timed out -> sweeper must skip.
|
||||||
|
assert (
|
||||||
|
await svc._session_still_timed_out(sid, timeout_seconds=1, max_time_window=None)
|
||||||
|
is False
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_edit_message_or_raise_succeeds(msg_setup: dict) -> None:
|
async def test_edit_message_or_raise_succeeds(msg_setup: dict) -> None:
|
||||||
svc = msg_setup["svc"]
|
svc = msg_setup["svc"]
|
||||||
|
|||||||
Reference in New Issue
Block a user