diff --git a/alembic/versions/e1a2b3c4d5e6_add_anonymous_vote_tables.py b/alembic/versions/e1a2b3c4d5e6_add_anonymous_vote_tables.py new file mode 100644 index 0000000..f23d6e6 --- /dev/null +++ b/alembic/versions/e1a2b3c4d5e6_add_anonymous_vote_tables.py @@ -0,0 +1,76 @@ +"""Add anonymous vote tables + +Revision ID: e1a2b3c4d5e6 +Revises: d4f8c2a6e1b7 +Create Date: 2026-07-31 13:53:00.000000 + +""" +import sqlalchemy as sa +from sqlalchemy.dialects import mysql + +from alembic import op + +# revision identifiers, used by Alembic. +revision = "e1a2b3c4d5e6" +down_revision = "d4f8c2a6e1b7" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "anonymous_vote_session", + sa.Column("id", sa.Integer(), nullable=False), + sa.Column("guild_id", mysql.BIGINT(display_width=18), nullable=False), + sa.Column("channel_id", mysql.BIGINT(display_width=18), nullable=False), + sa.Column("message_id", mysql.BIGINT(display_width=18), nullable=True), + sa.Column("topic", mysql.TEXT(), nullable=True), + sa.Column("created_by_id", mysql.BIGINT(display_width=18), nullable=False), + sa.Column("closes_at", mysql.TIMESTAMP(), nullable=False), + sa.Column("closed", sa.Boolean(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_table( + "anonymous_vote_candidate", + sa.Column("id", sa.Integer(), nullable=False), + sa.Column("session_id", sa.Integer(), nullable=False), + sa.Column("user_id", mysql.BIGINT(display_width=18), nullable=False), + sa.Column("display_name", mysql.TEXT(), nullable=False), + sa.ForeignKeyConstraint( + ["session_id"], + ["anonymous_vote_session.id"], + ondelete="CASCADE", + ), + sa.PrimaryKeyConstraint("id"), + ) + op.create_table( + "anonymous_vote_ballot", + sa.Column("id", sa.Integer(), nullable=False), + sa.Column("session_id", sa.Integer(), nullable=False), + sa.Column("candidate_id", sa.Integer(), nullable=False), + sa.Column("voter_id", mysql.BIGINT(display_width=18), nullable=False), + sa.Column("choice", sa.String(length=16), nullable=False), + sa.ForeignKeyConstraint( + ["candidate_id"], + ["anonymous_vote_candidate.id"], + ondelete="CASCADE", + ), + sa.ForeignKeyConstraint( + ["session_id"], + ["anonymous_vote_session.id"], + ondelete="CASCADE", + ), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "session_id", + "candidate_id", + "voter_id", + name="uq_anonymous_vote_ballot_session_candidate_voter", + ), + ) + + +def downgrade() -> None: + op.drop_table("anonymous_vote_ballot") + op.drop_table("anonymous_vote_candidate") + op.drop_table("anonymous_vote_session") diff --git a/src/bot.py b/src/bot.py index 6f004f5..9d378c8 100644 --- a/src/bot.py +++ b/src/bot.py @@ -95,6 +95,7 @@ async def on_ready(self) -> None: async def _register_persistent_views(self) -> None: """Re-register persistent UI views so buttons survive bot restarts.""" + from src.views.anonymous_vote import register_anonymous_vote_views from src.views.bandecisionview import register_ban_views try: @@ -102,6 +103,11 @@ async def _register_persistent_views(self) -> None: except Exception: logger.exception("Failed to register persistent ban decision views") + try: + await register_anonymous_vote_views(self) + except Exception: + logger.exception("Failed to register persistent anonymous vote views") + async def on_application_command(self, ctx: ApplicationContext) -> None: """A global handler cog.""" logger.debug(f"Command '{ctx.command}' received.") diff --git a/src/cmds/core/admin.py b/src/cmds/core/admin.py index 791482d..f378209 100644 --- a/src/cmds/core/admin.py +++ b/src/cmds/core/admin.py @@ -1,19 +1,48 @@ """Admin command group for bot administration commands.""" import logging +import re +from datetime import datetime import discord from discord import ApplicationContext, Interaction, Option, WebhookMessage from discord.ext import commands from discord.ext.commands import has_any_role +from sqlalchemy import select +from sqlalchemy.orm import selectinload from src.bot import Bot from src.core import settings +from src.database.models import AnonymousVoteCandidate, AnonymousVoteSession from src.database.models.dynamic_role import RoleCategory +from src.database.session import AsyncSessionLocal +from src.helpers.duration import validate_duration +from src.views.anonymous_vote import ( + AnonymousVoteView, + build_poll_embed, + schedule_vote_close, +) logger = logging.getLogger(__name__) CATEGORY_CHOICES = [c.value for c in RoleCategory] +_MEMBER_TOKEN_RE = re.compile(r"<@!?(\d+)>|^(\d+)$") + + +def _parse_member_ids(raw: str) -> list[int]: + """Parse space/comma-separated mentions or snowflake IDs into unique IDs.""" + ids: list[int] = [] + for part in re.split(r"[\s,]+", raw.strip()): + if not part: + continue + match = _MEMBER_TOKEN_RE.fullmatch(part) + if not match: + raise ValueError( + f"Could not parse `{part}`. Use mentions or numeric user IDs." + ) + ids.append(int(match.group(1) or match.group(2))) + # Preserve order, drop duplicates + return list(dict.fromkeys(ids)) class AdminCog(commands.Cog): @@ -149,6 +178,109 @@ async def reload(self, ctx: ApplicationContext) -> Interaction | WebhookMessage: await self.bot.role_manager.reload() return await ctx.respond("Dynamic roles reloaded from database.", ephemeral=True) + @admin.command( + name="vote", + description="Start an anonymous timed vote on multiple members.", + ) + @has_any_role(*settings.role_groups.get("VOTE_STARTERS")) + async def vote( + self, + ctx: ApplicationContext, + members: Option( + str, + "Nominees as mentions or user IDs (space/comma separated, max 25)", + ), + duration: Option(str, "How long the vote stays open (e.g. 12h, 1d, 30m)"), + topic: Option(str, "Optional topic shown on the poll", required=False), + ) -> Interaction | WebhookMessage: + """Start an anonymous vote; tallies reveal automatically when duration ends.""" + closes_at_ts, error = validate_duration(duration) + if error: + return await ctx.respond(error, ephemeral=True) + + try: + member_ids = _parse_member_ids(members) + except ValueError as exc: + return await ctx.respond(str(exc), ephemeral=True) + + if not member_ids: + return await ctx.respond("Provide at least one nominee.", ephemeral=True) + if len(member_ids) > 25: + return await ctx.respond( + "Discord select menus support at most 25 nominees.", + ephemeral=True, + ) + + resolved: list[tuple[int, str]] = [] + missing: list[str] = [] + for user_id in member_ids: + member = ctx.guild.get_member(user_id) + if member is None: + try: + member = await ctx.guild.fetch_member(user_id) + except discord.HTTPException: + missing.append(str(user_id)) + continue + resolved.append((member.id, member.display_name)) + + if missing: + return await ctx.respond( + "Could not find member(s) in this server: " + ", ".join(f"`{m}`" for m in missing), + ephemeral=True, + ) + + closes_at = datetime.fromtimestamp(closes_at_ts) + await ctx.defer(ephemeral=True) + + async with AsyncSessionLocal() as session: + vote_session = AnonymousVoteSession( + guild_id=ctx.guild.id, + channel_id=ctx.channel.id, + message_id=None, + topic=topic, + created_by_id=ctx.author.id, + closes_at=closes_at, + closed=False, + ) + session.add(vote_session) + await session.flush() + + for user_id, display_name in resolved: + session.add( + AnonymousVoteCandidate( + session_id=vote_session.id, + user_id=user_id, + display_name=display_name, + ) + ) + await session.commit() + + loaded = await session.scalar( + select(AnonymousVoteSession) + .where(AnonymousVoteSession.id == vote_session.id) + .options(selectinload(AnonymousVoteSession.candidates)) + ) + session_id = loaded.id + candidates = list(loaded.candidates) + poll_embed = build_poll_embed(loaded, candidates) + + view = AnonymousVoteView(session_id, self.bot, candidates) + self.bot.add_view(view) + message = await ctx.channel.send(embed=poll_embed, view=view) + + async with AsyncSessionLocal() as session: + vote_session = await session.get(AnonymousVoteSession, session_id) + if vote_session: + vote_session.message_id = message.id + await session.commit() + + schedule_vote_close(self.bot, session_id, closes_at) + return await ctx.followup.send( + f"Anonymous vote #{session_id} started in {ctx.channel.mention}. " + f"Closes {discord.utils.format_dt(closes_at, style='R')}.", + ephemeral=True, + ) + def setup(bot: Bot) -> None: """Load the AdminCog.""" diff --git a/src/core/config.py b/src/core/config.py index b468ebc..1685f1b 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -259,6 +259,19 @@ def role_groups(self) -> dict[str, list[int]]: ], "ALL_HTB_STAFF": [self.roles.HTB_STAFF], "ALL_HTB_SUPPORT": [self.roles.HTB_SUPPORT], + "VOTE_STARTERS": [ + self.roles.ADMINISTRATOR, + self.roles.COMMUNITY_MANAGER, + self.roles.COMMUNITY_TEAM, + ], + "VOTE_CASTERS": [ + self.roles.ADMINISTRATOR, + self.roles.COMMUNITY_MANAGER, + self.roles.COMMUNITY_TEAM, + self.roles.SR_MODERATOR, + self.roles.MODERATOR, + self.roles.JR_MODERATOR, + ], } diff --git a/src/database/models/__init__.py b/src/database/models/__init__.py index 15ab602..0515ee2 100644 --- a/src/database/models/__init__.py +++ b/src/database/models/__init__.py @@ -1,6 +1,7 @@ # flake8: noqa from src.database.base_class import Base # noqa +from .anonymous_vote import AnonymousVoteBallot, AnonymousVoteCandidate, AnonymousVoteSession from .ban import Ban from .ctf import Ctf from .dynamic_role import DynamicRole, RoleCategory diff --git a/src/database/models/anonymous_vote.py b/src/database/models/anonymous_vote.py new file mode 100644 index 0000000..d11bf9c --- /dev/null +++ b/src/database/models/anonymous_vote.py @@ -0,0 +1,73 @@ +# flake8: noqa: D101 +from datetime import datetime + +from sqlalchemy import Boolean, ForeignKey, Integer, String, UniqueConstraint +from sqlalchemy.dialects.mysql import BIGINT, TEXT, TIMESTAMP +from sqlalchemy.orm import Mapped, mapped_column, relationship + +from . import Base + + +class AnonymousVoteSession(Base): + """Timed anonymous vote session over one or more nominees.""" + + id: Mapped[int] = mapped_column(Integer, primary_key=True) + guild_id: Mapped[int] = mapped_column(BIGINT(18), nullable=False) + channel_id: Mapped[int] = mapped_column(BIGINT(18), nullable=False) + message_id: Mapped[int | None] = mapped_column(BIGINT(18), nullable=True) + topic: Mapped[str | None] = mapped_column(TEXT, nullable=True) + created_by_id: Mapped[int] = mapped_column(BIGINT(18), nullable=False) + closes_at: Mapped[datetime] = mapped_column(TIMESTAMP, nullable=False) + closed: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + + candidates: Mapped[list["AnonymousVoteCandidate"]] = relationship( + back_populates="session", + cascade="all, delete-orphan", + ) + ballots: Mapped[list["AnonymousVoteBallot"]] = relationship( + back_populates="session", + cascade="all, delete-orphan", + ) + + +class AnonymousVoteCandidate(Base): + """A nominee in an anonymous vote session.""" + + id: Mapped[int] = mapped_column(Integer, primary_key=True) + session_id: Mapped[int] = mapped_column( + Integer, ForeignKey("anonymous_vote_session.id", ondelete="CASCADE"), nullable=False + ) + user_id: Mapped[int] = mapped_column(BIGINT(18), nullable=False) + display_name: Mapped[str] = mapped_column(TEXT, nullable=False) + + session: Mapped["AnonymousVoteSession"] = relationship(back_populates="candidates") + ballots: Mapped[list["AnonymousVoteBallot"]] = relationship( + back_populates="candidate", + cascade="all, delete-orphan", + ) + + +class AnonymousVoteBallot(Base): + """A single voter's choice for one nominee. voter_id is never shown in Discord.""" + + __table_args__ = ( + UniqueConstraint( + "session_id", + "candidate_id", + "voter_id", + name="uq_anonymous_vote_ballot_session_candidate_voter", + ), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True) + session_id: Mapped[int] = mapped_column( + Integer, ForeignKey("anonymous_vote_session.id", ondelete="CASCADE"), nullable=False + ) + candidate_id: Mapped[int] = mapped_column( + Integer, ForeignKey("anonymous_vote_candidate.id", ondelete="CASCADE"), nullable=False + ) + voter_id: Mapped[int] = mapped_column(BIGINT(18), nullable=False) + choice: Mapped[str] = mapped_column(String(16), nullable=False) + + session: Mapped["AnonymousVoteSession"] = relationship(back_populates="ballots") + candidate: Mapped["AnonymousVoteCandidate"] = relationship(back_populates="ballots") diff --git a/src/views/anonymous_vote.py b/src/views/anonymous_vote.py new file mode 100644 index 0000000..0d1e585 --- /dev/null +++ b/src/views/anonymous_vote.py @@ -0,0 +1,380 @@ +"""Persistent UI and helpers for anonymous multi-user votes.""" + +from __future__ import annotations + +import logging +from collections import defaultdict +from datetime import datetime + +import discord +from discord import Interaction, SelectOption +from discord.ui import Button, Select, View +from sqlalchemy import select +from sqlalchemy.orm import selectinload + +from src.bot import Bot +from src.core import settings +from src.database.models import AnonymousVoteBallot, AnonymousVoteCandidate, AnonymousVoteSession +from src.database.session import AsyncSessionLocal +from src.helpers.schedule import schedule + +logger = logging.getLogger(__name__) + +CHOICE_APPROVE = "approve" +CHOICE_REJECT = "reject" +# Neutral marker shown per ballot cast (does not reveal approve vs reject). +VOTE_ACTIVITY_BOX = "⬜" +# Keep embed descriptions safely under Discord's 4096-char limit. +MAX_ACTIVITY_BOXES_PER_NOMINEE = 40 + +_pending_selection: dict[tuple[int, int], int] = {} + + +def _member_can_vote(member: discord.Member) -> bool: + voter_roles = set(settings.role_groups.get("VOTE_CASTERS", [])) + return bool(voter_roles.intersection({role.id for role in member.roles})) + + +def _ballot_counts_by_candidate(ballots: list[AnonymousVoteBallot]) -> dict[int, int]: + counts: dict[int, int] = defaultdict(int) + for ballot in ballots: + counts[ballot.candidate_id] += 1 + return counts + + +def _format_nominee_line(candidate: AnonymousVoteCandidate, vote_count: int) -> str: + """Format a nominee line with one activity box per cast ballot.""" + line = f"• **{candidate.display_name}** (`{candidate.user_id}`)" + if vote_count <= 0: + return line + shown = min(vote_count, MAX_ACTIVITY_BOXES_PER_NOMINEE) + boxes = VOTE_ACTIVITY_BOX * shown + if vote_count > MAX_ACTIVITY_BOXES_PER_NOMINEE: + boxes += "…" + return f"{line} {boxes}" + + +def build_poll_embed( + session: AnonymousVoteSession, + candidates: list[AnonymousVoteCandidate], + ballots: list[AnonymousVoteBallot] | None = None, +) -> discord.Embed: + """Build the public poll embed (no approve/reject tallies).""" + title = session.topic or "Anonymous vote" + counts = _ballot_counts_by_candidate(ballots or []) + lines = [_format_nominee_line(c, counts.get(c.id, 0)) for c in candidates] + embed = discord.Embed( + title=title, + description="\n".join(lines) if lines else "No nominees.", + color=0x9ACC14, + ) + embed.add_field( + name="How to vote", + value=( + "1. Select a nominee from the menu\n" + "2. Press **Approve** or **Reject**\n" + f"Each {VOTE_ACTIVITY_BOX} next to a name means someone voted on them " + "(not whether it was approve or reject). " + "Votes stay anonymous; exact tallies appear when the poll closes." + ), + inline=False, + ) + embed.add_field( + name="Closes", + value=discord.utils.format_dt(session.closes_at, style="F") + + f" ({discord.utils.format_dt(session.closes_at, style='R')})", + inline=False, + ) + embed.set_footer(text=f"Session #{session.id}") + return embed + + +def build_results_embed( + session: AnonymousVoteSession, + candidates: list[AnonymousVoteCandidate], + ballots: list[AnonymousVoteBallot], +) -> discord.Embed: + """Build results embed with totals only (no voter identities).""" + counts: dict[int, dict[str, int]] = defaultdict(lambda: {CHOICE_APPROVE: 0, CHOICE_REJECT: 0}) + for ballot in ballots: + if ballot.choice in (CHOICE_APPROVE, CHOICE_REJECT): + counts[ballot.candidate_id][ballot.choice] += 1 + + title = session.topic or "Anonymous vote" + lines = [] + for candidate in candidates: + tally = counts[candidate.id] + lines.append( + f"• **{candidate.display_name}** — ✓ {tally[CHOICE_APPROVE]} / ✗ {tally[CHOICE_REJECT]}" + ) + + embed = discord.Embed( + title=f"Results: {title}", + description="\n".join(lines) if lines else "No nominees.", + color=0x5865F2, + ) + embed.set_footer(text=f"Session #{session.id} • Closed • Voter identities are not shown") + return embed + + +class AnonymousVoteView(View): + """Persistent view: select a nominee, then approve/reject anonymously.""" + + def __init__( + self, + session_id: int, + bot: Bot, + candidates: list[AnonymousVoteCandidate] | None = None, + ): + super().__init__(timeout=None) + self.session_id = session_id + self.bot = bot + + options = [ + SelectOption( + label=(c.display_name[:100] or str(c.user_id)), + value=str(c.id), + description=f"ID {c.user_id}"[:100], + ) + for c in (candidates or []) + ] + if not options: + options = [SelectOption(label="No nominees", value="0", default=True)] + + select = Select( + placeholder="Select a nominee to vote on", + options=options[:25], + custom_id=f"anon_vote_select:{session_id}", + min_values=1, + max_values=1, + disabled=not candidates, + ) + + async def on_select(interaction: Interaction) -> None: + await self._handle_select(interaction, select.values) + + select.callback = on_select + self.add_item(select) + + approve_btn = Button( + label="Approve", + style=discord.ButtonStyle.success, + emoji="✅", + custom_id=f"anon_vote_approve:{session_id}", + ) + approve_btn.callback = self._on_approve + self.add_item(approve_btn) + + reject_btn = Button( + label="Reject", + style=discord.ButtonStyle.danger, + emoji="❌", + custom_id=f"anon_vote_reject:{session_id}", + ) + reject_btn.callback = self._on_reject + self.add_item(reject_btn) + + async def _handle_select(self, interaction: Interaction, values: list[str]) -> None: + if not isinstance(interaction.user, discord.Member) or not _member_can_vote(interaction.user): + await interaction.response.send_message( + "You are not authorized to vote in this poll.", ephemeral=True + ) + return + + if not values: + await interaction.response.send_message("No nominee selected.", ephemeral=True) + return + + candidate_id = int(values[0]) + if candidate_id <= 0: + await interaction.response.send_message("Invalid nominee.", ephemeral=True) + return + + async with AsyncSessionLocal() as session: + vote_session = await session.get(AnonymousVoteSession, self.session_id) + if not vote_session or vote_session.closed: + await interaction.response.send_message("This poll is closed.", ephemeral=True) + return + candidate = await session.get(AnonymousVoteCandidate, candidate_id) + if not candidate or candidate.session_id != self.session_id: + await interaction.response.send_message("Unknown nominee.", ephemeral=True) + return + display_name = candidate.display_name + + _pending_selection[(self.session_id, interaction.user.id)] = candidate_id + await interaction.response.send_message( + f"Selected **{display_name}**. Now press **Approve** or **Reject**.", + ephemeral=True, + ) + + async def _on_approve(self, interaction: Interaction) -> None: + await self._cast_vote(interaction, CHOICE_APPROVE) + + async def _on_reject(self, interaction: Interaction) -> None: + await self._cast_vote(interaction, CHOICE_REJECT) + + async def _cast_vote(self, interaction: Interaction, choice: str) -> None: + if not isinstance(interaction.user, discord.Member) or not _member_can_vote(interaction.user): + await interaction.response.send_message( + "You are not authorized to vote in this poll.", ephemeral=True + ) + return + + candidate_id = _pending_selection.get((self.session_id, interaction.user.id)) + if not candidate_id: + await interaction.response.send_message( + "Select a nominee from the menu first, then press Approve or Reject.", + ephemeral=True, + ) + return + + poll_embed: discord.Embed | None = None + async with AsyncSessionLocal() as session: + vote_session = await session.get(AnonymousVoteSession, self.session_id) + if not vote_session or vote_session.closed: + await interaction.response.send_message("This poll is closed.", ephemeral=True) + return + + candidate = await session.get(AnonymousVoteCandidate, candidate_id) + if not candidate or candidate.session_id != self.session_id: + await interaction.response.send_message("Unknown nominee.", ephemeral=True) + return + + stmt = select(AnonymousVoteBallot).where( + AnonymousVoteBallot.session_id == self.session_id, + AnonymousVoteBallot.candidate_id == candidate_id, + AnonymousVoteBallot.voter_id == interaction.user.id, + ) + ballot = await session.scalar(stmt) + if ballot: + ballot.choice = choice + action = "updated" + else: + session.add( + AnonymousVoteBallot( + session_id=self.session_id, + candidate_id=candidate_id, + voter_id=interaction.user.id, + choice=choice, + ) + ) + action = "recorded" + + display_name = candidate.display_name + await session.commit() + + loaded = await session.scalar( + select(AnonymousVoteSession) + .where(AnonymousVoteSession.id == self.session_id) + .options( + selectinload(AnonymousVoteSession.candidates), + selectinload(AnonymousVoteSession.ballots), + ) + ) + poll_embed = build_poll_embed( + loaded, + list(loaded.candidates), + list(loaded.ballots), + ) if loaded else None + + label = "Approve" if choice == CHOICE_APPROVE else "Reject" + await interaction.response.send_message( + f"Vote {action}: **{label}** for **{display_name}**. " + "Your choice is anonymous; only a neutral activity box is shown publicly.", + ephemeral=True, + ) + + if poll_embed is not None and interaction.message is not None: + try: + await interaction.message.edit(embed=poll_embed) + except discord.HTTPException: + logger.exception( + "Failed to refresh poll embed for session %s after vote.", + self.session_id, + ) + + +async def close_anonymous_vote(bot: Bot, session_id: int) -> None: + """Close a vote session, post totals, and disable controls.""" + async with AsyncSessionLocal() as session: + stmt = ( + select(AnonymousVoteSession) + .where(AnonymousVoteSession.id == session_id) + .options( + selectinload(AnonymousVoteSession.candidates), + selectinload(AnonymousVoteSession.ballots), + ) + ) + vote_session = await session.scalar(stmt) + if not vote_session: + logger.warning("Anonymous vote session %s not found for close.", session_id) + return + if vote_session.closed: + logger.debug("Anonymous vote session %s already closed.", session_id) + return + + vote_session.closed = True + candidates = list(vote_session.candidates) + ballots = list(vote_session.ballots) + channel_id = vote_session.channel_id + message_id = vote_session.message_id + results_embed = build_results_embed(vote_session, candidates, ballots) + await session.commit() + + channel = bot.get_channel(channel_id) + if channel is None: + try: + channel = await bot.fetch_channel(channel_id) + except discord.HTTPException: + logger.exception("Failed to fetch channel %s for vote session %s", channel_id, session_id) + return + + view = AnonymousVoteView(session_id, bot, candidates) + for item in view.children: + item.disabled = True + + if message_id: + try: + message = await channel.fetch_message(message_id) + await message.edit( + content="This anonymous vote is closed. Results below.", + embed=results_embed, + view=view, + ) + return + except discord.HTTPException: + logger.exception( + "Failed to edit poll message %s for session %s; posting results separately.", + message_id, + session_id, + ) + + await channel.send(embed=results_embed) + + +def schedule_vote_close(bot: Bot, session_id: int, closes_at: datetime) -> None: + """Schedule auto-close for a vote session on the bot event loop.""" + bot.loop.create_task(schedule(close_anonymous_vote(bot, session_id), closes_at)) + + +async def register_anonymous_vote_views(bot: Bot) -> None: + """Re-register open vote views and reschedule their closes after restart.""" + async with AsyncSessionLocal() as session: + stmt = ( + select(AnonymousVoteSession) + .where(AnonymousVoteSession.closed.is_(False)) + .options(selectinload(AnonymousVoteSession.candidates)) + ) + result = await session.scalars(stmt) + open_sessions = list(result.all()) + + now = datetime.now() + for vote_session in open_sessions: + bot.add_view(AnonymousVoteView(vote_session.id, bot, vote_session.candidates)) + if vote_session.closes_at <= now: + bot.loop.create_task(close_anonymous_vote(bot, vote_session.id)) + else: + schedule_vote_close(bot, vote_session.id, vote_session.closes_at) + + if open_sessions: + logger.info("Registered %d open anonymous vote session(s).", len(open_sessions)) diff --git a/tests/src/core/test_config.py b/tests/src/core/test_config.py index 0260bac..6730d84 100644 --- a/tests/src/core/test_config.py +++ b/tests/src/core/test_config.py @@ -81,6 +81,10 @@ def test_core_role_groups_present(self): self.assertIn("ALL_HTB_STAFF", settings.role_groups) self.assertIn("ALL_SR_MODS", settings.role_groups) self.assertIn("ALL_HTB_SUPPORT", settings.role_groups) + self.assertIn("VOTE_STARTERS", settings.role_groups) + self.assertIn("VOTE_CASTERS", settings.role_groups) + self.assertEqual(len(settings.role_groups["VOTE_STARTERS"]), 3) + self.assertEqual(len(settings.role_groups["VOTE_CASTERS"]), 6) def test_dynamic_role_groups_removed(self): """Test that dynamic role groups are no longer in settings."""