diff --git a/modules/region_warning.py b/modules/region_warning.py index 7bac72b..47d24cc 100644 --- a/modules/region_warning.py +++ b/modules/region_warning.py @@ -519,14 +519,14 @@ class RegionWarningMonitor: ) return - # The run has been spent whether or not the send itself succeeds; not - # resetting it would retry on the sender's very next message. - state.unscoped_seen = 0 - text = render_message(settings.message, sender_id, channel) if not text: return + # The run has been spent whether or not the send itself succeeds; not + # resetting it would retry on the sender's very next message. + state.unscoped_seen = 0 + if settings.dry_run: self._record_event(sender_id, sender_pubkey, channel, ACTION_DRY_RUN, text) self._mark_warning_sent(now) @@ -535,16 +535,26 @@ class RegionWarningMonitor: ) return + # Reserve the slot *before* awaiting the send. Every gate above and this + # reservation run without an await between them, so on the single event + # loop they are atomic: a second channel message arriving mid-send reads + # a cooldown and a cap row that already account for this warning. Doing + # it after the send instead let two concurrent messages both pass a cap + # of one and both transmit. + previous_mark = (self._last_warning_monotonic, self._last_warning_wall) + event_id = self._record_event( + sender_id, sender_pubkey, channel, ACTION_SENT, text) + self._mark_warning_sent(now) + sent, detail = await self._send_warning(sender_id, channel, text) - self._record_event( - sender_id, sender_pubkey, channel, - ACTION_SENT if sent else ACTION_FAILED, - detail, - ) if sent: - # Only a successful send starts the mesh cooldown; a failed one - # spent no airtime and should not silence a sender who would. - self._mark_warning_sent(now) + return + # Correct the optimistic reservation. The event row stays — the attempt + # still counts against the daily cap, because a send that reported + # failure may have put something on the air before it did — but a + # failure should not hold the mesh cooldown against the next sender. + self._update_event(event_id, ACTION_FAILED, detail) + self._last_warning_monotonic, self._last_warning_wall = previous_mark async def _send_warning( self, sender_id: str, channel: Optional[str], text: str @@ -714,19 +724,19 @@ class RegionWarningMonitor: channel: Optional[str], action: str, detail: str, - ) -> None: + ) -> Optional[int]: + """Write one decision row; returns its id so the outcome can be corrected.""" db_manager = getattr(self.bot, "db_manager", None) if not db_manager: - return + return None try: with db_manager.connection() as conn: - conn.execute( + cursor = conn.execute( "INSERT INTO region_warning_events " "(created_at, sender_id, sender_pubkey, channel, delivery, action, detail) " "VALUES (?, ?, ?, ?, ?, ?, ?)", ( - local_now(getattr(self.bot, "config", None), self.logger) - .isoformat(sep=" ", timespec="seconds"), + self._now().isoformat(sep=" ", timespec="seconds"), sender_id, (sender_pubkey or "")[:64] or None, channel, @@ -736,8 +746,25 @@ class RegionWarningMonitor: ), ) conn.commit() + return int(cursor.lastrowid) if cursor.lastrowid else None except Exception: self.logger.exception("Failed to record region warning event") + return None + + def _update_event(self, event_id: Optional[int], action: str, detail: str) -> None: + """Replace a reserved row's outcome once the send has resolved.""" + db_manager = getattr(self.bot, "db_manager", None) + if not db_manager or event_id is None: + return + try: + with db_manager.connection() as conn: + conn.execute( + "UPDATE region_warning_events SET action = ?, detail = ? WHERE id = ?", + (action, detail[:400] if detail else None, event_id), + ) + conn.commit() + except Exception: + self.logger.exception("Failed to update region warning event") def _parse_timestamp(raw: Any) -> Optional[datetime]: diff --git a/tests/unit/test_region_warning.py b/tests/unit/test_region_warning.py index 3c2fbb3..2317528 100644 --- a/tests/unit/test_region_warning.py +++ b/tests/unit/test_region_warning.py @@ -358,6 +358,40 @@ class TestWarningGates: await self._flood(monitor, 1, sender=name) assert monitor.bot.command_manager.send_dm.call_count == 2 + async def test_concurrent_messages_cannot_both_pass_a_cap_of_one(self): + """The slot is reserved before the send awaits, so a second message sees it.""" + import asyncio + + monitor = _monitor( + enabled="true", dry_run="false", min_unscoped_messages=1, + mesh_cooldown_minutes=0, per_sender_cooldown_hours=0, max_warnings_per_day=1) + + async def slow_send(*_args, **_kwargs): + await asyncio.sleep(0) + await asyncio.sleep(0) + return True + + monitor.bot.command_manager.send_dm = AsyncMock(side_effect=slow_send) + await asyncio.gather( + monitor.observe( + verdict=VERDICT_GLOBAL, sender_id="Ann", sender_pubkey="ab", channel="#gen"), + monitor.observe( + verdict=VERDICT_GLOBAL, sender_id="Bob", sender_pubkey="cd", channel="#gen"), + ) + assert monitor.bot.command_manager.send_dm.call_count == 1 + + async def test_failed_send_row_is_corrected_not_duplicated(self): + monitor = _monitor( + enabled="true", dry_run="false", min_unscoped_messages=1, + mesh_cooldown_minutes=0, per_sender_cooldown_hours=0) + monitor.bot.command_manager.send_dm = AsyncMock(return_value=False) + await self._flood(monitor, 1) + rows = monitor.bot.db_manager.execute_query( + "SELECT action, detail FROM region_warning_events") + assert len(rows) == 1 + assert rows[0]["action"] == ACTION_FAILED + assert "failed" in rows[0]["detail"] + async def test_sender_table_stays_bounded(self): monitor = _monitor(min_unscoped_messages=99) monitor.MAX_TRACKED_SENDERS = 10