Refactored polls module to use browser sessions for creation and voting routes.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (Python 3.12) (push) Successful in 2m47s
CI / Tests (Python 3.13) (push) Successful in 2m46s
CI / Tests (Python 3.14) (push) Successful in 2m44s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 6s

This commit is contained in:
2026-05-05 13:46:41 -04:00
parent 9ac4a17ac8
commit 4189b76f33
8 changed files with 136 additions and 295 deletions
+1 -1
Submodule docs updated: d8ce4b4570...fb9fb34e27
+1 -9
View File
@@ -114,8 +114,7 @@ async def on_user_joined(ctx: EventContext[UserJoinedEvent]) -> None:
"""Notify a user joining chat that a poll is active. """Notify a user joining chat that a poll is active.
Sends a private message with the poll question, options, and voting Sends a private message with the poll question, options, and voting
instructions. Also refreshes the user's token privileges in case their instructions.
authentication or moderator status has changed.
:param ctx: The event context. :param ctx: The event context.
""" """
@@ -124,13 +123,6 @@ async def on_user_joined(ctx: EventContext[UserJoinedEvent]) -> None:
if poll is None: if poll is None:
return return
event = ctx.event
manager.refresh_user_status(
event.user.id,
is_moderator=event.user.is_moderator,
is_authenticated=event.user.is_authenticated,
)
escaped_question = escape(poll.question) escaped_question = escape(poll.question)
option_lines = "".join( option_lines = "".join(
f"<strong>{i}.</strong> {escape(opt)}<br>" f"<strong>{i}.</strong> {escape(opt)}<br>"
+14 -28
View File
@@ -34,22 +34,19 @@ from .manager import PollError, get_manager
async def poll_command(ctx: CommandContext) -> None: async def poll_command(ctx: CommandContext) -> None:
"""Start the poll creation flow. """Start the poll creation flow.
Generates a moderator token and sends the poll creator a private link Sends the poll creator a private protected link to the creation form.
to the creation form.
:param ctx: The command context. :param ctx: The command context.
""" """
manager = get_manager(ctx.module) manager = get_manager(ctx.module)
try: if manager.active_poll is not None:
mod_token = await manager.begin_creation(ctx.user.id)
except PollError as e:
await ctx.owncast_client.send_system_message_to_client( await ctx.owncast_client.send_system_message_to_client(
ctx.chat_event.client_id, str(e) ctx.chat_event.client_id, "A poll is already active."
) )
return return
url = ctx.routes.url_for(f"/create/{mod_token}") url = ctx.session_url_for("/create")
await ctx.owncast_client.send_system_message_to_client( await ctx.owncast_client.send_system_message_to_client(
ctx.chat_event.client_id, ctx.chat_event.client_id,
f'<a href="{url}"><u>Click here to set up a poll</u></a>.', f'<a href="{url}"><u>Click here to set up a poll</u></a>.',
@@ -114,12 +111,7 @@ async def vote_command(ctx: CommandContext) -> None:
return return
user_id = ctx.user.id user_id = ctx.user.id
token = manager.get_voter_token( vote_url = ctx.session_url_for("/vote")
user_id,
is_moderator=ctx.user.is_moderator,
is_authenticated=ctx.user.is_authenticated,
)
vote_url = ctx.routes.url_for(f"/vote/{token.value}")
# Check authentication requirement. # Check authentication requirement.
if poll.requires_auth and not ctx.user.is_authenticated: if poll.requires_auth and not ctx.user.is_authenticated:
@@ -165,7 +157,10 @@ async def vote_command(ctx: CommandContext) -> None:
# Convert 1-based to 0-based index. # Convert 1-based to 0-based index.
selection = choice - 1 selection = choice - 1
await manager.record_vote(user_id, {selection}) await manager.record_vote(
ctx.user,
{selection},
)
await ctx.owncast_client.send_system_message_to_client( await ctx.owncast_client.send_system_message_to_client(
ctx.chat_event.client_id, ctx.chat_event.client_id,
f"Your vote for #{choice} ({escape(poll.options[selection])})" f"Your vote for #{choice} ({escape(poll.options[selection])})"
@@ -209,13 +204,6 @@ async def handle_bare_vote(ctx: EventContext[ChatEvent]) -> None:
event = ctx.event event = ctx.event
# Keep token privilege flags in sync with the user's current status.
manager.refresh_user_status(
event.user.id,
is_moderator=event.user.is_moderator,
is_authenticated=event.user.is_authenticated,
)
text = event.raw_body.strip() text = event.raw_body.strip()
try: try:
@@ -231,12 +219,7 @@ async def handle_bare_vote(ctx: EventContext[ChatEvent]) -> None:
await ctx.owncast_client.set_message_visibility([event.message_id], visible=False) await ctx.owncast_client.set_message_visibility([event.message_id], visible=False)
ctx.stop_propagation("bare number vote in active poll") ctx.stop_propagation("bare number vote in active poll")
token = manager.get_voter_token( vote_url = ctx.session_url_for("/vote")
event.user.id,
is_moderator=event.user.is_moderator,
is_authenticated=event.user.is_authenticated,
)
vote_url = ctx.routes.url_for(f"/vote/{token.value}")
# Check authentication requirement. # Check authentication requirement.
if poll.requires_auth and not event.user.is_authenticated: if poll.requires_auth and not event.user.is_authenticated:
@@ -274,7 +257,10 @@ async def handle_bare_vote(ctx: EventContext[ChatEvent]) -> None:
# Record the vote. # Record the vote.
selection = choice - 1 selection = choice - 1
await manager.record_vote(event.user.id, {selection}) await manager.record_vote(
event.user,
{selection},
)
await ctx.owncast_client.send_system_message_to_client( await ctx.owncast_client.send_system_message_to_client(
event.client_id, event.client_id,
f"Your vote for #{choice} ({escape(poll.options[selection])})" f"Your vote for #{choice} ({escape(poll.options[selection])})"
+40 -155
View File
@@ -22,14 +22,12 @@ from __future__ import annotations
import asyncio import asyncio
import contextlib import contextlib
import secrets
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from markupsafe import escape from markupsafe import escape
from .types import ( from .types import (
CREATION_TOKEN_TIMEOUT,
MAX_DURATION, MAX_DURATION,
MAX_OPTION_LENGTH, MAX_OPTION_LENGTH,
MAX_OPTIONS, MAX_OPTIONS,
@@ -41,8 +39,8 @@ from .types import (
STREAM_GRACE_PERIOD, STREAM_GRACE_PERIOD,
ActivePoll, ActivePoll,
CompletedPoll, CompletedPoll,
PollSSEClient,
SSEEvent, SSEEvent,
Token,
compute_results, compute_results,
sse_payload, sse_payload,
) )
@@ -50,7 +48,7 @@ from .types import (
if TYPE_CHECKING: if TYPE_CHECKING:
from aiohttp import web from aiohttp import web
from owlbot.api import ModuleContext from owlbot.api import ModuleContext, User
class PollError(Exception): class PollError(Exception):
@@ -70,8 +68,6 @@ class PollManager:
"""Initialize the manager with the module context.""" """Initialize the manager with the module context."""
self._ctx = ctx self._ctx = ctx
self._active_poll: ActivePoll | None = None self._active_poll: ActivePoll | None = None
self._creation_token: Token | None = None
self._creation_expiry_task: asyncio.Task[None] | None = None
self._grace_task: asyncio.Task[None] | None = None self._grace_task: asyncio.Task[None] | None = None
self._last_result: CompletedPoll | None = None self._last_result: CompletedPoll | None = None
self._result_expiry_task: asyncio.Task[None] | None = None self._result_expiry_task: asyncio.Task[None] | None = None
@@ -81,57 +77,11 @@ class PollManager:
"""The currently active poll, if any.""" """The currently active poll, if any."""
return self._active_poll return self._active_poll
@property
def creation_token_value(self) -> str | None:
"""The pending token value for poll creation (or None)."""
return self._creation_token.value if self._creation_token else None
@property @property
def last_result(self) -> CompletedPoll | None: def last_result(self) -> CompletedPoll | None:
"""The most recently completed poll result, if any.""" """The most recently completed poll result, if any."""
return self._last_result return self._last_result
async def begin_creation(self, creator_id: str) -> str:
"""Begin the poll creation flow by generating a moderator token.
The token expires after :data:`CREATION_TOKEN_TIMEOUT` seconds
if not used.
:param creator_id: The user ID of the poll creator.
:return: The generated moderator token.
:raises PollError: If a poll is already active.
"""
if self._active_poll is not None:
raise PollError("A poll is already active.")
# Cancel any previous token and its expiry task.
self.cancel_creation_token()
self._creation_token = Token(
user_id=creator_id,
is_moderator=True,
is_authenticated=True,
)
async def _expire_creation_token() -> None:
await asyncio.sleep(CREATION_TOKEN_TIMEOUT)
self._creation_token = None
self._creation_expiry_task = None
self._ctx.logger.debug(
"Creation token expired after %ds.",
CREATION_TOKEN_TIMEOUT,
)
self._creation_expiry_task = asyncio.create_task(_expire_creation_token())
return self._creation_token.value
def cancel_creation_token(self) -> None:
"""Cancel any pending creation token and its expiry task."""
if self._creation_expiry_task is not None:
self._creation_expiry_task.cancel()
self._creation_expiry_task = None
self._creation_token = None
def _cancel_result_expiry(self) -> None: def _cancel_result_expiry(self) -> None:
"""Cancel any running result expiry task.""" """Cancel any running result expiry task."""
if self._result_expiry_task is not None: if self._result_expiry_task is not None:
@@ -141,7 +91,6 @@ class PollManager:
async def create_poll( async def create_poll(
self, self,
*, *,
token: str,
question: str, question: str,
options: list[str], options: list[str],
hidden: bool, hidden: bool,
@@ -153,10 +102,9 @@ class PollManager:
) -> ActivePoll: ) -> ActivePoll:
"""Create and activate a new poll. """Create and activate a new poll.
Validates the moderator token, constructs an :class:`ActivePoll`, Constructs an :class:`ActivePoll`, starts the auto-end timer task, and
starts the auto-end timer task, and announces in chat. announces in chat.
:param token: Moderator token from :meth:`begin_creation`.
:param question: The poll question. :param question: The poll question.
:param options: List of poll option strings. :param options: List of poll option strings.
:param hidden: Whether vote results are hidden until the poll ends. :param hidden: Whether vote results are hidden until the poll ends.
@@ -166,8 +114,7 @@ class PollManager:
:param max_selections: Maximum selections allowed (multi-select). :param max_selections: Maximum selections allowed (multi-select).
:param duration: Poll duration in seconds. :param duration: Poll duration in seconds.
:return: The newly created :class:`ActivePoll`. :return: The newly created :class:`ActivePoll`.
:raises PollError: If validation fails, the token is invalid, :raises PollError: If validation fails or a poll is already active.
or a poll is already active.
""" """
self._last_result = None self._last_result = None
self._cancel_result_expiry() self._cancel_result_expiry()
@@ -175,11 +122,6 @@ class PollManager:
if self._active_poll is not None: if self._active_poll is not None:
raise PollError("A poll is already active.") raise PollError("A poll is already active.")
if self._creation_token is None or not secrets.compare_digest(
self._creation_token.value, token
):
raise PollError("Invalid or expired creation token.")
# Validate inputs. # Validate inputs.
if not question: if not question:
raise PollError("Question is required.") raise PollError("Question is required.")
@@ -249,13 +191,9 @@ class PollManager:
max_selections=max_selections, max_selections=max_selections,
duration=duration, duration=duration,
created_at=now, created_at=now,
tokens={token: self._creation_token},
) )
self._active_poll = poll self._active_poll = poll
# The token object now lives in poll.tokens as a regular token,
# so clearing the creation reference has no effect on it.
self.cancel_creation_token()
# Start auto-end timer task with reminder before expiry. # Start auto-end timer task with reminder before expiry.
async def _poll_timer() -> None: async def _poll_timer() -> None:
@@ -321,60 +259,14 @@ class PollManager:
return poll return poll
def get_voter_token( async def record_vote(
self, user_id: str, *, is_moderator: bool, is_authenticated: bool self,
) -> Token: user: User,
"""Get or create a token for the given user. selections: set[int],
:param user_id: The user's unique identifier.
:param is_moderator: Whether the user has moderator privileges.
:param is_authenticated: Whether the user has an authenticated Owncast account.
:return: The Token object.
:raises PollError: If no poll is currently active.
"""
poll = self._active_poll
if poll is None:
raise PollError("No active poll.")
# Return existing token if user already has one.
for token in poll.tokens.values():
if token.user_id == user_id:
return token
# Create new token with privilege based on mod status.
token = Token(
user_id=user_id,
is_moderator=is_moderator,
is_authenticated=is_authenticated,
)
poll.tokens[token.value] = token
return token
def refresh_user_status(
self, user_id: str, *, is_moderator: bool, is_authenticated: bool
) -> None: ) -> None:
"""Update the privilege flags on an existing token for the given user.
No-op if there is no active poll or the user has no token.
:param user_id: The user's unique identifier.
:param is_moderator: Current moderator status.
:param is_authenticated: Current authentication status.
"""
poll = self._active_poll
if poll is None:
return
for token in poll.tokens.values():
if token.user_id == user_id:
token.is_moderator = is_moderator
token.is_authenticated = is_authenticated
return
async def record_vote(self, user_id: str, selections: set[int]) -> None:
"""Record a user's vote and broadcast an SSE update. """Record a user's vote and broadcast an SSE update.
:param user_id: The voter's unique identifier. :param user: The voter.
:param selections: Set of selected option indices (0-based). :param selections: Set of selected option indices (0-based).
:raises PollError: If no poll is active or the selections are invalid. :raises PollError: If no poll is active or the selections are invalid.
""" """
@@ -383,16 +275,10 @@ class PollManager:
raise PollError("No active poll.") raise PollError("No active poll.")
# Check authentication requirement. # Check authentication requirement.
if poll.requires_auth: if poll.requires_auth and not user.is_authenticated:
authenticated = False raise PollError("This poll requires an authenticated account to vote.")
for token_obj in poll.tokens.values():
if token_obj.user_id == user_id:
authenticated = token_obj.is_authenticated
break
if not authenticated:
raise PollError("This poll requires an authenticated account to vote.")
if user_id in poll.votes: if user.id in poll.votes:
raise PollError("You have already voted.") raise PollError("You have already voted.")
if not selections: if not selections:
@@ -414,7 +300,7 @@ class PollManager:
if poll.max_selections is not None and len(selections) > poll.max_selections: if poll.max_selections is not None and len(selections) > poll.max_selections:
raise PollError(f"At most {poll.max_selections} selections are allowed.") raise PollError(f"At most {poll.max_selections} selections are allowed.")
poll.votes[user_id] = selections poll.votes[user.id] = selections
# Broadcast SSE update. # Broadcast SSE update.
results = compute_results(poll.options, poll.votes) results = compute_results(poll.options, poll.votes)
@@ -553,14 +439,17 @@ class PollManager:
await self._close_sse_clients(poll) await self._close_sse_clients(poll)
async def register_sse_client( async def register_sse_client(
self, token: str, response: web.StreamResponse self,
user: User,
*,
response: web.StreamResponse,
) -> None: ) -> None:
"""Register an SSE client connection, replacing any existing one for the token. """Register an SSE client connection for a user.
Sends an initial ``tally`` event so the client has current Sends an initial ``tally`` event so the client has current
results without needing them baked into the page. results without needing them baked into the page.
:param token: The client's token (mod or voter). :param user: Owncast user associated with the browser session.
:param response: The SSE stream response object. :param response: The SSE stream response object.
:raises PollError: If no poll is currently active. :raises PollError: If no poll is currently active.
""" """
@@ -568,14 +457,14 @@ class PollManager:
if poll is None: if poll is None:
raise PollError("No active poll.") raise PollError("No active poll.")
if token not in poll.tokens: poll.sse_clients[user.id] = PollSSEClient(
raise PollError("Invalid token.") user_id=user.id,
is_moderator=user.is_moderator,
poll.sse_clients[token] = response response=response,
)
# Send initial vote counts if this client should see results. # Send initial vote counts if this client should see results.
token_obj = poll.tokens[token] if user.is_moderator or not poll.hidden:
if token_obj.is_moderator or not poll.hidden:
results = compute_results(poll.options, poll.votes) results = compute_results(poll.options, poll.votes)
payload = sse_payload( payload = sse_payload(
SSEEvent.TALLY, SSEEvent.TALLY,
@@ -591,8 +480,7 @@ class PollManager:
"""Start the stream-stop grace period. """Start the stream-stop grace period.
If the stream does not restart within :data:`STREAM_GRACE_PERIOD` If the stream does not restart within :data:`STREAM_GRACE_PERIOD`
seconds, any pending creation token is cleared and any active poll seconds, any active poll is cancelled.
is cancelled.
""" """
self.cancel_stream_grace() self.cancel_stream_grace()
@@ -602,7 +490,6 @@ class PollManager:
STREAM_GRACE_PERIOD, STREAM_GRACE_PERIOD,
) )
await asyncio.sleep(STREAM_GRACE_PERIOD) await asyncio.sleep(STREAM_GRACE_PERIOD)
self.cancel_creation_token()
if self._active_poll is not None: if self._active_poll is not None:
await self.cancel() await self.cancel()
self._ctx.logger.info( self._ctx.logger.info(
@@ -621,10 +508,9 @@ class PollManager:
async def teardown(self) -> None: async def teardown(self) -> None:
"""Clean up on module unload. """Clean up on module unload.
Cancels grace period, creation token, and any active poll. Cancels grace period and any active poll.
""" """
self.cancel_stream_grace() self.cancel_stream_grace()
self.cancel_creation_token()
self._cancel_result_expiry() self._cancel_result_expiry()
if self._active_poll is not None: if self._active_poll is not None:
await self.cancel() await self.cancel()
@@ -634,8 +520,8 @@ class PollManager:
) -> None: ) -> None:
"""Broadcast an SSE event to connected clients. """Broadcast an SSE event to connected clients.
For ``tally`` events on hidden polls, only the moderator client For ``tally`` events on hidden polls, only moderator clients receive
receives the update. All other event types are sent to every client. the update. All other event types are sent to every client.
:param poll: The active poll. :param poll: The active poll.
:param event_type: The SSE event type. :param event_type: The SSE event type.
@@ -646,12 +532,11 @@ class PollManager:
# Determine which clients to send to. # Determine which clients to send to.
if poll.hidden and event_type == SSEEvent.TALLY: if poll.hidden and event_type == SSEEvent.TALLY:
# Only privileged tokens receive vote updates for hidden polls. # Only moderators receive vote updates for hidden polls.
targets = { targets = {
token_value: client user_id: client
for token_value, client in clients.items() for user_id, client in clients.items()
if (token_obj := poll.tokens.get(token_value)) if client.is_moderator
and token_obj.is_moderator
} }
else: else:
targets = dict(clients) targets = dict(clients)
@@ -660,18 +545,18 @@ class PollManager:
return return
# Send to all targets concurrently. # Send to all targets concurrently.
tokens = list(targets.keys()) user_ids = list(targets.keys())
responses = list(targets.values()) responses = list(targets.values())
results = await asyncio.gather( results = await asyncio.gather(
*(r.write(payload) for r in responses), *(client.response.write(payload) for client in responses),
return_exceptions=True, return_exceptions=True,
) )
# Remove failed clients. # Remove failed clients.
for token, result in zip(tokens, results, strict=True): for user_id, result in zip(user_ids, results, strict=True):
if isinstance(result, BaseException): if isinstance(result, BaseException):
self._ctx.logger.debug("Removing disconnected SSE client: %s", token) self._ctx.logger.debug("Removing disconnected SSE client: %s", user_id)
clients.pop(token, None) clients.pop(user_id, None)
async def _close_sse_clients(self, poll: ActivePoll) -> None: async def _close_sse_clients(self, poll: ActivePoll) -> None:
"""Close all SSE client connections for a poll. """Close all SSE client connections for a poll.
@@ -680,7 +565,7 @@ class PollManager:
""" """
for client in poll.sse_clients.values(): for client in poll.sse_clients.values():
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
await client.write_eof() await client.response.write_eof()
poll.sse_clients.clear() poll.sse_clients.clear()
+58 -74
View File
@@ -18,10 +18,9 @@ from __future__ import annotations
import asyncio import asyncio
import contextlib import contextlib
import secrets
from http import HTTPStatus from http import HTTPStatus
from pathlib import Path from pathlib import Path
from typing import Any from typing import TYPE_CHECKING, Any
from aiohttp import web from aiohttp import web
@@ -39,6 +38,9 @@ from .types import (
sse_payload, sse_payload,
) )
if TYPE_CHECKING:
from owlbot.sessions import BrowserSession
_STATIC_DIR = Path(__file__).resolve().parent / "static" _STATIC_DIR = Path(__file__).resolve().parent / "static"
_POLLS_CSS_PATH = _STATIC_DIR / "polls.css" _POLLS_CSS_PATH = _STATIC_DIR / "polls.css"
_POLLS_JS_PATH = _STATIC_DIR / "polls.js" _POLLS_JS_PATH = _STATIC_DIR / "polls.js"
@@ -57,6 +59,14 @@ def _error_page(ctx: RouteContext, status: int, message: str) -> web.Response:
return web.Response(status=status, text=page, content_type="text/html") return web.Response(status=status, text=page, content_type="text/html")
def _require_session(ctx: RouteContext) -> BrowserSession:
"""Return the browser session for a protected polls route."""
session = ctx.session
if session is None:
raise RuntimeError("Protected polls route called without a browser session.")
return session
def _default_form() -> dict[str, Any]: def _default_form() -> dict[str, Any]:
"""Return default form values for the poll creation template.""" """Return default form values for the poll creation template."""
return { return {
@@ -73,7 +83,6 @@ def _default_form() -> dict[str, Any]:
def _create_form_page( def _create_form_page(
ctx: RouteContext, ctx: RouteContext,
token: str,
*, *,
form: dict[str, Any], form: dict[str, Any],
error: str | None = None, error: str | None = None,
@@ -82,7 +91,6 @@ def _create_form_page(
"""Render the poll creation form. """Render the poll creation form.
:param ctx: The route context. :param ctx: The route context.
:param token: The creation token.
:param form: Form field values to populate. :param form: Form field values to populate.
:param error: Optional error message to display. :param error: Optional error message to display.
:param status: HTTP status code. :param status: HTTP status code.
@@ -97,40 +105,31 @@ def _create_form_page(
} }
page = ctx.templates.render( page = ctx.templates.render(
"create.html", "create.html",
token=token,
form=form, form=form,
error=error, error=error,
constraints=constraints, constraints=constraints,
submit_url=ctx.routes.url_for(f"/create/{token}"), submit_url=ctx.routes.url_for("/create"),
create_js_url=ctx.routes.url_for("/static/create.js"), create_js_url=ctx.routes.url_for("/static/create.js"),
) )
return web.Response(status=status, text=page, content_type="text/html") return web.Response(status=status, text=page, content_type="text/html")
@on_route("/create/{token}", methods=["GET"]) @on_route("/create", methods=["GET"], requires_moderator=True)
async def create_page(ctx: RouteContext) -> web.Response: async def create_page(ctx: RouteContext) -> web.Response:
"""Serve the poll creation form. """Serve the poll creation form.
:param ctx: The route context. :param ctx: The route context.
:return: HTML response with the creation form. :return: HTML response with the creation form.
""" """
token = ctx.match_info["token"]
manager = get_manager(ctx.module) manager = get_manager(ctx.module)
if manager.active_poll is not None: if manager.active_poll is not None:
return _error_page(ctx, HTTPStatus.CONFLICT, "A poll is already active.") return _error_page(ctx, HTTPStatus.CONFLICT, "A poll is already active.")
if manager.creation_token_value is None or not secrets.compare_digest( return _create_form_page(ctx, form=_default_form())
manager.creation_token_value, token
):
return _error_page(
ctx, HTTPStatus.FORBIDDEN, "Invalid or expired creation link."
)
return _create_form_page(ctx, token, form=_default_form())
@on_route("/create/{token}", methods=["POST"]) @on_route("/create", methods=["POST"], requires_moderator=True)
async def create_submit(ctx: RouteContext) -> web.Response: async def create_submit(ctx: RouteContext) -> web.Response:
"""Handle poll creation form submission. """Handle poll creation form submission.
@@ -140,7 +139,6 @@ async def create_submit(ctx: RouteContext) -> web.Response:
:param ctx: The route context. :param ctx: The route context.
:return: Redirect to the voting page, or error response. :return: Redirect to the voting page, or error response.
""" """
token = ctx.match_info["token"]
data = await ctx.request.post() data = await ctx.request.post()
# Parse form inputs. # Parse form inputs.
@@ -188,7 +186,6 @@ async def create_submit(ctx: RouteContext) -> web.Response:
manager = get_manager(ctx.module) manager = get_manager(ctx.module)
try: try:
await manager.create_poll( await manager.create_poll(
token=token,
question=question, question=question,
options=options, options=options,
hidden=hidden, hidden=hidden,
@@ -200,20 +197,20 @@ async def create_submit(ctx: RouteContext) -> web.Response:
) )
except PollError as e: except PollError as e:
return _create_form_page( return _create_form_page(
ctx, token, form=form, error=str(e), status=HTTPStatus.BAD_REQUEST ctx, form=form, error=str(e), status=HTTPStatus.BAD_REQUEST
) )
# Redirect to the mod's voting page. # Redirect to the mod's voting page.
vote_url = ctx.routes.url_for(f"/vote/{token}") vote_url = ctx.routes.url_for("/vote")
raise web.HTTPFound(vote_url) raise web.HTTPFound(vote_url)
@on_route("/vote/{token}", methods=["GET"]) @on_route("/vote", methods=["GET"], requires_session=True)
async def vote_page(ctx: RouteContext) -> web.Response: async def vote_page(ctx: RouteContext) -> web.Response:
"""Serve the voting page. """Serve the voting page.
Validates the token and renders the voting page with appropriate Renders the voting page with appropriate visibility and permission controls
visibility and permission controls based on the token's privileges. based on the connected viewer's session.
:param ctx: The route context. :param ctx: The route context.
:return: HTML response with the voting page. :return: HTML response with the voting page.
@@ -224,34 +221,35 @@ async def vote_page(ctx: RouteContext) -> web.Response:
if poll is None: if poll is None:
return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.") return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.")
token = ctx.match_info["token"] session = _require_session(ctx)
token_obj = poll.tokens.get(token) user_id = session.user.id
has_voted = user_id in poll.votes
if token_obj is None: current_votes = poll.votes.get(user_id, set())
return _error_page(ctx, HTTPStatus.FORBIDDEN, "Invalid or expired voting link.") vote_url = ctx.routes.url_for("/vote")
has_voted = token_obj.user_id in poll.votes
current_votes = poll.votes.get(token_obj.user_id, set())
page = ctx.templates.render( page = ctx.templates.render(
"vote.html", "vote.html",
poll=poll, poll=poll,
token=token_obj, viewer_is_authenticated=session.is_authenticated,
viewer_is_moderator=session.is_moderator,
has_voted=has_voted, has_voted=has_voted,
current_votes=current_votes, current_votes=current_votes,
vote_url=ctx.routes.url_for(f"/vote/{token}"), vote_url=vote_url,
events_url=ctx.routes.url_for("/vote/events"),
end_url=ctx.routes.url_for("/vote/end"),
cancel_url=ctx.routes.url_for("/vote/cancel"),
polls_css_url=ctx.routes.url_for("/static/polls.css"), polls_css_url=ctx.routes.url_for("/static/polls.css"),
polls_js_url=ctx.routes.url_for("/static/polls.js"), polls_js_url=ctx.routes.url_for("/static/polls.js"),
) )
return web.Response(text=page, content_type="text/html") return web.Response(text=page, content_type="text/html")
@on_route("/vote/{token}", methods=["POST"]) @on_route("/vote", methods=["POST"], requires_session=True)
async def vote_submit(ctx: RouteContext) -> web.Response: async def vote_submit(ctx: RouteContext) -> web.Response:
"""Handle vote form submission. """Handle vote form submission.
Validates the token, parses selections, delegates to the manager, Parses selections, delegates to the manager, and redirects back to the
and redirects back to the voting page. voting page.
:param ctx: The route context. :param ctx: The route context.
:return: Redirect back to the voting page, or error response. :return: Redirect back to the voting page, or error response.
@@ -262,21 +260,15 @@ async def vote_submit(ctx: RouteContext) -> web.Response:
if poll is None: if poll is None:
return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.") return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.")
token = ctx.match_info["token"] session = _require_session(ctx)
token_obj = poll.tokens.get(token)
if token_obj is None: if poll.requires_auth and not session.is_authenticated:
return _error_page(ctx, HTTPStatus.FORBIDDEN, "Invalid voting link.")
if poll.requires_auth and not token_obj.is_authenticated:
return _error_page( return _error_page(
ctx, ctx,
HTTPStatus.FORBIDDEN, HTTPStatus.FORBIDDEN,
"This poll requires an authenticated account to vote.", "This poll requires an authenticated account to vote.",
) )
voter_id = token_obj.user_id
data = await ctx.request.post() data = await ctx.request.post()
# Parse selections from form checkboxes/radio buttons. # Parse selections from form checkboxes/radio buttons.
@@ -287,16 +279,19 @@ async def vote_submit(ctx: RouteContext) -> web.Response:
return _error_page(ctx, HTTPStatus.BAD_REQUEST, "Invalid selection.") return _error_page(ctx, HTTPStatus.BAD_REQUEST, "Invalid selection.")
try: try:
await manager.record_vote(voter_id, selections) await manager.record_vote(
session.user,
selections,
)
except PollError as e: except PollError as e:
return _error_page(ctx, HTTPStatus.BAD_REQUEST, str(e)) return _error_page(ctx, HTTPStatus.BAD_REQUEST, str(e))
# Redirect back to the voting page. # Redirect back to the voting page.
vote_url = ctx.routes.url_for(f"/vote/{token}") vote_url = ctx.routes.url_for("/vote")
raise web.HTTPFound(vote_url) raise web.HTTPFound(vote_url)
@on_route("/vote/{token}/events", streaming=True) @on_route("/vote/events", streaming=True, requires_session=True)
async def events_stream(ctx: RouteContext) -> web.StreamResponse | web.Response: async def events_stream(ctx: RouteContext) -> web.StreamResponse | web.Response:
"""SSE endpoint for live poll updates. """SSE endpoint for live poll updates.
@@ -311,9 +306,8 @@ async def events_stream(ctx: RouteContext) -> web.StreamResponse | web.Response:
if poll is None: if poll is None:
return web.Response(status=HTTPStatus.NOT_FOUND, text="No active poll.") return web.Response(status=HTTPStatus.NOT_FOUND, text="No active poll.")
token = ctx.match_info["token"] session = _require_session(ctx)
if token not in poll.tokens: user_id = session.user.id
return web.Response(status=HTTPStatus.FORBIDDEN, text="Invalid token.")
response = web.StreamResponse( response = web.StreamResponse(
status=HTTPStatus.OK, status=HTTPStatus.OK,
@@ -326,7 +320,10 @@ async def events_stream(ctx: RouteContext) -> web.StreamResponse | web.Response:
) )
await response.prepare(ctx.request) await response.prepare(ctx.request)
await manager.register_sse_client(token, response) await manager.register_sse_client(
session.user,
response=response,
)
try: try:
while manager.active_poll is poll: while manager.active_poll is poll:
@@ -340,24 +337,25 @@ async def events_stream(ctx: RouteContext) -> web.StreamResponse | web.Response:
payload = sse_payload(SSEEvent.KEEPALIVE, {"time_remaining": remaining}) payload = sse_payload(SSEEvent.KEEPALIVE, {"time_remaining": remaining})
await response.write(payload) await response.write(payload)
except Exception: # noqa: BLE001 # SSE disconnect; no single exception covers all transport failures except Exception: # noqa: BLE001 # SSE disconnect; no single exception covers all transport failures
ctx.logger.debug("SSE keepalive failed for %s...", token[:8]) ctx.logger.debug("SSE keepalive failed for %s.", user_id)
break break
finally: finally:
# Only remove ourselves if we're still the registered client. # Only remove ourselves if we're still the registered client.
# A page reload creates a new SSE connection under the same token, # A page reload creates a new SSE connection for the same user,
# so the old handler must not evict the replacement. # so the old handler must not evict the replacement.
if poll.sse_clients.get(token) is response: client = poll.sse_clients.get(user_id)
poll.sse_clients.pop(token, None) if client is not None and client.response is response:
poll.sse_clients.pop(user_id, None)
ctx.logger.debug( ctx.logger.debug(
"SSE client disconnected: %s... (%d remaining).", "SSE client disconnected: %s (%d remaining).",
token[:8], user_id,
len(poll.sse_clients), len(poll.sse_clients),
) )
return response return response
@on_route("/vote/{token}/end", methods=["POST"]) @on_route("/vote/end", methods=["POST"], requires_moderator=True)
async def mod_end(ctx: RouteContext) -> web.Response: async def mod_end(ctx: RouteContext) -> web.Response:
"""End the poll via the moderator web UI. """End the poll via the moderator web UI.
@@ -370,19 +368,12 @@ async def mod_end(ctx: RouteContext) -> web.Response:
if poll is None: if poll is None:
return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.") return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.")
token = ctx.match_info["token"]
token_obj = poll.tokens.get(token)
if token_obj is None or not token_obj.is_moderator:
return _error_page(
ctx, HTTPStatus.FORBIDDEN, "Only moderators can end the poll."
)
await manager.end() await manager.end()
results_url = ctx.routes.url_for("/results") results_url = ctx.routes.url_for("/results")
raise web.HTTPFound(results_url) raise web.HTTPFound(results_url)
@on_route("/vote/{token}/cancel", methods=["POST"]) @on_route("/vote/cancel", methods=["POST"], requires_moderator=True)
async def mod_cancel(ctx: RouteContext) -> web.Response: async def mod_cancel(ctx: RouteContext) -> web.Response:
"""Cancel the poll via the moderator web UI. """Cancel the poll via the moderator web UI.
@@ -395,13 +386,6 @@ async def mod_cancel(ctx: RouteContext) -> web.Response:
if poll is None: if poll is None:
return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.") return _error_page(ctx, HTTPStatus.NOT_FOUND, "No active poll.")
token = ctx.match_info["token"]
token_obj = poll.tokens.get(token)
if token_obj is None or not token_obj.is_moderator:
return _error_page(
ctx, HTTPStatus.FORBIDDEN, "Only moderators can cancel the poll."
)
await manager.cancel() await manager.cancel()
return _error_page(ctx, HTTPStatus.OK, "This poll has been cancelled.") return _error_page(ctx, HTTPStatus.OK, "This poll has been cancelled.")
+4 -4
View File
@@ -16,8 +16,8 @@
"use strict"; "use strict";
const config = window.POLL_CONFIG; const config = window.POLL_CONFIG;
const tokenInfo = window.TOKEN_INFO; const viewerInfo = window.VIEWER_INFO;
if (!config || !tokenInfo) return; if (!config || !viewerInfo) return;
const timerEl = document.getElementById("timer"); const timerEl = document.getElementById("timer");
let remaining = config.timeRemaining; let remaining = config.timeRemaining;
@@ -71,7 +71,7 @@
// Show vote counts immediately for non-hidden polls so the UI // Show vote counts immediately for non-hidden polls so the UI
// doesn't look identical to a hidden poll before SSE delivers data. // doesn't look identical to a hidden poll before SSE delivers data.
if (!config.hidden || tokenInfo.isModerator) { if (!config.hidden || viewerInfo.isModerator) {
for (const el of document.querySelectorAll(".vote-count")) { for (const el of document.querySelectorAll(".vote-count")) {
el.classList.remove("d-none"); el.classList.remove("d-none");
} }
@@ -79,7 +79,7 @@
// --- SSE connection --- // --- SSE connection ---
const evtSource = new EventSource(tokenInfo.eventsUrl); const evtSource = new EventSource(viewerInfo.eventsUrl);
const parse = (e) => { const parse = (e) => {
try { try {
@@ -8,7 +8,7 @@
<div id="poll-status" class="mb-4"> <div id="poll-status" class="mb-4">
<p class="text-body-secondary mb-1">Time remaining: <strong id="timer">--:--</strong></p> <p class="text-body-secondary mb-1">Time remaining: <strong id="timer">--:--</strong></p>
{% if poll.hidden %} {% if poll.hidden %}
{% if token.is_moderator %} {% if viewer_is_moderator %}
<p class="text-body-secondary mb-0">Results are hidden until this poll ends. As a moderator, you can still see the real-time results.</p> <p class="text-body-secondary mb-0">Results are hidden until this poll ends. As a moderator, you can still see the real-time results.</p>
{% else %} {% else %}
<p class="text-body-secondary mb-0">Results are hidden until this poll ends.</p> <p class="text-body-secondary mb-0">Results are hidden until this poll ends.</p>
@@ -27,11 +27,11 @@
{% if poll.allow_multiple %} {% if poll.allow_multiple %}
<input type="checkbox" name="option" value="{{ loop.index0 }}" class="form-check-input me-2" <input type="checkbox" name="option" value="{{ loop.index0 }}" class="form-check-input me-2"
{% if loop.index0 in current_votes %}checked{% endif %} {% if loop.index0 in current_votes %}checked{% endif %}
{% if has_voted or (not token.is_authenticated and poll.requires_auth) %}disabled{% endif %}> {% if has_voted or (not viewer_is_authenticated and poll.requires_auth) %}disabled{% endif %}>
{% else %} {% else %}
<input type="radio" name="option" value="{{ loop.index0 }}" class="form-check-input me-2" required <input type="radio" name="option" value="{{ loop.index0 }}" class="form-check-input me-2" required
{% if loop.index0 in current_votes %}checked{% endif %} {% if loop.index0 in current_votes %}checked{% endif %}
{% if has_voted or (not token.is_authenticated and poll.requires_auth) %}disabled{% endif %}> {% if has_voted or (not viewer_is_authenticated and poll.requires_auth) %}disabled{% endif %}>
{% endif %} {% endif %}
{{ option }} {{ option }}
</label> </label>
@@ -42,7 +42,7 @@
{% endfor %} {% endfor %}
</div> </div>
{% if poll.requires_auth and not token.is_authenticated %} {% if poll.requires_auth and not viewer_is_authenticated %}
<div class="alert alert-warning mb-0" role="alert">This poll requires an authenticated account to vote. You are viewing results only.</div> <div class="alert alert-warning mb-0" role="alert">This poll requires an authenticated account to vote. You are viewing results only.</div>
{% elif not has_voted %} {% elif not has_voted %}
<button type="submit" class="btn btn-primary me-2">Vote</button> <button type="submit" class="btn btn-primary me-2">Vote</button>
@@ -51,15 +51,15 @@
{% endif %} {% endif %}
</form> </form>
{% if token.is_moderator %} {% if viewer_is_moderator %}
<hr class="my-4"> <hr class="my-4">
<h5>Manage Poll</h5> <h5>Manage Poll</h5>
<p class="text-body-secondary">As a moderator, you can perform the following actions on this poll:</p> <p class="text-body-secondary">As a moderator, you can perform the following actions on this poll:</p>
<div class="d-flex gap-2"> <div class="d-flex gap-2">
<form method="POST" action="{{ vote_url }}/end"> <form method="POST" action="{{ end_url }}">
<button type="submit" class="btn btn-success">End Poll</button> <button type="submit" class="btn btn-success">End Poll</button>
</form> </form>
<form method="POST" action="{{ vote_url }}/cancel"> <form method="POST" action="{{ cancel_url }}">
<button type="submit" class="btn btn-danger">Cancel Poll</button> <button type="submit" class="btn btn-danger">Cancel Poll</button>
</form> </form>
</div> </div>
@@ -74,9 +74,9 @@ window.POLL_CONFIG = {
minSelections: {{ poll.min_selections | tojson }}, minSelections: {{ poll.min_selections | tojson }},
maxSelections: {{ poll.max_selections | tojson }} maxSelections: {{ poll.max_selections | tojson }}
}; };
window.TOKEN_INFO = { window.VIEWER_INFO = {
isModerator: {{ token.is_moderator | tojson }}, isModerator: {{ viewer_is_moderator | tojson }},
eventsUrl: {{ (vote_url + "/events") | tojson }} eventsUrl: {{ events_url | tojson }}
}; };
</script> </script>
<script src="{{ polls_js_url }}?v={{ owlbot_version }}"></script> <script src="{{ polls_js_url }}?v={{ owlbot_version }}"></script>
+8 -14
View File
@@ -17,7 +17,6 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import secrets
import time import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import StrEnum from enum import StrEnum
@@ -39,7 +38,6 @@ DEFAULT_DURATION = 240
MIN_DURATION = 30 MIN_DURATION = 30
MAX_DURATION = 600 MAX_DURATION = 600
REMINDER_THRESHOLD = 60 REMINDER_THRESHOLD = 60
CREATION_TOKEN_TIMEOUT = 15 * 60
STREAM_GRACE_PERIOD = 5 * 60 STREAM_GRACE_PERIOD = 5 * 60
RESULT_EXPIRY = 3600 RESULT_EXPIRY = 3600
@@ -64,19 +62,17 @@ def sse_payload(event_type: SSEEvent, data: dict[str, Any]) -> bytes:
@dataclass(slots=True) @dataclass(slots=True)
class Token: class PollSSEClient:
"""A token for poll access. """A connected live-results stream for a poll viewer.
:param user_id: The Owncast user ID this token belongs to. :param user_id: Owncast user ID associated with the browser session.
:param is_moderator: Whether the user has moderator privileges. :param is_moderator: Whether this viewer can see hidden live tallies.
:param is_authenticated: Whether the user has an authenticated Owncast account. :param response: The active SSE response.
:param value: The token string used in URLs (auto-generated if not provided).
""" """
user_id: str user_id: str
is_moderator: bool is_moderator: bool
is_authenticated: bool response: web.StreamResponse
value: str = field(default_factory=lambda: secrets.token_urlsafe(32))
@dataclass(slots=True) @dataclass(slots=True)
@@ -93,9 +89,8 @@ class ActivePoll:
:param duration: Poll duration in seconds. :param duration: Poll duration in seconds.
:param created_at: Timestamp when the poll was created. :param created_at: Timestamp when the poll was created.
:param votes: Mapping of user ID to set of selected option indices. :param votes: Mapping of user ID to set of selected option indices.
:param sse_clients: Mapping of client token to SSE stream response. :param sse_clients: Mapping of user ID to SSE client response metadata.
:param timer_task: Asyncio task for the poll timer. :param timer_task: Asyncio task for the poll timer.
:param tokens: Mapping of token string to Token object.
""" """
question: str question: str
@@ -108,9 +103,8 @@ class ActivePoll:
duration: int duration: int
created_at: datetime created_at: datetime
votes: dict[str, set[int]] = field(default_factory=dict) votes: dict[str, set[int]] = field(default_factory=dict)
sse_clients: dict[str, web.StreamResponse] = field(default_factory=dict) sse_clients: dict[str, PollSSEClient] = field(default_factory=dict)
timer_task: asyncio.Task[None] | None = None timer_task: asyncio.Task[None] | None = None
tokens: dict[str, Token] = field(default_factory=dict)
done_event: asyncio.Event = field(default_factory=asyncio.Event) done_event: asyncio.Event = field(default_factory=asyncio.Event)
_started_at_mono: float = field(default_factory=time.monotonic) _started_at_mono: float = field(default_factory=time.monotonic)