Refactored OwncastSentry internals and API validation.
CI / Formatting (push) Failing after 7s
CI / Linting (push) Successful in 12s
CI / Tests (push) Successful in 47s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 7s

This commit is contained in:
2026-05-17 14:46:44 -04:00
parent 179d087e33
commit 0620c675d9
25 changed files with 2730 additions and 1569 deletions
+9 -4
View File
@@ -23,12 +23,12 @@ from prometheus_client.exposition import choose_encoder
from .commands import CommandHandler from .commands import CommandHandler
from .config import Config from .config import Config
from .database import StreamRepository, SubscriptionRepository
from .metrics import ErrorSource, MetricsService from .metrics import ErrorSource, MetricsService
from .migrations import get_upgrade_table
from .notification_service import NotificationService from .notification_service import NotificationService
from .owncast_client import OwncastClient from .owncast_client import OwncastClient
from .repository import StreamRepository, SubscriptionRepository, get_upgrade_table
from .stream_monitor import StreamMonitor from .stream_monitor import StreamMonitor
from .subscription_manager import SubscriptionManager
if TYPE_CHECKING: if TYPE_CHECKING:
from mautrix.util.async_db import Database, UpgradeTable from mautrix.util.async_db import Database, UpgradeTable
@@ -97,14 +97,19 @@ class OwncastSentry(Plugin):
metrics=self.metrics_service, metrics=self.metrics_service,
) )
# Initialize command handler # Initialize subscription manager
self.command_handler = CommandHandler( self.subscription_manager = SubscriptionManager(
self.owncast_client, self.owncast_client,
self.stream_repo, self.stream_repo,
self.subscription_repo, self.subscription_repo,
self.log, self.log,
) )
# Initialize command handler
self.command_handler = CommandHandler(
self.subscription_manager,
)
# Schedule periodic stream state updates every 60 seconds # Schedule periodic stream state updates every 60 seconds
self.sched.run_periodically(60, self._update_all_stream_states) self.sched.run_periodically(60, self._update_all_stream_states)
+112 -139
View File
@@ -14,20 +14,71 @@
"""Command handlers for OwncastSentry bot commands.""" """Command handlers for OwncastSentry bot commands."""
import sqlite3
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .models import StreamStatus from .types import (
from .utils import domainify, sanitize_for_markdown AlreadySubscribedError,
InvalidOwncastInstanceError,
NotSubscribedError,
StreamStatus,
)
if TYPE_CHECKING: if TYPE_CHECKING:
import logging
from maubot import MessageEvent # type: ignore[attr-defined] from maubot import MessageEvent # type: ignore[attr-defined]
from .database import StreamRepository, SubscriptionRepository from .subscription_manager import SubscriptionManager
from .owncast_client import OwncastClient
_MARKDOWN_ESCAPE_TABLE = str.maketrans({c: f"\\{c}" for c in r"\*_[]()~`#+-=|{}.!<>&"})
def _sanitize_for_plain_text(text: str) -> str:
"""Sanitize text before Markdown escaping."""
if not text:
return text
sanitized = text.replace("\n", " ").replace("\r", " ")
return " ".join(sanitized.split())
def _escape_markdown(text: str) -> str:
"""Escape Markdown special characters in untrusted text."""
if not text:
return text
return text.translate(_MARKDOWN_ESCAPE_TABLE)
def _sanitize_for_markdown(text: str) -> str:
"""Sanitize text for safe Markdown rendering."""
if not text:
return text
return _escape_markdown(_sanitize_for_plain_text(text))
def _format_duration(timestamp_str: str, now: datetime) -> str:
"""Calculate and format the duration from a timestamp to now."""
try:
timestamp = datetime.fromisoformat(timestamp_str)
delta = now - timestamp
seconds = int(delta.total_seconds())
if seconds < 0:
return "unknown duration"
if seconds < 60:
return f"{seconds} second{'s' if seconds != 1 else ''}"
if seconds < 3600:
minutes = seconds // 60
return f"{minutes} minute{'s' if minutes != 1 else ''}"
if seconds < 86400:
hours = seconds // 3600
return f"{hours} hour{'s' if hours != 1 else ''}"
days = seconds // 86400
return f"{days} day{'s' if days != 1 else ''}"
except (TypeError, ValueError):
return "unknown duration"
class CommandHandler: class CommandHandler:
@@ -35,22 +86,13 @@ class CommandHandler:
def __init__( def __init__(
self, self,
owncast_client: OwncastClient, subscription_manager: SubscriptionManager,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
logger: logging.Logger,
) -> None: ) -> None:
"""Initialize the command handler. """Initialize the command handler.
:param owncast_client: Client for making API calls to Owncast instances. :param subscription_manager: Subscription domain workflow coordinator.
:param stream_repo: Repository for stream data.
:param subscription_repo: Repository for subscription data.
:param logger: Logger instance for debugging.
""" """
self.owncast_client = owncast_client self.subscription_manager = subscription_manager
self.stream_repo = stream_repo
self.subscription_repo = subscription_repo
self.log = logger
async def subscribe(self, evt: MessageEvent, url: str) -> None: async def subscribe(self, evt: MessageEvent, url: str) -> None:
"""Subscribe a room to a stream's notifications. """Subscribe a room to a stream's notifications.
@@ -58,46 +100,22 @@ class CommandHandler:
:param evt: MessageEvent of the message calling the command. :param evt: MessageEvent of the message calling the command.
:param url: User supplied URL to a stream to subscribe to. :param url: User supplied URL to a stream to subscribe to.
""" """
# Convert the user input to only a domain
stream_domain = domainify(url)
# How many subscriptions already exist for this domain?
subscription_count = await self.subscription_repo.count_by_domain(stream_domain)
if subscription_count == 0:
# No subscriptions; validate this is an Owncast stream.
is_valid = await self.owncast_client.validate_instance(stream_domain)
if not is_valid:
# Fetch returned nothing. Probably not Owncast.
await evt.reply(
"The URL you supplied does not appear to "
"be a valid Owncast instance. You may have "
"specified an invalid domain, or the "
"instance is offline."
)
return
# Try to add a new subscription for this stream in this room
try: try:
await self.subscription_repo.add(stream_domain, evt.room_id) stream_domain = await self.subscription_manager.subscribe(evt.room_id, url)
except sqlite3.IntegrityError: except InvalidOwncastInstanceError:
# Room is already subscribed.
await evt.reply( await evt.reply(
f"This room is already subscribed to notifications for {stream_domain}." "The URL you supplied does not appear to "
"be a valid Owncast instance. You may have "
"specified an invalid domain, or the "
"instance is offline."
)
return
except AlreadySubscribedError as e:
await evt.reply(
f"This room is already subscribed to notifications for {e.domain}."
) )
return return
# Try to add a placeholder row for the stream's state.
try:
await self.stream_repo.create(stream_domain)
# First time seeing this stream. Log it.
self.log.info(f"[{stream_domain}] Discovered new stream!")
except sqlite3.IntegrityError:
# Adding rows for known streams is expected.
pass
# All went well! Tell the user.
self.log.info(f"[{stream_domain}] Subscription added for room {evt.room_id}.")
await evt.reply( await evt.reply(
f"Subscription added! This room will receive " f"Subscription added! This room will receive "
f"notifications when {stream_domain} goes live." f"notifications when {stream_domain} goes live."
@@ -109,66 +127,32 @@ class CommandHandler:
:param evt: MessageEvent of the message calling the command. :param evt: MessageEvent of the message calling the command.
:param url: User supplied URL to a stream to unsubscribe from. :param url: User supplied URL to a stream to unsubscribe from.
""" """
# Convert the user input to only a domain try:
stream_domain = domainify(url) stream_domain = await self.subscription_manager.unsubscribe(
evt.room_id, url
# Attempt to delete the requested subscription
result = await self.subscription_repo.remove(stream_domain, evt.room_id)
# Did it work?
if result == 1:
# Yes, one row was deleted. Tell the user.
self.log.info(
f"[{stream_domain}] Subscription removed for room {evt.room_id}."
) )
await evt.reply( except NotSubscribedError as e:
f"Subscription removed! This room will no "
f"longer receive notifications for {stream_domain}."
)
else:
# No, nothing changed. Tell the user.
await evt.reply( await evt.reply(
"This room is already not subscribed to " "This room is already not subscribed to "
f"notifications for {stream_domain}." f"notifications for {e.domain}."
) )
return
def _format_duration(self, timestamp_str: str) -> str: await evt.reply(
"""Calculate and format the duration from a timestamp to now. f"Subscription removed! This room will no "
f"longer receive notifications for {stream_domain}."
:param timestamp_str: ISO 8601 timestamp string. )
:return: Formatted duration string (e.g., "1 hour", "2 days").
"""
try:
timestamp = datetime.fromisoformat(timestamp_str)
now = datetime.now(UTC)
delta = now - timestamp
seconds = int(delta.total_seconds())
if seconds < 60:
return f"{seconds} second{'s' if seconds != 1 else ''}"
if seconds < 3600:
minutes = seconds // 60
return f"{minutes} minute{'s' if minutes != 1 else ''}"
if seconds < 86400:
hours = seconds // 3600
return f"{hours} hour{'s' if hours != 1 else ''}"
days = seconds // 86400
return f"{days} day{'s' if days != 1 else ''}"
except ValueError:
return "unknown duration"
async def subscriptions(self, evt: MessageEvent) -> None: async def subscriptions(self, evt: MessageEvent) -> None:
"""List all stream subscriptions in the current room. """List all stream subscriptions in the current room.
:param evt: MessageEvent of the message calling the command. :param evt: MessageEvent of the message calling the command.
""" """
# Get all stream domains this room is subscribed to subscriptions = await self.subscription_manager.list_room_subscriptions(
subscribed_domains = ( evt.room_id
await self.subscription_repo.get_subscribed_streams_for_room(evt.room_id)
) )
# Check if there are no subscriptions if not subscriptions:
if not subscribed_domains:
await evt.reply( await evt.reply(
"This room is not subscribed to any Owncast " "This room is not subscribed to any Owncast "
"instances.\n\nTo subscribe to an Owncast " "instances.\n\nTo subscribe to an Owncast "
@@ -178,36 +162,33 @@ class CommandHandler:
return return
# Build the response message body as Markdown # Build the response message body as Markdown
count = len(subscribed_domains) count = len(subscriptions)
parts = [f"**Subscriptions for this room ({count}):**\n\n"] parts = [f"**Subscriptions for this room ({count}):**\n\n"]
now = datetime.now(UTC)
for domain in subscribed_domains: for subscription in subscriptions:
# Get the stream state from the database domain = subscription.domain
stream_state = await self.stream_repo.get_by_domain(domain) stream_state = subscription.stream_state
if stream_state is None:
continue
# Determine stream name (use domain as fallback)
stream_name = stream_state.name or domain stream_name = stream_state.name or domain
safe_stream_name = sanitize_for_markdown(stream_name) safe_stream_name = _sanitize_for_markdown(stream_name)
# Start building this stream's entry with stream name as main bullet # Start building this stream's entry with stream name as main bullet
parts.append(f"- **{safe_stream_name}** \n") parts.append(f"- **{safe_stream_name}** \n")
# Add title if stream is online (as a sub-bullet) # Add title if stream is online (as a sub-bullet)
if stream_state.status == StreamStatus.ONLINE and stream_state.title: if stream_state.status == StreamStatus.ONLINE and stream_state.title:
safe_title = sanitize_for_markdown(stream_state.title) safe_title = _sanitize_for_markdown(stream_state.title)
parts.append(f" - Title: {safe_title} \n") parts.append(f" - Title: {safe_title} \n")
# Determine status and duration (as a sub-bullet) # Determine status and duration (as a sub-bullet)
match stream_state.status: match stream_state.status:
case StreamStatus.ONLINE if stream_state.last_connect_time: case StreamStatus.ONLINE if stream_state.last_connect_time:
duration = self._format_duration(stream_state.last_connect_time) duration = _format_duration(stream_state.last_connect_time, now)
parts.append(f" - Status: Online for {duration} \n") parts.append(f" - Status: Online for {duration} \n")
case StreamStatus.UNKNOWN: case StreamStatus.UNKNOWN:
parts.append(" - Status: Unknown (instance unreachable) \n") parts.append(" - Status: Unknown (instance unreachable) \n")
case StreamStatus.OFFLINE if stream_state.last_disconnect_time: case StreamStatus.OFFLINE if stream_state.last_disconnect_time:
duration = self._format_duration(stream_state.last_disconnect_time) duration = _format_duration(stream_state.last_disconnect_time, now)
parts.append(f" - Status: Offline for {duration} \n") parts.append(f" - Status: Offline for {duration} \n")
case StreamStatus.OFFLINE: case StreamStatus.OFFLINE:
parts.append(" - Status: Offline \n") parts.append(" - Status: Offline \n")
@@ -229,30 +210,20 @@ class CommandHandler:
:param evt: MessageEvent of the message calling the command. :param evt: MessageEvent of the message calling the command.
""" """
# Get all stream domains this room is subscribed to live_streams = await self.subscription_manager.list_live_room_subscriptions(
subscribed_domains = ( evt.room_id
await self.subscription_repo.get_subscribed_streams_for_room(evt.room_id)
) )
# Check if there are no subscriptions
if not subscribed_domains:
await evt.reply(
"This room is not subscribed to any Owncast "
"instances.\n\nTo subscribe to an Owncast "
"instance, use `!subscribe <domain>`",
markdown=True,
)
return
# Filter for only live streams (exclude unknown status)
live_streams = []
for domain in subscribed_domains:
stream_state = await self.stream_repo.get_by_domain(domain)
if stream_state and stream_state.status == StreamStatus.ONLINE:
live_streams.append((domain, stream_state))
# Check if there are no live streams
if not live_streams: if not live_streams:
if not await self.subscription_manager.has_room_subscriptions(evt.room_id):
await evt.reply(
"This room is not subscribed to any Owncast "
"instances.\n\nTo subscribe to an Owncast "
"instance, use `!subscribe <domain>`",
markdown=True,
)
return
await evt.reply( await evt.reply(
"No subscribed Owncast instances are currently " "No subscribed Owncast instances are currently "
"live.\n\nUse `!subscriptions` to list all " "live.\n\nUse `!subscriptions` to list all "
@@ -264,23 +235,25 @@ class CommandHandler:
# Build the response message body as Markdown # Build the response message body as Markdown
count = len(live_streams) count = len(live_streams)
parts = [f"**Live Owncast instances ({count}):**\n\n"] parts = [f"**Live Owncast instances ({count}):**\n\n"]
now = datetime.now(UTC)
for domain, stream_state in live_streams: for subscription in live_streams:
# Determine stream name (use domain as fallback) domain = subscription.domain
stream_state = subscription.stream_state
stream_name = stream_state.name or domain stream_name = stream_state.name or domain
safe_stream_name = sanitize_for_markdown(stream_name) safe_stream_name = _sanitize_for_markdown(stream_name)
# Start building this stream's entry with stream name as main bullet # Start building this stream's entry with stream name as main bullet
parts.append(f"- **{safe_stream_name}** \n") parts.append(f"- **{safe_stream_name}** \n")
# Add title (should be present for live streams) # Add title (should be present for live streams)
if stream_state.title: if stream_state.title:
safe_title = sanitize_for_markdown(stream_state.title) safe_title = _sanitize_for_markdown(stream_state.title)
parts.append(f" - Title: {safe_title} \n") parts.append(f" - Title: {safe_title} \n")
# Add status with duration # Add status with duration
if stream_state.last_connect_time: if stream_state.last_connect_time:
duration = self._format_duration(stream_state.last_connect_time) duration = _format_duration(stream_state.last_connect_time, now)
parts.append(f" - Online for {duration} \n") parts.append(f" - Online for {duration} \n")
# Add stream link # Add stream link
-199
View File
@@ -1,199 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Database repository classes for OwncastSentry."""
from typing import TYPE_CHECKING
from .models import StreamState
if TYPE_CHECKING:
from mautrix.util.async_db import Database
class StreamRepository:
"""Repository for managing stream data in the database."""
def __init__(self, database: Database):
"""Initialize the stream repository.
:param database: The maubot database instance.
"""
self.db = database
async def get_by_domain(self, domain: str) -> StreamState | None:
"""Get a stream's state by domain.
:param domain: The stream domain.
:return: StreamState if found, None otherwise.
"""
query = "SELECT * FROM streams WHERE domain=$1"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
row = await conn.fetchrow(query, domain)
return StreamState.from_db_row(row) if row else None
async def create(self, domain: str) -> None:
"""Create a new stream entry in the database.
:param domain: The stream domain.
"""
query = "INSERT INTO streams (domain) VALUES ($1)"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
await conn.execute(query, domain)
async def update(self, state: StreamState) -> None:
"""Update a stream's state in the database.
:param state: The StreamState to save.
"""
query = """UPDATE streams
SET name=$1, title=$2, last_connect_time=$3, last_disconnect_time=$4
WHERE domain=$5"""
async with self.db.acquire() as conn: # type: ignore[var-annotated]
await conn.execute(
query,
state.name,
state.title,
state.last_connect_time,
state.last_disconnect_time,
state.domain,
)
async def exists(self, domain: str) -> bool:
"""Check if a stream exists in the database.
:param domain: The stream domain.
:return: True if exists, False otherwise.
"""
result = await self.get_by_domain(domain)
return result is not None
async def increment_failure_counter(self, domain: str) -> None:
"""Increment the failure counter for a stream by 1.
:param domain: The stream domain.
"""
query = """UPDATE streams
SET failure_counter = failure_counter + 1
WHERE domain=$1"""
async with self.db.acquire() as conn: # type: ignore[var-annotated]
await conn.execute(query, domain)
async def reset_failure_counter(self, domain: str) -> None:
"""Reset the failure counter for a stream to 0.
:param domain: The stream domain.
"""
query = """UPDATE streams
SET failure_counter = 0
WHERE domain=$1"""
async with self.db.acquire() as conn: # type: ignore[var-annotated]
await conn.execute(query, domain)
async def delete(self, domain: str) -> None:
"""Delete a stream record from the database.
:param domain: The stream domain.
"""
query = "DELETE FROM streams WHERE domain=$1"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
await conn.execute(query, domain)
class SubscriptionRepository:
"""Repository for managing stream subscriptions in the database."""
def __init__(self, database: Database):
"""Initialize the subscription repository.
:param database: The maubot database instance.
"""
self.db = database
async def add(self, domain: str, room_id: str) -> None:
"""Add a subscription for a room to a stream.
:param domain: The stream domain.
:param room_id: The Matrix room ID.
:raises sqlite3.IntegrityError: If subscription already exists.
"""
query = "INSERT INTO subscriptions (stream_domain, room_id) VALUES ($1, $2)"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
await conn.execute(query, domain, room_id)
async def remove(self, domain: str, room_id: str) -> int:
"""Remove a subscription for a room from a stream.
:param domain: The stream domain.
:param room_id: The Matrix room ID.
:return: Number of rows deleted (0 or 1).
"""
query = "DELETE FROM subscriptions WHERE stream_domain=$1 AND room_id=$2"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
result = await conn.execute(query, domain, room_id)
return int(result.rowcount)
async def get_subscribed_rooms(self, domain: str) -> list[str]:
"""Get all room IDs subscribed to a stream.
:param domain: The stream domain.
:return: List of room IDs.
"""
query = "SELECT room_id FROM subscriptions WHERE stream_domain=$1"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
results = await conn.fetch(query, domain)
return [row["room_id"] for row in results]
async def get_subscribed_streams_for_room(self, room_id: str) -> list[str]:
"""Get all stream domains that a room is subscribed to.
:param room_id: The Matrix room ID.
:return: List of stream domains.
"""
query = "SELECT stream_domain FROM subscriptions WHERE room_id=$1"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
results = await conn.fetch(query, room_id)
return [row["stream_domain"] for row in results]
async def get_all_subscribed_domains(self) -> list[str]:
"""Get all unique stream domains that have at least one subscription.
:return: List of stream domains.
"""
query = "SELECT DISTINCT stream_domain FROM subscriptions"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
results = await conn.fetch(query)
return [row["stream_domain"] for row in results]
async def count_by_domain(self, domain: str) -> int:
"""Count the number of subscriptions for a given stream domain.
:param domain: The stream domain.
:return: Number of subscriptions.
"""
query = "SELECT COUNT(*) FROM subscriptions WHERE stream_domain=$1"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
result = await conn.fetchrow(query, domain)
return int(result[0])
async def delete_all_for_domain(self, domain: str) -> int:
"""Delete all subscriptions for a given stream domain.
:param domain: The stream domain.
:return: Number of subscriptions deleted.
"""
query = "DELETE FROM subscriptions WHERE stream_domain=$1"
async with self.db.acquire() as conn: # type: ignore[var-annotated]
result = await conn.execute(query, domain)
return int(result.rowcount)
+1 -1
View File
@@ -21,7 +21,7 @@ from typing import TYPE_CHECKING
from prometheus_client import CollectorRegistry, Counter, Gauge, Info from prometheus_client import CollectorRegistry, Counter, Gauge, Info
from .models import StreamStatus from .types import StreamStatus
class NotificationType(StrEnum): class NotificationType(StrEnum):
-100
View File
@@ -1,100 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Database migration definitions for OwncastSentry."""
from mautrix.util.async_db import Connection, UpgradeTable
upgrade_table = UpgradeTable()
@upgrade_table.register(description="Initial revision") # type: ignore[arg-type, call-arg, untyped-decorator]
async def upgrade_v1(conn: Connection) -> None:
"""Create the initial database schema.
Creates the streams and subscriptions tables.
:param conn: A connection to run the v1 database migration on.
"""
await conn.execute(
"""CREATE TABLE "streams" (
"domain" TEXT NOT NULL UNIQUE,
"name" TEXT,
"title" TEXT,
"last_connect_time" TEXT,
"last_disconnect_time" TEXT,
PRIMARY KEY("domain")
)"""
)
await conn.execute(
"""CREATE TABLE "subscriptions" (
"stream_domain" INTEGER NOT NULL,
"room_id" TEXT NOT NULL,
UNIQUE("room_id","stream_domain")
)"""
)
@upgrade_table.register( # type: ignore[arg-type, call-arg, untyped-decorator]
description="Fix stream_domain column type from INTEGER to TEXT"
)
async def upgrade_v2(conn: Connection) -> None:
"""Upgrade database schema to version 2 format.
Fixes the stream_domain column type in the subscriptions table
from INTEGER to TEXT.
:param conn: A connection to run the v2 database migration on.
"""
# Create new subscriptions table with correct schema
await conn.execute(
"""CREATE TABLE "subscriptions_new" (
"stream_domain" TEXT NOT NULL,
"room_id" TEXT NOT NULL,
UNIQUE("room_id","stream_domain")
)"""
)
# Copy all existing data from old table to new table
await conn.execute(
"""INSERT INTO subscriptions_new (stream_domain, room_id)
SELECT stream_domain, room_id FROM subscriptions"""
)
# Drop the old table and rename new table to original name
await conn.execute("DROP TABLE subscriptions")
await conn.execute("ALTER TABLE subscriptions_new RENAME TO subscriptions")
@upgrade_table.register( # type: ignore[arg-type, call-arg, untyped-decorator]
description="Add failure_counter column for backoff and auto-cleanup"
)
async def upgrade_v3(conn: Connection) -> None:
"""Upgrade database schema to version 3 format.
Adds the failure_counter column to track connection failures
for backoff and auto-cleanup.
:param conn: A connection to run the v3 database migration on.
"""
# Add failure_counter column with default value of 0
await conn.execute(
"""ALTER TABLE streams ADD COLUMN failure_counter INTEGER DEFAULT 0"""
)
def get_upgrade_table() -> UpgradeTable:
"""Return the upgrade table with registered migrations."""
return upgrade_table
-124
View File
@@ -1,124 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Data models for OwncastSentry."""
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
from .utils import (
MAX_INSTANCE_TITLE_LENGTH,
MAX_STREAM_TITLE_LENGTH,
MAX_TAG_LENGTH,
UNKNOWN_STATUS_THRESHOLD,
truncate,
)
class StreamStatus(Enum):
"""Represents the status of a stream."""
ONLINE = "online"
OFFLINE = "offline"
UNKNOWN = "unknown"
@dataclass
class StreamState:
"""Represents the state of an Owncast stream."""
domain: str
name: str | None = None
title: str | None = None
last_connect_time: str | None = None
last_disconnect_time: str | None = None
failure_counter: int = 0
@property
def status(self) -> StreamStatus:
"""Derive stream status from failure count and connect times.
Returns UNKNOWN if failures exceed the threshold, ONLINE if a
connect time is present, or OFFLINE otherwise.
"""
if self.failure_counter > UNKNOWN_STATUS_THRESHOLD:
return StreamStatus.UNKNOWN
if self.last_connect_time is not None:
return StreamStatus.ONLINE
return StreamStatus.OFFLINE
@classmethod
def from_api_response(cls, response: dict[str, Any], domain: str) -> StreamState:
"""Create a StreamState from an API response.
:param response: API response as a dictionary (camelCase keys).
:param domain: The stream domain.
:return: StreamState instance.
"""
return cls(
domain=domain,
title=truncate(response.get("streamTitle", ""), MAX_STREAM_TITLE_LENGTH),
last_connect_time=response.get("lastConnectTime"),
last_disconnect_time=response.get("lastDisconnectTime"),
)
@classmethod
def from_db_row(cls, row: dict[str, Any]) -> StreamState:
"""Create a StreamState from a database row.
:param row: Database row as a dictionary.
:return: StreamState instance.
"""
return cls(
domain=row["domain"],
name=row["name"],
title=row["title"],
last_connect_time=row["last_connect_time"],
last_disconnect_time=row["last_disconnect_time"],
failure_counter=row["failure_counter"],
)
@dataclass
class UpdateResult:
"""Result of a stream update cycle."""
total_streams: int
successful_checks: int
failed_checks: int
@dataclass
class StreamConfig:
"""Represents the configuration of an Owncast stream."""
name: str = ""
tags: list[str] = field(default_factory=list)
@classmethod
def from_api_response(cls, response: dict[str, Any]) -> StreamConfig:
"""Create a StreamConfig from an API response.
:param response: API response as a dictionary.
:return: StreamConfig instance.
"""
# Truncate instance name to max length
name = truncate(response.get("name", ""), MAX_INSTANCE_TITLE_LENGTH)
# Truncate each tag to max length
raw_tags = response.get("tags", [])
tags = [truncate(tag, MAX_TAG_LENGTH) for tag in raw_tags]
return cls(name=name, tags=tags)
+124 -93
View File
@@ -21,18 +21,27 @@ from typing import TYPE_CHECKING, Any
from mautrix.types import MessageType, TextMessageEventContent from mautrix.types import MessageType, TextMessageEventContent
from .metrics import NotificationType from .metrics import NotificationType
from .utils import (
CLEANUP_DELETE_DAYS,
CLEANUP_WARNING_DAYS,
SECONDS_BETWEEN_NOTIFICATIONS,
sanitize_for_plain_text,
)
if TYPE_CHECKING: if TYPE_CHECKING:
import logging import logging
from collections.abc import Sequence
from .database import SubscriptionRepository
from .metrics import MetricsService from .metrics import MetricsService
from .repository import SubscriptionRepository
_SECONDS_BETWEEN_NOTIFICATIONS = 20 * 60
_CLEANUP_WARNING_DAYS = 83
_CLEANUP_DELETE_DAYS = 90
def _sanitize_for_plain_text(text: str) -> str:
"""Sanitize text for plain text rendering."""
if not text:
return text
return " ".join(text.split())
class NotificationService: class NotificationService:
@@ -65,7 +74,7 @@ class NotificationService:
domain: str, domain: str,
name: str, name: str,
title: str, title: str,
tags: list[str], tags: Sequence[str],
*, *,
title_change: bool = False, title_change: bool = False,
) -> None: ) -> None:
@@ -83,28 +92,32 @@ class NotificationService:
time.monotonic() - self.notification_timers_cache[domain] time.monotonic() - self.notification_timers_cache[domain]
) )
self.log.info( self.log.info(
f"[{domain}] Not sending notifications. Only " "[%s] Not sending notifications. Only %s of required "
f"{seconds_since_last} of required " "%s seconds have passed since last notification.",
f"{SECONDS_BETWEEN_NOTIFICATIONS} seconds have " domain,
f"passed since last notification." seconds_since_last,
_SECONDS_BETWEEN_NOTIFICATIONS,
) )
return return
# Record that we're sending a notification now
self._record_notification(domain)
# Build the notification message # Build the notification message
body_text = self._format_message(name, title, domain, tags, title_change) body_text = self._format_message(name, title, domain, tags, title_change)
# Send notifications to all subscribed rooms in parallel # Send notifications to all subscribed rooms in parallel
successful, failed = await self._broadcast_to_rooms(domain, body_text) successful, failed = await self._broadcast_to_rooms(domain, body_text)
# Record that a notification was sent if at least one room received it.
if successful > 0:
self._record_notification(domain)
# Log completion # Log completion
notification_type = "title change" if title_change else "going live" notification_type = "title change" if title_change else "going live"
self.log.info( self.log.info(
f"[{domain}] Completed sending {notification_type} " "[%s] Completed sending %s notifications! %s succeeded, %s failed.",
f"notifications! {successful} succeeded, " domain,
f"{failed} failed." notification_type,
successful,
failed,
) )
self.metrics.record_delivery( self.metrics.record_delivery(
@@ -113,6 +126,75 @@ class NotificationService:
failed=failed, failed=failed,
) )
async def send_cleanup_warning(self, domain: str) -> None:
"""Send cleanup warning notification to all subscribed rooms.
:param domain: The stream domain.
"""
remaining_days = _CLEANUP_DELETE_DAYS - _CLEANUP_WARNING_DAYS
body_text = (
"⚠️ Warning: Subscription Cleanup Scheduled\n\n"
f"The Owncast instance at {domain} has been "
f"unreachable for {_CLEANUP_WARNING_DAYS} days. If it remains "
f"unreachable for {remaining_days} more days "
f"({_CLEANUP_DELETE_DAYS} days total), this subscription "
f"will be automatically removed."
)
successful, failed = await self._broadcast_to_rooms(domain, body_text)
self.log.info(
"[%s] Sent cleanup warning to %s rooms (%s failed).",
domain,
successful,
failed,
)
self.metrics.record_delivery(
NotificationType.CLEANUP_WARNING, successful=successful, failed=failed
)
async def send_cleanup_deletion(self, domain: str) -> None:
"""Send cleanup deletion notification to all subscribed rooms.
:param domain: The stream domain.
"""
body_text = (
"🗑️ Subscription Automatically Removed\n\n"
f"The Owncast instance at {domain} has been "
f"unreachable for {_CLEANUP_DELETE_DAYS} days and has been "
f"automatically removed from subscriptions in this "
f"room.\n\n"
f"If the instance comes online again and you want to "
f"resubscribe, run `!subscribe {domain}`."
)
successful, failed = await self._broadcast_to_rooms(domain, body_text)
self.log.info(
"[%s] Sent cleanup deletion notice to %s rooms (%s failed).",
domain,
successful,
failed,
)
self.metrics.record_delivery(
NotificationType.CLEANUP_DELETION, successful=successful, failed=failed
)
def get_last_notification_time(self, domain: str) -> float:
"""Get the timestamp of the last notification sent for a domain.
:param domain: The stream domain.
:return: Unix timestamp of last notification, or 0 if never notified.
"""
return self.notification_timers_cache.get(domain, 0)
def clear_notification_state(self, domain: str) -> None:
"""Clear cached notification state for a deleted domain.
:param domain: The stream domain to remove from local caches.
"""
self.notification_timers_cache.pop(domain, None)
async def _send_notification( async def _send_notification(
self, room_id: str, body_text: str, domain: str self, room_id: str, body_text: str, domain: str
) -> None: ) -> None:
@@ -128,13 +210,20 @@ class NotificationService:
await self.client.send_message(room_id, content) await self.client.send_message(room_id, content)
except Exception as exception: except Exception as exception:
self.log.warning( self.log.warning(
f"[{domain}] Failed to send notification " "[%s] Failed to send notification message to room [%s]: %s",
f"message to room [{room_id}]: {exception}" domain,
room_id,
exception,
) )
raise raise
def _format_message( def _format_message(
self, name: str, title: str, domain: str, tags: list[str], title_change: bool self,
name: str,
title: str,
domain: str,
tags: Sequence[str],
title_change: bool,
) -> str: ) -> str:
"""Format the notification message body. """Format the notification message body.
@@ -147,7 +236,7 @@ class NotificationService:
""" """
# Use name if available, fallback to domain # Use name if available, fallback to domain
stream_name = name or domain stream_name = name or domain
safe_stream_name = sanitize_for_plain_text(stream_name) safe_stream_name = _sanitize_for_plain_text(stream_name)
# Choose message based on notification type # Choose message based on notification type
if title_change: if title_change:
@@ -157,7 +246,7 @@ class NotificationService:
# Add title if present # Add title if present
if title: if title:
safe_title = sanitize_for_plain_text(title) safe_title = _sanitize_for_plain_text(title)
parts.append(f"\nStream Title: {safe_title}") parts.append(f"\nStream Title: {safe_title}")
# Add stream URL # Add stream URL
@@ -165,39 +254,30 @@ class NotificationService:
# Add tags if present # Add tags if present
if tags: if tags:
safe_tags = [ tag_text = " ".join(
safe_tag f"#{safe_tag}"
for tag in tags for tag in tags
if (safe_tag := sanitize_for_plain_text(tag)) if (safe_tag := _sanitize_for_plain_text(tag))
and not safe_tag.startswith(".") and not safe_tag.startswith(".")
] )
if safe_tags: if tag_text:
parts.append(f"\n\n{' '.join(f'#{tag}' for tag in safe_tags)}") parts.append(f"\n\n{tag_text}")
return "".join(parts) return "".join(parts)
def get_last_notification_time(self, domain: str) -> float:
"""Get the timestamp of the last notification sent for a domain.
:param domain: The stream domain.
:return: Unix timestamp of last notification, or 0 if never notified.
"""
return self.notification_timers_cache.get(domain, 0)
def _can_notify(self, domain: str) -> bool: def _can_notify(self, domain: str) -> bool:
"""Check if enough time has passed to send another notification. """Check if enough time has passed to send another notification.
:param domain: The stream domain. :param domain: The stream domain.
:return: True if notification can be sent, False otherwise. :return: True if notification can be sent, False otherwise.
""" """
if domain not in self.notification_timers_cache: last_notification_time = self.notification_timers_cache.get(domain)
return True return (
last_notification_time is None
seconds_since_last = round( or time.monotonic() - last_notification_time
time.monotonic() - self.notification_timers_cache[domain] >= _SECONDS_BETWEEN_NOTIFICATIONS
) )
return seconds_since_last >= SECONDS_BETWEEN_NOTIFICATIONS
def _record_notification(self, domain: str) -> None: def _record_notification(self, domain: str) -> None:
"""Record that a notification was sent at the current time. """Record that a notification was sent at the current time.
@@ -218,55 +298,6 @@ class NotificationService:
self._send_notification(room_id, body_text, domain) for room_id in room_ids self._send_notification(room_id, body_text, domain) for room_id in room_ids
] ]
results = await asyncio.gather(*tasks, return_exceptions=True) results = await asyncio.gather(*tasks, return_exceptions=True)
failed = sum(1 for r in results if isinstance(r, Exception)) failed = sum(1 for r in results if isinstance(r, BaseException))
successful = len(results) - failed successful = len(results) - failed
return successful, failed return successful, failed
async def send_cleanup_warning(self, domain: str) -> None:
"""Send cleanup warning notification to all subscribed rooms.
:param domain: The stream domain.
"""
remaining_days = CLEANUP_DELETE_DAYS - CLEANUP_WARNING_DAYS
body_text = (
"⚠️ Warning: Subscription Cleanup Scheduled\n\n"
f"The Owncast instance at {domain} has been "
f"unreachable for {CLEANUP_WARNING_DAYS} days. If it remains "
f"unreachable for {remaining_days} more days "
f"({CLEANUP_DELETE_DAYS} days total), this subscription "
f"will be automatically removed."
)
successful, failed = await self._broadcast_to_rooms(domain, body_text)
self.log.info(
f"[{domain}] Sent cleanup warning to {successful} rooms ({failed} failed)."
)
self.metrics.record_delivery(
NotificationType.CLEANUP_WARNING, successful=successful, failed=failed
)
async def send_cleanup_deletion(self, domain: str) -> None:
"""Send cleanup deletion notification to all subscribed rooms.
:param domain: The stream domain.
"""
body_text = (
"🗑️ Subscription Automatically Removed\n\n"
f"The Owncast instance at {domain} has been "
f"unreachable for {CLEANUP_DELETE_DAYS} days and has been "
f"automatically removed from subscriptions in this "
f"room.\n\n"
f"If the instance comes online again and you want to "
f"resubscribe, run `!subscribe {domain}`."
)
successful, failed = await self._broadcast_to_rooms(domain, body_text)
self.log.info(
f"[{domain}] Sent cleanup deletion notice to "
f"{successful} rooms ({failed} failed)."
)
self.metrics.record_delivery(
NotificationType.CLEANUP_DELETION, successful=successful, failed=failed
)
+112 -33
View File
@@ -14,17 +14,12 @@
"""HTTP client for querying Owncast instance APIs.""" """HTTP client for querying Owncast instance APIs."""
import json
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import aiohttp import aiohttp
from .models import StreamConfig, StreamState from .types import InvalidApiResponseError, StreamConfig, StreamState
from .utils import (
OWNCAST_CONFIG_PATH,
OWNCAST_STATUS_PATH,
REQUIRED_STATUS_FIELDS,
user_agent,
)
if TYPE_CHECKING: if TYPE_CHECKING:
import logging import logging
@@ -32,6 +27,48 @@ if TYPE_CHECKING:
from .metrics import MetricsService from .metrics import MetricsService
_OWNCAST_STATUS_PATH = "/api/status"
_OWNCAST_CONFIG_PATH = "/api/config"
_MAX_JSON_RESPONSE_BYTES = 1024 * 1024
_JSON_READ_CHUNK_BYTES = 64 * 1024
_HTTP_CONNECTION_LIMIT = 1000
_HTTP_CONNECTION_LIMIT_PER_HOST = 1
_HTTP_KEEPALIVE_TIMEOUT_SECONDS = 120
_HTTP_CONNECT_TIMEOUT_SECONDS = 5
_HTTP_READ_TIMEOUT_SECONDS = 5
def _user_agent(version: str) -> str:
"""Build the User-Agent header string for HTTP requests."""
return (
f"OwncastSentry/{version}"
" (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)"
)
async def _read_limited_response_body(
response: aiohttp.ClientResponse,
) -> bytearray | None:
"""Read a response body while enforcing the maximum JSON response size."""
# Check Content-Length first when the server provides it so clearly
# oversized responses can be rejected before buffering any body bytes.
if (
response.content_length is not None
and response.content_length > _MAX_JSON_RESPONSE_BYTES
):
return None
body = bytearray()
# Read until EOF instead of using one read(n) call. aiohttp's read(n)
# may return a partial body as soon as data is available.
async for chunk in response.content.iter_chunked(_JSON_READ_CHUNK_BYTES):
body.extend(chunk)
if len(body) > _MAX_JSON_RESPONSE_BYTES:
return None
return body
class OwncastClient: class OwncastClient:
"""HTTP client for communicating with Owncast instances.""" """HTTP client for communicating with Owncast instances."""
@@ -51,15 +88,18 @@ class OwncastClient:
self.metrics = metrics self.metrics = metrics
# Set up HTTP session configuration # Set up HTTP session configuration
headers = {"User-Agent": user_agent(version)} headers = {"User-Agent": _user_agent(version)}
cookie_jar = aiohttp.DummyCookieJar() cookie_jar = aiohttp.DummyCookieJar()
connector = aiohttp.TCPConnector( connector = aiohttp.TCPConnector(
use_dns_cache=False, use_dns_cache=False,
limit=1000, limit=_HTTP_CONNECTION_LIMIT,
limit_per_host=1, limit_per_host=_HTTP_CONNECTION_LIMIT_PER_HOST,
keepalive_timeout=120, keepalive_timeout=_HTTP_KEEPALIVE_TIMEOUT_SECONDS,
)
timeout = aiohttp.ClientTimeout(
sock_connect=_HTTP_CONNECT_TIMEOUT_SECONDS,
sock_read=_HTTP_READ_TIMEOUT_SECONDS,
) )
timeout = aiohttp.ClientTimeout(sock_connect=5, sock_read=5)
self.session = aiohttp.ClientSession( self.session = aiohttp.ClientSession(
headers=headers, headers=headers,
@@ -80,23 +120,48 @@ class OwncastClient:
async with self.session.get(url, allow_redirects=False) as response: async with self.session.get(url, allow_redirects=False) as response:
if response.status != 200: if response.status != 200:
self.log.warning( self.log.warning(
f"[{domain}] Response to request on " "[%s] Response to request on %s was not 200, "
f"{path} was not 200, " "got %s instead.",
f"got {response.status} instead." domain,
path,
response.status,
) )
return None return None
try: try:
result: dict[str, Any] = await response.json() body = await _read_limited_response_body(response)
if body is None:
self.log.warning(
"[%s] Rejecting response to request on %s as it "
"was larger than %s bytes.",
domain,
path,
_MAX_JSON_RESPONSE_BYTES,
)
return None
result = json.loads(body)
if not isinstance(result, dict):
self.log.warning(
"[%s] Rejecting response to request on %s as JSON "
"was not an object.",
domain,
path,
)
return None
return result return result
except (ValueError, aiohttp.ContentTypeError) as e: except ValueError as e:
self.log.warning( self.log.warning(
f"[{domain}] Rejecting response to request on " "[%s] Rejecting response to request on %s as could not "
f"{path} as could not be " "be interpreted as JSON: %s",
f"interpreted as JSON: {e}" domain,
path,
e,
) )
return None return None
except (aiohttp.ClientError, TimeoutError, OSError) as e: except (aiohttp.ClientError, TimeoutError, OSError) as e:
self.log.warning(f"[{domain}] Error making GET request to {path}: {e}") self.log.warning(
"[%s] Error making GET request to %s: %s", domain, path, e
)
return None return None
async def get_stream_state(self, domain: str) -> StreamState | None: async def get_stream_state(self, domain: str) -> StreamState | None:
@@ -108,25 +173,27 @@ class OwncastClient:
:param domain: The domain (not URL) where the stream is hosted. :param domain: The domain (not URL) where the stream is hosted.
:return: A StreamState if available, None on error. :return: A StreamState if available, None on error.
""" """
self.log.debug(f"[{domain}] Fetching current stream state...") self.log.debug("[%s] Fetching current stream state...", domain)
with self.metrics.response_timer(domain) as timer: with self.metrics.response_timer(domain) as timer:
new_state = await self._fetch_json(domain, OWNCAST_STATUS_PATH) new_state = await self._fetch_json(domain, _OWNCAST_STATUS_PATH)
if new_state is None: if new_state is None:
return None return None
# Validate the response contains all basic info needed try:
missing = REQUIRED_STATUS_FIELDS - new_state.keys() stream_state = StreamState.from_api_response(new_state, domain)
if missing: except InvalidApiResponseError as e:
self.log.warning( self.log.warning(
f"[{domain}] Rejecting response to request on " "[%s] Rejecting response to request on %s as response "
f"{OWNCAST_STATUS_PATH} as it is missing " "shape is invalid: %s",
f"fields: {', '.join(sorted(missing))}" domain,
_OWNCAST_STATUS_PATH,
e,
) )
return None return None
timer.success() timer.success()
return StreamState.from_api_response(new_state, domain) return stream_state
async def get_stream_config(self, domain: str) -> StreamConfig | None: async def get_stream_config(self, domain: str) -> StreamConfig | None:
"""Get the current stream config for a given domain. """Get the current stream config for a given domain.
@@ -137,14 +204,26 @@ class OwncastClient:
:param domain: The domain (not URL) where the stream is hosted. :param domain: The domain (not URL) where the stream is hosted.
:return: A StreamConfig, or None if fetch failed. :return: A StreamConfig, or None if fetch failed.
""" """
self.log.debug(f"[{domain}] Fetching current stream config...") self.log.debug("[%s] Fetching current stream config...", domain)
with self.metrics.response_timer(domain) as timer: with self.metrics.response_timer(domain) as timer:
config = await self._fetch_json(domain, OWNCAST_CONFIG_PATH) config = await self._fetch_json(domain, _OWNCAST_CONFIG_PATH)
if config is None: if config is None:
return None return None
try:
stream_config = StreamConfig.from_api_response(config)
except InvalidApiResponseError as e:
self.log.warning(
"[%s] Rejecting response to request on %s as response "
"shape is invalid: %s",
domain,
_OWNCAST_CONFIG_PATH,
e,
)
return None
timer.success() timer.success()
return StreamConfig.from_api_response(config) return stream_config
async def validate_instance(self, domain: str) -> bool: async def validate_instance(self, domain: str) -> bool:
"""Validate that a domain is a valid Owncast instance. """Validate that a domain is a valid Owncast instance.
+370
View File
@@ -0,0 +1,370 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Repository and schema upgrade definitions for OwncastSentry."""
from typing import TYPE_CHECKING, Any
from mautrix.util.async_db import Connection, UpgradeTable
from .types import (
UNKNOWN_STATUS_THRESHOLD,
AlreadySubscribedError,
NotSubscribedError,
RoomSubscription,
StreamState,
)
if TYPE_CHECKING:
from mautrix.util.async_db import Database
upgrade_table = UpgradeTable()
@upgrade_table.register( # type: ignore[arg-type, call-arg, untyped-decorator]
description="Initial revision"
)
async def upgrade_v1(conn: Connection) -> None:
"""Create the initial database schema.
Creates the streams and subscriptions tables.
:param conn: A connection to run the v1 database migration on.
"""
await conn.execute(
"""CREATE TABLE "streams" (
"domain" TEXT NOT NULL UNIQUE,
"name" TEXT,
"title" TEXT,
"last_connect_time" TEXT,
"last_disconnect_time" TEXT,
PRIMARY KEY("domain")
)"""
)
await conn.execute(
"""CREATE TABLE "subscriptions" (
"stream_domain" INTEGER NOT NULL,
"room_id" TEXT NOT NULL,
UNIQUE("room_id","stream_domain")
)"""
)
@upgrade_table.register( # type: ignore[arg-type, call-arg, untyped-decorator]
description="Fix stream_domain column type from INTEGER to TEXT"
)
async def upgrade_v2(conn: Connection) -> None:
"""Upgrade database schema to version 2 format.
Fixes the stream_domain column type in the subscriptions table
from INTEGER to TEXT.
:param conn: A connection to run the v2 database migration on.
"""
await conn.execute(
"""CREATE TABLE "subscriptions_new" (
"stream_domain" TEXT NOT NULL,
"room_id" TEXT NOT NULL,
UNIQUE("room_id","stream_domain")
)"""
)
await conn.execute(
"""INSERT INTO subscriptions_new (stream_domain, room_id)
SELECT stream_domain, room_id FROM subscriptions"""
)
await conn.execute("DROP TABLE subscriptions")
await conn.execute("ALTER TABLE subscriptions_new RENAME TO subscriptions")
@upgrade_table.register( # type: ignore[arg-type, call-arg, untyped-decorator]
description="Add failure_counter column for backoff and auto-cleanup"
)
async def upgrade_v3(conn: Connection) -> None:
"""Upgrade database schema to version 3 format.
Adds the failure_counter column to track connection failures
for backoff and auto-cleanup.
:param conn: A connection to run the v3 database migration on.
"""
await conn.execute(
"""ALTER TABLE streams ADD COLUMN failure_counter INTEGER DEFAULT 0"""
)
def get_upgrade_table() -> UpgradeTable:
"""Return the repository upgrade table with registered migrations."""
return upgrade_table
class StreamRepository:
"""Repository for managing stream data in the database."""
def __init__(self, database: Database) -> None:
"""Initialize the stream repository.
:param database: The maubot database instance.
"""
self.db: Any = database
async def create(self, domain: str) -> bool:
"""Create a new stream entry in the database.
:param domain: The stream domain.
:return: True if created, False if the stream already existed.
"""
query = """INSERT INTO streams (domain)
VALUES ($1)
ON CONFLICT (domain) DO NOTHING"""
async with self.db.acquire() as conn:
result = await conn.execute(query, domain)
return int(result.rowcount) > 0
async def get_by_domain(self, domain: str) -> StreamState | None:
"""Get a stream's state by domain.
:param domain: The stream domain.
:return: StreamState if found, None otherwise.
"""
query = "SELECT * FROM streams WHERE domain=$1"
async with self.db.acquire() as conn:
row = await conn.fetchrow(query, domain)
return StreamState.from_db_row(row) if row else None
async def exists(self, domain: str) -> bool:
"""Check if a stream exists in the database.
:param domain: The stream domain.
:return: True if exists, False otherwise.
"""
result = await self.get_by_domain(domain)
return result is not None
async def update(self, state: StreamState) -> None:
"""Update a stream's state in the database.
This updates display/state fields only. Failure counters are
updated through dedicated methods.
:param state: The StreamState to save.
"""
query = """UPDATE streams
SET name=$1, title=$2, last_connect_time=$3, last_disconnect_time=$4
WHERE domain=$5"""
async with self.db.acquire() as conn:
await conn.execute(
query,
state.name,
state.title,
state.last_connect_time,
state.last_disconnect_time,
state.domain,
)
async def delete(self, domain: str) -> None:
"""Delete a stream record from the database.
:param domain: The stream domain.
"""
query = "DELETE FROM streams WHERE domain=$1"
async with self.db.acquire() as conn:
await conn.execute(query, domain)
async def increment_failure_counter(self, domain: str) -> None:
"""Increment the failure counter for a stream by 1.
:param domain: The stream domain.
"""
query = """UPDATE streams
SET failure_counter = failure_counter + 1
WHERE domain=$1"""
async with self.db.acquire() as conn:
await conn.execute(query, domain)
async def reset_failure_counter(self, domain: str) -> None:
"""Reset the failure counter for a stream to 0.
:param domain: The stream domain.
"""
query = """UPDATE streams
SET failure_counter = 0
WHERE domain=$1 AND failure_counter != 0"""
async with self.db.acquire() as conn:
await conn.execute(query, domain)
class SubscriptionRepository:
"""Repository for managing stream subscriptions in the database."""
def __init__(self, database: Database) -> None:
"""Initialize the subscription repository.
:param database: The maubot database instance.
"""
self.db: Any = database
async def add(self, domain: str, room_id: str) -> None:
"""Add a subscription for a room to a stream.
:param domain: The stream domain.
:param room_id: The Matrix room ID.
:raises AlreadySubscribedError: If subscription already exists.
"""
query = """INSERT INTO subscriptions (stream_domain, room_id)
VALUES ($1, $2)
ON CONFLICT (room_id, stream_domain) DO NOTHING"""
async with self.db.acquire() as conn:
result = await conn.execute(query, domain, room_id)
if int(result.rowcount) == 0:
raise AlreadySubscribedError(domain)
async def remove(self, domain: str, room_id: str) -> None:
"""Remove a subscription for a room from a stream.
:param domain: The stream domain.
:param room_id: The Matrix room ID.
:raises NotSubscribedError: If no subscription exists.
"""
query = "DELETE FROM subscriptions WHERE stream_domain=$1 AND room_id=$2"
async with self.db.acquire() as conn:
result = await conn.execute(query, domain, room_id)
if int(result.rowcount) == 0:
raise NotSubscribedError(domain)
async def delete_all_for_domain(self, domain: str) -> int:
"""Delete all subscriptions for a given stream domain.
:param domain: The stream domain.
:return: Number of subscriptions deleted.
"""
query = "DELETE FROM subscriptions WHERE stream_domain=$1"
async with self.db.acquire() as conn:
result = await conn.execute(query, domain)
return int(result.rowcount)
async def get_subscribed_rooms(self, domain: str) -> list[str]:
"""Get all room IDs subscribed to a stream.
:param domain: The stream domain.
:return: List of room IDs.
"""
query = "SELECT room_id FROM subscriptions WHERE stream_domain=$1"
async with self.db.acquire() as conn:
results = await conn.fetch(query, domain)
return [row["room_id"] for row in results]
async def get_subscribed_streams_for_room(self, room_id: str) -> list[str]:
"""Get all stream domains that a room is subscribed to.
:param room_id: The Matrix room ID.
:return: List of stream domains.
"""
query = "SELECT stream_domain FROM subscriptions WHERE room_id=$1"
async with self.db.acquire() as conn:
results = await conn.fetch(query, room_id)
return [row["stream_domain"] for row in results]
async def has_room_subscriptions(self, room_id: str) -> bool:
"""Check whether a room has any subscriptions."""
query = "SELECT 1 FROM subscriptions WHERE room_id=$1 LIMIT 1"
async with self.db.acquire() as conn:
result = await conn.fetchrow(query, room_id)
return result is not None
async def get_room_subscriptions(self, room_id: str) -> list[RoomSubscription]:
"""Get resolved stream subscriptions for a room ordered by domain.
:param room_id: The Matrix room ID.
:return: Subscriptions with stream state attached.
"""
query = """SELECT streams.*
FROM subscriptions
JOIN streams ON streams.domain = subscriptions.stream_domain
WHERE subscriptions.room_id=$1
ORDER BY streams.domain"""
async with self.db.acquire() as conn:
results = await conn.fetch(query, room_id)
return [
RoomSubscription(
domain=row["domain"],
stream_state=StreamState.from_db_row(row),
)
for row in results
]
async def get_live_room_subscriptions(self, room_id: str) -> list[RoomSubscription]:
"""Get resolved live stream subscriptions for a room ordered by domain."""
query = """SELECT streams.*
FROM subscriptions
JOIN streams ON streams.domain = subscriptions.stream_domain
WHERE subscriptions.room_id=$1
AND streams.last_connect_time IS NOT NULL
AND streams.failure_counter <= $2
ORDER BY streams.domain"""
async with self.db.acquire() as conn:
results = await conn.fetch(query, room_id, UNKNOWN_STATUS_THRESHOLD)
return [
RoomSubscription(
domain=row["domain"],
stream_state=StreamState.from_db_row(row),
)
for row in results
]
async def get_all_subscribed_domains(self) -> list[str]:
"""Get all unique stream domains that have at least one subscription.
:return: List of stream domains.
"""
query = "SELECT DISTINCT stream_domain FROM subscriptions"
async with self.db.acquire() as conn:
results = await conn.fetch(query)
return [row["stream_domain"] for row in results]
async def count_by_domain(self, domain: str) -> int:
"""Count the number of subscriptions for a given stream domain.
:param domain: The stream domain.
:return: Number of subscriptions.
"""
query = "SELECT COUNT(*) FROM subscriptions WHERE stream_domain=$1"
async with self.db.acquire() as conn:
result = await conn.fetchrow(query, domain)
return int(result[0])
async def count_by_domains(self, domains: list[str]) -> dict[str, int]:
"""Count subscriptions for each requested stream domain.
:param domains: The stream domains to count subscriptions for.
:return: Mapping from each requested domain to its subscription count.
"""
if not domains:
return {}
counts = dict.fromkeys(domains, 0)
query = """SELECT stream_domain, COUNT(*) AS subscription_count
FROM subscriptions
GROUP BY stream_domain"""
async with self.db.acquire() as conn:
results = await conn.fetch(query)
for row in results:
domain = row["stream_domain"]
if domain in counts:
counts[domain] = int(row["subscription_count"])
return counts
+75 -45
View File
@@ -18,21 +18,34 @@ import asyncio
import time import time
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .models import StreamState, StreamStatus, UpdateResult from .types import StreamState, StreamStatus, UpdateResult
from .utils import (
CLEANUP_DELETE_THRESHOLD,
CLEANUP_WARNING_THRESHOLD,
TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN,
should_query_stream,
)
if TYPE_CHECKING: if TYPE_CHECKING:
import logging import logging
from .database import StreamRepository, SubscriptionRepository
from .metrics import MetricsService from .metrics import MetricsService
from .notification_service import NotificationService from .notification_service import NotificationService
from .owncast_client import OwncastClient from .owncast_client import OwncastClient
from .repository import StreamRepository, SubscriptionRepository
_TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN = 7 * 60
_CLEANUP_WARNING_THRESHOLD = 83 * 24 * 60
_CLEANUP_DELETE_THRESHOLD = 90 * 24 * 60
def _should_query_stream(failure_counter: int) -> bool:
"""Determine if a stream should be queried based on failure count."""
if failure_counter <= 4:
return True
if failure_counter <= 9:
return failure_counter % 2 == 0
if failure_counter <= 14:
return failure_counter % 3 == 0
if failure_counter <= 29:
return failure_counter % 5 == 0
return failure_counter % 15 == 0
class StreamMonitor: class StreamMonitor:
@@ -92,7 +105,8 @@ class StreamMonitor:
for domain, result in zip(subscribed_domains, results, strict=True): for domain, result in zip(subscribed_domains, results, strict=True):
if isinstance(result, BaseException): if isinstance(result, BaseException):
self.log.exception( self.log.exception(
f"[{domain}] Unhandled exception during stream update.", "[%s] Unhandled exception during stream update.",
domain,
exc_info=result, exc_info=result,
) )
failed_checks += 1 failed_checks += 1
@@ -102,13 +116,19 @@ class StreamMonitor:
failed_checks += 1 failed_checks += 1
self.log.debug( self.log.debug(
f"Update complete. {successful_checks}/{total_streams} succeeded, " "Update complete. %s/%s succeeded, %s failed.",
f"{failed_checks} failed." successful_checks,
total_streams,
failed_checks,
) )
subscription_counts = await self.subscription_repo.count_by_domains(
subscribed_domains
)
for domain in subscribed_domains: for domain in subscribed_domains:
count = await self.subscription_repo.count_by_domain(domain) self.metrics.set_subscription_count(
self.metrics.set_subscription_count(domain, count) domain, subscription_counts.get(domain, 0)
)
return UpdateResult( return UpdateResult(
total_streams=total_streams, total_streams=total_streams,
@@ -134,19 +154,20 @@ class StreamMonitor:
return True return True
# Check if we should query this stream based on backoff schedule # Check if we should query this stream based on backoff schedule
if not should_query_stream(failure_counter): if not _should_query_stream(failure_counter):
# Skip this cycle, increment counter to track time passage # Skip this cycle, increment counter to track time passage
await self.stream_repo.increment_failure_counter(domain) await self.stream_repo.increment_failure_counter(domain)
self.log.debug( self.log.debug(
f"[{domain}] Skipping query due to backoff " "[%s] Skipping query due to backoff (counter=%s)",
f"(counter={failure_counter + 1})" domain,
failure_counter + 1,
) )
# Check cleanup thresholds even when skipping query # Check cleanup thresholds even when skipping query
await self._check_cleanup_thresholds(domain, failure_counter + 1) await self._check_cleanup_thresholds(domain, failure_counter + 1)
updated_state = await self.stream_repo.get_by_domain(domain) updated_state = await self.stream_repo.get_by_domain(domain)
if updated_state is not None: if updated_state is not None:
self.metrics.set_stream_status(domain, updated_state.status) self.metrics.set_stream_status(domain, updated_state.status)
self.metrics.set_check_failures(domain, failure_counter + 1) self.metrics.set_check_failures(domain, failure_counter + 1)
# Backoff is expected behavior, not a failure # Backoff is expected behavior, not a failure
return True return True
@@ -168,14 +189,16 @@ class StreamMonitor:
if new_state is None: if new_state is None:
await self.stream_repo.increment_failure_counter(domain) await self.stream_repo.increment_failure_counter(domain)
self.log.warning( self.log.warning(
f"[{domain}] Connection failure (counter={failure_counter + 1})" "[%s] Connection failure (counter=%s)",
domain,
failure_counter + 1,
) )
# Check cleanup thresholds after connection failure # Check cleanup thresholds after connection failure
await self._check_cleanup_thresholds(domain, failure_counter + 1) await self._check_cleanup_thresholds(domain, failure_counter + 1)
updated_state = await self.stream_repo.get_by_domain(domain) updated_state = await self.stream_repo.get_by_domain(domain)
if updated_state is not None: if updated_state is not None:
self.metrics.set_stream_status(domain, updated_state.status) self.metrics.set_stream_status(domain, updated_state.status)
self.metrics.set_check_failures(domain, failure_counter + 1) self.metrics.set_check_failures(domain, failure_counter + 1)
# Actual connection failure # Actual connection failure
return False return False
@@ -204,7 +227,7 @@ class StreamMonitor:
update_database = True update_database = True
stream_config = await self.owncast_client.get_stream_config(domain) stream_config = await self.owncast_client.get_stream_config(domain)
self.log.info(f"[{domain}] Stream is now live!") self.log.info("[%s] Stream is now live!", domain)
# Calculate seconds since the stream last went offline # Calculate seconds since the stream last went offline
seconds_since_last_offline = round( seconds_since_last_offline = round(
@@ -215,10 +238,13 @@ class StreamMonitor:
if not first_update: if not first_update:
# Use fallback values if config fetch failed # Use fallback values if config fetch failed
stream_name = stream_config.name if stream_config else domain stream_name = stream_config.name if stream_config else domain
stream_tags = stream_config.tags if stream_config else [] stream_tags = stream_config.tags if stream_config else ()
# Has this stream been offline for a short time? # Has this stream been offline for a short time?
if seconds_since_last_offline < TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN: if (
seconds_since_last_offline
< _TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN
):
# Did the stream title change? # Did the stream title change?
if old_state.title != new_state.title: if old_state.title != new_state.title:
# Stream was briefly down; send title # Stream was briefly down; send title
@@ -233,13 +259,12 @@ class StreamMonitor:
else: else:
# Briefly offline, no title change. Skip. # Briefly offline, no title change. Skip.
self.log.info( self.log.info(
f"[{domain}] Not sending " "[%s] Not sending notifications. Stream was only "
f"notifications. Stream was only " "offline for %s of %s seconds and did not change "
f"offline for " "its title.",
f"{seconds_since_last_offline} of " domain,
f"{TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN}" seconds_since_last_offline,
f" seconds and did not change its " _TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN,
f"title."
) )
else: else:
# Offline for a while. Send a normal notification. # Offline for a while. Send a normal notification.
@@ -253,9 +278,9 @@ class StreamMonitor:
else: else:
# No, this is the first time we're querying # No, this is the first time we're querying
self.log.info( self.log.info(
f"[{domain}] Not sending notifications. " "[%s] Not sending notifications. This is the first state "
f"This is the first state update for " "update for this stream.",
f"this stream." domain,
) )
if ( if (
@@ -264,13 +289,13 @@ class StreamMonitor:
): ):
# Did the stream title change mid-session? # Did the stream title change mid-session?
if old_state.title != new_state.title: if old_state.title != new_state.title:
self.log.info(f"[{domain}] Stream title was changed!") self.log.info("[%s] Stream title was changed!", domain)
update_database = True update_database = True
stream_config = await self.owncast_client.get_stream_config(domain) stream_config = await self.owncast_client.get_stream_config(domain)
# Use fallback values if config fetch failed # Use fallback values if config fetch failed
stream_name = stream_config.name if stream_config else domain stream_name = stream_config.name if stream_config else domain
stream_tags = stream_config.tags if stream_config else [] stream_tags = stream_config.tags if stream_config else ()
# Was the last notification sent before the stream # Was the last notification sent before the stream
# last went offline? If so, send a regular go-live # last went offline? If so, send a regular go-live
@@ -303,7 +328,7 @@ class StreamMonitor:
# Yep. This stream is now offline. Log it. # Yep. This stream is now offline. Log it.
update_database = True update_database = True
self.offline_timer_cache[domain] = time.monotonic() self.offline_timer_cache[domain] = time.monotonic()
self.log.info(f"[{domain}] Stream is now offline.") self.log.info("[%s] Stream is now offline.", domain)
# Update the database with current stream state, if needed. # Update the database with current stream state, if needed.
if update_database: if update_database:
@@ -314,7 +339,7 @@ class StreamMonitor:
# Use fallback value if config fetch failed # Use fallback value if config fetch failed
stream_name = stream_config.name if stream_config else "" stream_name = stream_config.name if stream_config else ""
self.log.debug(f"[{domain}] Updating stream state in database...") self.log.debug("[%s] Updating stream state in database...", domain)
# Create updated state object (title already truncated in new_state) # Create updated state object (title already truncated in new_state)
updated_state = StreamState( updated_state = StreamState(
@@ -328,7 +353,7 @@ class StreamMonitor:
await self.stream_repo.update(updated_state) await self.stream_repo.update(updated_state)
# All done. # All done.
self.log.debug(f"[{domain}] State update completed.") self.log.debug("[%s] State update completed.", domain)
if new_state.last_connect_time is not None: if new_state.last_connect_time is not None:
self.metrics.set_stream_status(domain, StreamStatus.ONLINE) self.metrics.set_stream_status(domain, StreamStatus.ONLINE)
else: else:
@@ -342,17 +367,19 @@ class StreamMonitor:
:param counter: The current failure counter value. :param counter: The current failure counter value.
""" """
# Check for 83-day warning threshold # Check for 83-day warning threshold
if counter == CLEANUP_WARNING_THRESHOLD: if counter == _CLEANUP_WARNING_THRESHOLD:
self.log.warning( self.log.warning(
f"[{domain}] Reached 83-day warning threshold. Sending cleanup warning." "[%s] Reached 83-day warning threshold. Sending cleanup warning.",
domain,
) )
await self.notification_service.send_cleanup_warning(domain) await self.notification_service.send_cleanup_warning(domain)
# Check for 90-day deletion threshold # Check for 90-day deletion threshold
if counter >= CLEANUP_DELETE_THRESHOLD: if counter >= _CLEANUP_DELETE_THRESHOLD:
self.log.warning( self.log.warning(
f"[{domain}] Reached 90-day deletion threshold." "[%s] Reached 90-day deletion threshold. "
f" Removing all subscriptions." "Removing all subscriptions.",
domain,
) )
# Send deletion notification # Send deletion notification
await self.notification_service.send_cleanup_deletion(domain) await self.notification_service.send_cleanup_deletion(domain)
@@ -362,10 +389,13 @@ class StreamMonitor:
# Delete the stream record # Delete the stream record
await self.stream_repo.delete(domain) await self.stream_repo.delete(domain)
self.offline_timer_cache.pop(domain, None)
self.notification_service.clear_notification_state(domain)
self.log.info( self.log.info(
f"[{domain}] Cleanup complete. " "[%s] Cleanup complete. Deleted %s subscriptions "
f"Deleted {deleted_count} subscriptions " "and stream record.",
f"and stream record." domain,
deleted_count,
) )
self.metrics.remove_stream(domain) self.metrics.remove_stream(domain)
+123
View File
@@ -0,0 +1,123 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Business logic for managing room stream subscriptions."""
import re
from typing import TYPE_CHECKING
from urllib.parse import urlparse
from .types import (
InvalidOwncastInstanceError,
RoomSubscription,
)
if TYPE_CHECKING:
import logging
from .owncast_client import OwncastClient
from .repository import StreamRepository, SubscriptionRepository
_DOMAIN_CLEANUP_RE = re.compile(r"[^a-z0-9.-]")
def _domainify(url: str) -> str:
"""Extract and sanitize a domain from user input."""
url = url.strip()
if "@" in url:
url = url.rsplit("@", 1)[1]
if not url.startswith(("http://", "https://", "//")):
url = f"//{url}"
parsed = urlparse(url)
domain = (parsed.netloc or parsed.path).lower()
domain = domain.partition(":")[0].partition("/")[0]
return _DOMAIN_CLEANUP_RE.sub("", domain).strip(".-")
class SubscriptionManager:
"""Coordinates subscription use cases between handlers and repositories."""
def __init__(
self,
owncast_client: OwncastClient,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
logger: logging.Logger,
) -> None:
"""Initialize the subscription manager."""
self.owncast_client = owncast_client
self.stream_repo = stream_repo
self.subscription_repo = subscription_repo
self.log = logger
async def subscribe(self, room_id: str, url: str) -> str:
"""Subscribe a room to stream notifications and return the stream domain.
:param room_id: Matrix room ID to subscribe.
:param url: User-supplied Owncast URL, domain, or Fediverse-style address.
:return: Normalized stream domain.
:raises InvalidOwncastInstanceError: If first-time validation fails.
:raises AlreadySubscribedError: If the room is already subscribed.
"""
stream_domain = _domainify(url)
subscription_count = await self.subscription_repo.count_by_domain(stream_domain)
if subscription_count == 0:
is_valid = await self.owncast_client.validate_instance(stream_domain)
if not is_valid:
raise InvalidOwncastInstanceError(stream_domain)
await self.subscription_repo.add(stream_domain, room_id)
if await self.stream_repo.create(stream_domain):
self.log.info("[%s] Discovered new stream!", stream_domain)
self.log.info(
"[%s] Subscription added for room %s.", stream_domain, room_id
)
return stream_domain
async def unsubscribe(self, room_id: str, url: str) -> str:
"""Remove a room subscription and return the stream domain.
:param room_id: Matrix room ID to unsubscribe.
:param url: User-supplied Owncast URL, domain, or Fediverse-style address.
:return: Normalized stream domain.
:raises NotSubscribedError: If no subscription was removed.
"""
stream_domain = _domainify(url)
await self.subscription_repo.remove(stream_domain, room_id)
self.log.info(
"[%s] Subscription removed for room %s.", stream_domain, room_id
)
return stream_domain
async def list_room_subscriptions(self, room_id: str) -> list[RoomSubscription]:
"""Return stream subscriptions for a room with display state attached."""
return await self.subscription_repo.get_room_subscriptions(room_id)
async def has_room_subscriptions(self, room_id: str) -> bool:
"""Return whether the room has any stream subscriptions."""
return await self.subscription_repo.has_room_subscriptions(room_id)
async def list_live_room_subscriptions(
self, room_id: str
) -> list[RoomSubscription]:
"""Return only currently online stream subscriptions for a room."""
return await self.subscription_repo.get_live_room_subscriptions(room_id)
+228
View File
@@ -0,0 +1,228 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Data containers and domain errors for OwncastSentry."""
from dataclasses import dataclass
from enum import Enum
from typing import Any
UNKNOWN_STATUS_THRESHOLD = 15
# Maximum field lengths based on Owncast's configuration
# Source: https://github.com/owncast/owncast/blob/master/
# web/utils/config-constants.tsx
_MAX_INSTANCE_TITLE_LENGTH = 255 # Server Name (line 81)
_MAX_STREAM_TITLE_LENGTH = 100 # Stream Title (line 91)
_MAX_TAG_LENGTH = 24 # Per tag (line 208)
class InvalidApiResponseError(ValueError):
"""The Owncast API response did not match the expected shape."""
def _require_field(response: dict[str, Any], field: str) -> Any:
"""Return a required API response field or raise if missing."""
try:
return response[field]
except KeyError as e:
raise InvalidApiResponseError(f"missing field: {field}") from e
def _require_str(response: dict[str, Any], field: str) -> str:
"""Return a required string API response field."""
value = _require_field(response, field)
if not isinstance(value, str):
raise InvalidApiResponseError(f"{field} must be a string")
return value
def _require_nullable_str(response: dict[str, Any], field: str) -> str | None:
"""Return a required nullable string API response field."""
value = _require_field(response, field)
if value is not None and not isinstance(value, str):
raise InvalidApiResponseError(f"{field} must be a string or null")
return value
def _optional_config_str(response: dict[str, Any], field: str) -> str:
"""Return an optional config string, defaulting to empty when absent."""
value = response.get(field, "")
if not isinstance(value, str):
raise InvalidApiResponseError(f"{field} must be a string")
return value
def _optional_tag_list(response: dict[str, Any]) -> list[str]:
"""Return optional config tags, defaulting to an empty list when absent."""
value = response.get("tags", [])
if not isinstance(value, list):
raise InvalidApiResponseError("tags must be a list")
if not all(isinstance(tag, str) for tag in value):
raise InvalidApiResponseError("tags must contain only strings")
return value
def _truncate(text: str, max_length: int) -> str:
"""Truncate text to a maximum length."""
if len(text) <= max_length:
return text
return text[:max_length]
class StreamStatus(Enum):
"""Represents the status of a stream."""
ONLINE = "online"
OFFLINE = "offline"
UNKNOWN = "unknown"
@dataclass(frozen=True, slots=True)
class StreamState:
"""Represents the state of an Owncast stream."""
domain: str
name: str | None = None
title: str | None = None
last_connect_time: str | None = None
last_disconnect_time: str | None = None
failure_counter: int = 0
@property
def status(self) -> StreamStatus:
"""Derive stream status from failure count and connect times.
Returns UNKNOWN if failures exceed the threshold, ONLINE if a
connect time is present, or OFFLINE otherwise.
"""
if self.failure_counter > UNKNOWN_STATUS_THRESHOLD:
return StreamStatus.UNKNOWN
if self.last_connect_time is not None:
return StreamStatus.ONLINE
return StreamStatus.OFFLINE
@classmethod
def from_api_response(cls, response: dict[str, Any], domain: str) -> StreamState:
"""Create a StreamState from an API response.
:param response: API response as a dictionary (camelCase keys).
:param domain: The stream domain.
:return: StreamState instance.
:raises InvalidApiResponseError: If the response shape is invalid.
"""
stream_title = _require_str(response, "streamTitle")
last_connect_time = _require_nullable_str(response, "lastConnectTime")
last_disconnect_time = _require_nullable_str(response, "lastDisconnectTime")
online = _require_field(response, "online")
if not isinstance(online, bool):
raise InvalidApiResponseError("online must be a boolean")
return cls(
domain=domain,
title=_truncate(stream_title, _MAX_STREAM_TITLE_LENGTH),
last_connect_time=last_connect_time,
last_disconnect_time=last_disconnect_time,
)
@classmethod
def from_db_row(cls, row: dict[str, Any]) -> StreamState:
"""Create a StreamState from a database row.
:param row: Database row as a dictionary.
:return: StreamState instance.
"""
return cls(
domain=row["domain"],
name=row["name"],
title=row["title"],
last_connect_time=row["last_connect_time"],
last_disconnect_time=row["last_disconnect_time"],
failure_counter=row["failure_counter"],
)
@dataclass(frozen=True, slots=True)
class StreamConfig:
"""Represents the configuration of an Owncast stream."""
name: str = ""
tags: tuple[str, ...] = ()
@classmethod
def from_api_response(cls, response: dict[str, Any]) -> StreamConfig:
"""Create a StreamConfig from an API response.
:param response: API response as a dictionary.
:return: StreamConfig instance.
:raises InvalidApiResponseError: If the response shape is invalid.
"""
# Truncate instance name to max length
name = _truncate(
_optional_config_str(response, "name"), _MAX_INSTANCE_TITLE_LENGTH
)
# Truncate each tag to max length
raw_tags = _optional_tag_list(response)
tags = tuple([_truncate(tag, _MAX_TAG_LENGTH) for tag in raw_tags])
return cls(name=name, tags=tags)
@dataclass(frozen=True, slots=True)
class UpdateResult:
"""Result of a stream update cycle."""
total_streams: int
successful_checks: int
failed_checks: int
class SubscriptionError(Exception):
"""Base class for subscription domain errors."""
class InvalidOwncastInstanceError(SubscriptionError):
"""The requested domain is not a reachable Owncast instance."""
def __init__(self, domain: str) -> None:
"""Initialize with the rejected stream domain."""
self.domain = domain
super().__init__(f"invalid Owncast instance: {domain}")
class AlreadySubscribedError(SubscriptionError):
"""The room is already subscribed to the stream."""
def __init__(self, domain: str) -> None:
"""Initialize with the duplicate stream domain."""
self.domain = domain
super().__init__(f"already subscribed: {domain}")
class NotSubscribedError(SubscriptionError):
"""The room is not subscribed to the stream."""
def __init__(self, domain: str) -> None:
"""Initialize with the missing stream domain."""
self.domain = domain
super().__init__(f"not subscribed: {domain}")
@dataclass(frozen=True, slots=True)
class RoomSubscription:
"""A stream subscription resolved with the stream state used for display."""
domain: str
stream_state: StreamState
-205
View File
@@ -1,205 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Utility functions and constants for OwncastSentry."""
import re
from urllib.parse import urlparse
# Path to the GetStatus API call on Owncast instances
OWNCAST_STATUS_PATH = "/api/status"
# Path to GetWebConfig API call on Owncast instances
OWNCAST_CONFIG_PATH = "/api/config"
# Fields that must be present in an Owncast status API response
REQUIRED_STATUS_FIELDS = frozenset(
{
"lastConnectTime",
"lastDisconnectTime",
"streamTitle",
"online",
}
)
def user_agent(version: str) -> str:
"""Build the User-Agent header string for HTTP requests.
:param version: The plugin version string.
:return: A formatted User-Agent string.
"""
return (
f"OwncastSentry/{version}"
" (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)"
)
# Hard minimum amount of time between when notifications can be sent
# for a stream. Prevents spamming notifications for glitchy or
# malicious streams.
SECONDS_BETWEEN_NOTIFICATIONS = 20 * 60 # 20 minutes in seconds
# After a stream goes offline, a timer is started. Then, ...
# - If a stream comes back online with the same title within this
# time, no notification is sent.
# - If a stream comes back online with a different title, a rename
# notification is sent.
# - If this time period passes entirely and a stream comes back
# online after, it's treated as regular going live.
TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN = 7 * 60 # 7 min in seconds
# Auto-cleanup timing (days of continuous unreachability)
CLEANUP_WARNING_DAYS = 83
CLEANUP_DELETE_DAYS = 90
# Counter thresholds derived from days (60-second polling intervals)
CLEANUP_WARNING_THRESHOLD = CLEANUP_WARNING_DAYS * 24 * 60
CLEANUP_DELETE_THRESHOLD = CLEANUP_DELETE_DAYS * 24 * 60
# Failure counter threshold for treating stream status as "unknown"
UNKNOWN_STATUS_THRESHOLD = 15
# Maximum field lengths based on Owncast's configuration
# Source: https://github.com/owncast/owncast/blob/master/
# web/utils/config-constants.tsx
MAX_INSTANCE_TITLE_LENGTH = 255 # Server Name (line 81)
MAX_STREAM_TITLE_LENGTH = 100 # Stream Title (line 91)
MAX_TAG_LENGTH = 24 # Per tag (line 208)
def should_query_stream(failure_counter: int) -> bool:
"""Determine if a stream should be queried based on failure count.
Implements progressive backoff with increasing intervals:
- Counters 0-4: every 60s (first 5 minutes)
- Counters 5-9: every 2min (next 5 minutes)
- Counters 10-14: every 3min (next 5 minutes)
- Counters 15-29: every 5min (next 15 minutes)
- Counters 30+: every 15min
:param failure_counter: The current failure counter value.
:return: True if the stream should be queried this cycle.
"""
if failure_counter <= 4:
# Query every cycle for first 5 minutes (counters 0-4)
return True
if failure_counter <= 9:
# Query every 2nd cycle for next 5 minutes (counters 5-9)
return failure_counter % 2 == 0
if failure_counter <= 14:
# Query every 3rd cycle for next 5 minutes (counters 10-14)
return failure_counter % 3 == 0
if failure_counter <= 29:
# Query every 5th cycle for next 15 minutes (counters 15-29)
return failure_counter % 5 == 0
# Query every 15th cycle after 30 minutes (counter 30+)
return failure_counter % 15 == 0
def domainify(url: str) -> str:
"""Extract and sanitize a domain from user input.
Handles URLs, bare domains, and email-style input (user@domain).
Only allows valid domain characters (alphanumeric, hyphens, periods).
:param url: URL, domain, or email-style string
:return: Sanitized domain
"""
# Handle email-style format first (e.g., "notify@stream.logal.dev")
if "@" in url:
url = url.split("@")[-1]
# Prepend // if no scheme so urlparse treats input as netloc
if not url.startswith(("http://", "https://", "//")):
url = f"//{url}"
parsed = urlparse(url)
domain = (parsed.netloc or parsed.path).lower()
# Strip port and path
domain = domain.split(":")[0].split("/")[0]
# Allow only valid domain characters
return re.sub(r"[^a-z0-9.-]", "", domain).strip(".-")
def truncate(text: str, max_length: int) -> str:
"""Truncate text to a maximum length.
:param text: The text to truncate
:param max_length: Maximum allowed length
:return: Truncated text, or original if within limit
"""
if not text or len(text) <= max_length:
return text
return text[:max_length]
_MARKDOWN_ESCAPE_TABLE = str.maketrans({c: f"\\{c}" for c in r"\*_[]()~`#+-=|{}.!<>&"})
def escape_markdown(text: str) -> str:
"""Escape Markdown special characters to prevent injection attacks.
This function sanitizes untrusted external input (like stream names and titles)
before embedding them in Markdown-formatted messages. It prevents malicious
actors from injecting arbitrary Markdown/HTML content.
:param text: The text to escape
:return: The escaped text safe for Markdown rendering
"""
if not text:
return text
return text.translate(_MARKDOWN_ESCAPE_TABLE)
def sanitize_for_plain_text(text: str) -> str:
"""Sanitize text for plain text rendering.
Remove newlines and normalize whitespace without escaping
special characters. Use this for plain text notifications where
escaping would show literal backslashes.
:param text: The text to sanitize
:return: Sanitized text
"""
if not text:
return text
# Remove newlines and carriage returns to prevent multi-line injection
sanitized = text.replace("\n", " ").replace("\r", " ")
# Collapse multiple spaces into single space
return " ".join(sanitized.split())
def sanitize_for_markdown(text: str) -> str:
"""Sanitize text for safe Markdown rendering.
Remove newlines, normalize whitespace, and escape Markdown special
characters. Use this for any untrusted external content before
embedding in Markdown messages.
Note: This function does not truncate. Size limits should be
enforced at the model layer (e.g., in from_api_response methods).
:param text: The text to sanitize
:return: Sanitized and escaped text safe for Markdown rendering
"""
if not text:
return text
return escape_markdown(sanitize_for_plain_text(text))
+6 -3
View File
@@ -23,15 +23,18 @@ from prometheus_client import generate_latest
from owncastsentry import OwncastSentry from owncastsentry import OwncastSentry
from owncastsentry.config import Config from owncastsentry.config import Config
from owncastsentry.database import StreamRepository, SubscriptionRepository from owncastsentry.repository import (
from owncastsentry.migrations import get_upgrade_table StreamRepository,
SubscriptionRepository,
get_upgrade_table,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from pathlib import Path from pathlib import Path
from owncastsentry.metrics import MetricsService from owncastsentry.metrics import MetricsService
from owncastsentry.models import StreamConfig, StreamState from owncastsentry.types import StreamConfig, StreamState
def generate_metrics_output(metrics: MetricsService) -> str: def generate_metrics_output(metrics: MetricsService) -> str:
+102 -32
View File
@@ -15,28 +15,57 @@
"""Tests for bot command handlers.""" """Tests for bot command handlers."""
import json import json
import logging
from datetime import UTC, datetime, timedelta from datetime import UTC, datetime, timedelta
from unittest.mock import MagicMock
import pytest import pytest
import time_machine import time_machine
from aioresponses import aioresponses from aioresponses import aioresponses
from owncastsentry.commands import CommandHandler from owncastsentry.commands import (
from owncastsentry.models import StreamState _escape_markdown,
from owncastsentry.utils import OWNCAST_STATUS_PATH, UNKNOWN_STATUS_THRESHOLD _format_duration,
_sanitize_for_markdown,
)
from owncastsentry.owncast_client import _OWNCAST_STATUS_PATH
from owncastsentry.types import UNKNOWN_STATUS_THRESHOLD, StreamState
from tests.conftest import VALID_STATUS_RESPONSE from tests.conftest import VALID_STATUS_RESPONSE
def _make_command_handler() -> CommandHandler: class TestEscapeMarkdown:
"""Build a CommandHandler with dummy dependencies for pure logic tests.""" """Markdown special character escaping."""
return CommandHandler(
owncast_client=MagicMock(), @pytest.mark.parametrize(
stream_repo=MagicMock(), ("input_text", "expected"),
subscription_repo=MagicMock(), [
logger=logging.getLogger("test"), pytest.param("hello", "hello", id="plain-text-unchanged"),
pytest.param("*bold*", "\\*bold\\*", id="asterisks"),
pytest.param("_italic_", "\\_italic\\_", id="underscores"),
pytest.param("[link](url)", "\\[link\\]\\(url\\)", id="link-syntax"),
pytest.param("`code`", "\\`code\\`", id="backticks"),
pytest.param("# heading", "\\# heading", id="heading"),
pytest.param("> quote", "\\> quote", id="blockquote"),
pytest.param("<html>", "\\<html\\>", id="angle-brackets"),
pytest.param("a & b", "a \\& b", id="ampersand"),
pytest.param("a\\b", "a\\\\b", id="backslash"),
pytest.param("", "", id="empty-string"),
],
) )
def test_escapes_special_chars(self, input_text: str, expected: str) -> None:
"""Escape the given Markdown special character."""
assert _escape_markdown(input_text) == expected
class TestSanitizeForMarkdown:
"""Markdown sanitization combining newline removal and escaping."""
def test_removes_newlines_and_escapes(self) -> None:
"""Remove newlines and escape Markdown special characters."""
result = _sanitize_for_markdown("*bold*\nnew line")
assert result == "\\*bold\\* new line"
def test_empty_string(self) -> None:
"""Return empty string unchanged."""
assert _sanitize_for_markdown("") == ""
class TestFormatDuration: class TestFormatDuration:
@@ -57,18 +86,24 @@ class TestFormatDuration:
pytest.param(172800, "2 days", id="plural-days"), pytest.param(172800, "2 days", id="plural-days"),
], ],
) )
@time_machine.travel(_NOW)
def test_formats_duration(self, seconds_ago: int, expected: str) -> None: def test_formats_duration(self, seconds_ago: int, expected: str) -> None:
"""Format a timestamp into a human-readable duration.""" """Format a timestamp into a human-readable duration."""
handler = _make_command_handler()
timestamp = (self._NOW - timedelta(seconds=seconds_ago)).isoformat() timestamp = (self._NOW - timedelta(seconds=seconds_ago)).isoformat()
result = handler._format_duration(timestamp) result = _format_duration(timestamp, self._NOW)
assert result == expected assert result == expected
def test_invalid_timestamp(self) -> None: def test_invalid_timestamp(self) -> None:
"""Return 'unknown duration' for unparsable timestamps.""" """Return 'unknown duration' for unparsable timestamps."""
handler = _make_command_handler() assert _format_duration("not-a-timestamp", self._NOW) == "unknown duration"
assert handler._format_duration("not-a-timestamp") == "unknown duration"
def test_naive_timestamp(self) -> None:
"""Return 'unknown duration' for timestamps without timezone information."""
assert _format_duration("2026-03-13T11:59:00", self._NOW) == "unknown duration"
def test_future_timestamp(self) -> None:
"""Return 'unknown duration' for timestamps in the future."""
timestamp = (self._NOW + timedelta(seconds=1)).isoformat()
assert _format_duration(timestamp, self._NOW) == "unknown duration"
class TestSubscribeCommand: class TestSubscribeCommand:
@@ -76,7 +111,7 @@ class TestSubscribeCommand:
async def test_subscribe_valid_stream(self, maubot_test_bot, maubot_plugin) -> None: async def test_subscribe_valid_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Subscribe to a valid Owncast stream.""" """Subscribe to a valid Owncast stream."""
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -94,7 +129,7 @@ class TestSubscribeCommand:
self, maubot_test_bot, maubot_plugin self, maubot_test_bot, maubot_plugin
) -> None: ) -> None:
"""Reject subscription to an invalid Owncast instance.""" """Reject subscription to an invalid Owncast instance."""
status_url = f"https://invalid.com{OWNCAST_STATUS_PATH}" status_url = f"https://invalid.com{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get(status_url, status=404) mocked.get(status_url, status=404)
await maubot_test_bot.send("!subscribe invalid.com") await maubot_test_bot.send("!subscribe invalid.com")
@@ -111,7 +146,7 @@ class TestSubscribeCommand:
self, maubot_test_bot, maubot_plugin self, maubot_test_bot, maubot_plugin
) -> None: ) -> None:
"""Reject duplicate subscription in the same room.""" """Reject duplicate subscription in the same room."""
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -131,7 +166,7 @@ class TestSubscribeCommand:
self, maubot_test_bot, maubot_plugin self, maubot_test_bot, maubot_plugin
) -> None: ) -> None:
"""Skip instance validation when subscribing from a new room.""" """Skip instance validation when subscribing from a new room."""
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -159,7 +194,7 @@ class TestUnsubscribeCommand:
async def test_unsubscribe_existing(self, maubot_test_bot, maubot_plugin) -> None: async def test_unsubscribe_existing(self, maubot_test_bot, maubot_plugin) -> None:
"""Unsubscribe from a subscribed stream.""" """Unsubscribe from a subscribed stream."""
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -205,7 +240,7 @@ class TestSubscriptionsCommand:
async def test_shows_online_stream(self, maubot_test_bot, maubot_plugin) -> None: async def test_shows_online_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show stream details including title and duration.""" """Show stream details including title and duration."""
# Subscribe first # Subscribe first
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -236,11 +271,46 @@ class TestSubscriptionsCommand:
"instances, use `!unsubscribe <domain>`" "instances, use `!unsubscribe <domain>`"
) )
@time_machine.travel(datetime(2026, 3, 13, 12, 0, 0, tzinfo=UTC))
async def test_escapes_markdown_in_stream_name_and_title(
self, maubot_test_bot, maubot_plugin
) -> None:
"""Render stream name and title as literal text in command output."""
status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked:
mocked.get(
status_url,
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
)
await maubot_test_bot.send("!subscribe stream.logal.dev")
await maubot_plugin.stream_repo.update(
StreamState(
domain="stream.logal.dev",
name="*Bold* [link](https://evil.example)\nName",
title="`code` > quote #tag",
last_connect_time="2026-01-01T12:00:00Z",
)
)
await maubot_test_bot.send("!subscriptions")
content = maubot_test_bot.responded[1].content
assert "● ***Bold* [link](https://evil.example) Name**" in content.body
assert " ○ Title: `code` > quote #tag" in content.body
assert content.formatted_body is not None
assert '<a href="https://evil.example">' not in content.formatted_body
assert (
"<strong>*Bold* [link](https://evil.example) Name</strong>"
in content.formatted_body
)
assert "Title: `code` &gt; quote #tag" in content.formatted_body
@time_machine.travel(datetime(2026, 3, 13, 12, 0, 0, tzinfo=UTC)) @time_machine.travel(datetime(2026, 3, 13, 12, 0, 0, tzinfo=UTC))
async def test_shows_offline_stream(self, maubot_test_bot, maubot_plugin) -> None: async def test_shows_offline_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show offline status for non-live streams.""" """Show offline status for non-live streams."""
# Subscribe first # Subscribe first
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -273,7 +343,7 @@ class TestSubscriptionsCommand:
self, maubot_test_bot, maubot_plugin self, maubot_test_bot, maubot_plugin
) -> None: ) -> None:
"""Show offline status without duration before first poll completes.""" """Show offline status without duration before first poll completes."""
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -296,7 +366,7 @@ class TestSubscriptionsCommand:
async def test_shows_unknown_stream(self, maubot_test_bot, maubot_plugin) -> None: async def test_shows_unknown_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show unknown status when instance has been unreachable.""" """Show unknown status when instance has been unreachable."""
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -330,11 +400,11 @@ class TestSubscriptionsCommand:
# Subscribe in reverse alphabetical order to verify sorted output # Subscribe in reverse alphabetical order to verify sorted output
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
f"https://beta.com{OWNCAST_STATUS_PATH}", f"https://beta.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(), body=json.dumps(VALID_STATUS_RESPONSE).encode(),
) )
mocked.get( mocked.get(
f"https://alpha.com{OWNCAST_STATUS_PATH}", f"https://alpha.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(), body=json.dumps(VALID_STATUS_RESPONSE).encode(),
) )
await maubot_test_bot.send("!subscribe beta.com") await maubot_test_bot.send("!subscribe beta.com")
@@ -393,7 +463,7 @@ class TestLiveCommand:
async def test_no_live_streams(self, maubot_test_bot, maubot_plugin) -> None: async def test_no_live_streams(self, maubot_test_bot, maubot_plugin) -> None:
"""Show 'no live' message when all streams are offline.""" """Show 'no live' message when all streams are offline."""
# Subscribe first # Subscribe first
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -423,7 +493,7 @@ class TestLiveCommand:
async def test_shows_live_stream(self, maubot_test_bot, maubot_plugin) -> None: async def test_shows_live_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show live stream with title and duration.""" """Show live stream with title and duration."""
# Subscribe first # Subscribe first
status_url = f"https://stream.logal.dev{OWNCAST_STATUS_PATH}" status_url = f"https://stream.logal.dev{_OWNCAST_STATUS_PATH}"
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
status_url, status_url,
@@ -460,11 +530,11 @@ class TestLiveCommand:
# Subscribe in reverse alphabetical order to verify sorted output # Subscribe in reverse alphabetical order to verify sorted output
with aioresponses() as mocked: with aioresponses() as mocked:
mocked.get( mocked.get(
f"https://beta.com{OWNCAST_STATUS_PATH}", f"https://beta.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(), body=json.dumps(VALID_STATUS_RESPONSE).encode(),
) )
mocked.get( mocked.get(
f"https://alpha.com{OWNCAST_STATUS_PATH}", f"https://alpha.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(), body=json.dumps(VALID_STATUS_RESPONSE).encode(),
) )
await maubot_test_bot.send("!subscribe beta.com") await maubot_test_bot.send("!subscribe beta.com")
-148
View File
@@ -1,148 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for database repository classes."""
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from owncastsentry.database import StreamRepository, SubscriptionRepository
class TestStreamExists:
"""Stream existence checks."""
async def test_returns_true_for_existing_stream(
self, stream_repo: StreamRepository
) -> None:
"""Return True when the stream exists in the database."""
await stream_repo.create("example.com")
assert await stream_repo.exists("example.com") is True
async def test_returns_false_for_missing_stream(
self, stream_repo: StreamRepository
) -> None:
"""Return False when the stream does not exist in the database."""
assert await stream_repo.exists("missing.com") is False
class TestStreamDelete:
"""Stream record deletion."""
async def test_removes_stream_record(self, stream_repo: StreamRepository) -> None:
"""Remove the stream record so get_by_domain returns None."""
await stream_repo.create("example.com")
await stream_repo.delete("example.com")
assert await stream_repo.get_by_domain("example.com") is None
class TestGetSubscribedStreamsForRoom:
"""Subscribed stream lookup by room."""
async def test_returns_all_domains_for_room(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return all domains a room is subscribed to."""
await stream_repo.create("alpha.com")
await stream_repo.create("beta.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("beta.com", "!room1:example.com")
result = await subscription_repo.get_subscribed_streams_for_room(
"!room1:example.com"
)
assert sorted(result) == ["alpha.com", "beta.com"]
async def test_returns_empty_list_for_unsubscribed_room(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return an empty list when the room has no subscriptions."""
result = await subscription_repo.get_subscribed_streams_for_room(
"!nobody:example.com"
)
assert result == []
class TestGetAllSubscribedDomains:
"""Unique subscribed domain retrieval."""
async def test_returns_each_domain_once(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return each domain once even with multiple subscriptions."""
await stream_repo.create("alpha.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
result = await subscription_repo.get_all_subscribed_domains()
assert result == ["alpha.com"]
async def test_returns_empty_list_with_no_subscriptions(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return an empty list when there are no subscriptions."""
result = await subscription_repo.get_all_subscribed_domains()
assert result == []
class TestCountByDomain:
"""Subscription count by domain."""
async def test_returns_correct_count(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return the correct subscription count for a domain."""
await stream_repo.create("alpha.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
assert await subscription_repo.count_by_domain("alpha.com") == 2
async def test_returns_zero_for_unknown_domain(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return 0 for a domain with no subscriptions."""
assert await subscription_repo.count_by_domain("unknown.com") == 0
class TestDeleteAllForDomain:
"""Bulk subscription deletion by domain."""
async def test_deletes_all_subscriptions_and_returns_count(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Delete all subscriptions for the domain and return the count."""
await stream_repo.create("alpha.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
deleted = await subscription_repo.delete_all_for_domain("alpha.com")
assert deleted == 2
rooms = await subscription_repo.get_subscribed_rooms("alpha.com")
assert rooms == []
async def test_returns_zero_for_unknown_domain(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return 0 when deleting subscriptions for an unknown domain."""
assert await subscription_repo.delete_all_for_domain("unknown.com") == 0
+1 -1
View File
@@ -17,7 +17,7 @@
import pytest import pytest
from owncastsentry.metrics import ErrorSource, MetricsService, NotificationType from owncastsentry.metrics import ErrorSource, MetricsService, NotificationType
from owncastsentry.models import StreamStatus from owncastsentry.types import StreamStatus
from tests.conftest import generate_metrics_output from tests.conftest import generate_metrics_output
-193
View File
@@ -1,193 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for data models."""
import pytest
from owncastsentry.models import StreamConfig, StreamState, StreamStatus
from owncastsentry.utils import (
MAX_INSTANCE_TITLE_LENGTH,
MAX_STREAM_TITLE_LENGTH,
MAX_TAG_LENGTH,
UNKNOWN_STATUS_THRESHOLD,
)
class TestStreamStateStatus:
"""Stream status derivation from state fields."""
@pytest.mark.parametrize(
("failure_counter", "last_connect_time", "expected"),
[
pytest.param(
UNKNOWN_STATUS_THRESHOLD + 1,
None,
StreamStatus.UNKNOWN,
id="above-threshold-offline-returns-unknown",
),
pytest.param(
UNKNOWN_STATUS_THRESHOLD + 1,
"2026-01-01T00:00:00Z",
StreamStatus.UNKNOWN,
id="above-threshold-online-returns-unknown",
),
pytest.param(
0,
"2026-01-01T00:00:00Z",
StreamStatus.ONLINE,
id="zero-failures-with-connect-time-returns-online",
),
pytest.param(
0,
None,
StreamStatus.OFFLINE,
id="zero-failures-no-connect-time-returns-offline",
),
pytest.param(
UNKNOWN_STATUS_THRESHOLD,
"2026-01-01T00:00:00Z",
StreamStatus.ONLINE,
id="at-threshold-with-connect-time-returns-online",
),
pytest.param(
UNKNOWN_STATUS_THRESHOLD,
None,
StreamStatus.OFFLINE,
id="at-threshold-no-connect-time-returns-offline",
),
],
)
def test_status(
self,
failure_counter: int,
last_connect_time: str | None,
expected: StreamStatus,
) -> None:
"""Return the correct status based on failure counter and connect time."""
state = StreamState(
domain="example.com",
failure_counter=failure_counter,
last_connect_time=last_connect_time,
)
assert state.status is expected
class TestStreamStateFromApiResponse:
"""StreamState construction from an API response dictionary."""
def test_typical_response(self) -> None:
"""Populate all fields from a complete API response."""
response = {
"streamTitle": "My Stream",
"lastConnectTime": "2026-01-01T00:00:00Z",
"lastDisconnectTime": "2025-12-31T23:00:00Z",
}
state = StreamState.from_api_response(response, "example.com")
assert state.domain == "example.com"
assert state.title == "My Stream"
assert state.last_connect_time == "2026-01-01T00:00:00Z"
assert state.last_disconnect_time == "2025-12-31T23:00:00Z"
assert state.name is None
assert state.failure_counter == 0
def test_empty_response_defaults(self) -> None:
"""Use defaults when optional fields are missing."""
state = StreamState.from_api_response({}, "bare.example.com")
assert state.domain == "bare.example.com"
assert state.title == ""
assert state.last_connect_time is None
assert state.last_disconnect_time is None
def test_title_truncation(self) -> None:
"""Truncate the stream title to MAX_STREAM_TITLE_LENGTH."""
long_title = "A" * (MAX_STREAM_TITLE_LENGTH + 50)
response = {"streamTitle": long_title}
state = StreamState.from_api_response(response, "example.com")
assert len(state.title) == MAX_STREAM_TITLE_LENGTH
assert state.title == "A" * MAX_STREAM_TITLE_LENGTH
class TestStreamStateFromDbRow:
"""StreamState construction from a database row dictionary."""
def test_typical_row(self) -> None:
"""Populate all fields from a complete database row."""
row = {
"domain": "example.com",
"name": "Test Instance",
"title": "Live Now",
"last_connect_time": "2026-01-01T00:00:00Z",
"last_disconnect_time": "2025-12-31T23:00:00Z",
"failure_counter": 3,
}
state = StreamState.from_db_row(row)
assert state.domain == "example.com"
assert state.name == "Test Instance"
assert state.title == "Live Now"
assert state.last_connect_time == "2026-01-01T00:00:00Z"
assert state.last_disconnect_time == "2025-12-31T23:00:00Z"
assert state.failure_counter == 3
def test_row_with_none_optional_fields(self) -> None:
"""Accept None for optional fields in a database row."""
row = {
"domain": "example.com",
"name": None,
"title": None,
"last_connect_time": None,
"last_disconnect_time": None,
"failure_counter": 0,
}
state = StreamState.from_db_row(row)
assert state.domain == "example.com"
assert state.name is None
assert state.title is None
assert state.last_connect_time is None
assert state.last_disconnect_time is None
assert state.failure_counter == 0
class TestStreamConfigFromApiResponse:
"""StreamConfig construction from an API response dictionary."""
def test_typical_response(self) -> None:
"""Populate name and tags from a complete API response."""
response = {"name": "My Instance", "tags": ["gaming", "music"]}
config = StreamConfig.from_api_response(response)
assert config.name == "My Instance"
assert config.tags == ["gaming", "music"]
def test_missing_keys_defaults(self) -> None:
"""Use defaults when name and tags keys are missing."""
config = StreamConfig.from_api_response({})
assert config.name == ""
assert config.tags == []
def test_name_truncation(self) -> None:
"""Truncate the instance name to MAX_INSTANCE_TITLE_LENGTH."""
long_name = "B" * (MAX_INSTANCE_TITLE_LENGTH + 50)
response = {"name": long_name, "tags": []}
config = StreamConfig.from_api_response(response)
assert len(config.name) == MAX_INSTANCE_TITLE_LENGTH
assert config.name == "B" * MAX_INSTANCE_TITLE_LENGTH
def test_tag_truncation(self) -> None:
"""Truncate each tag to MAX_TAG_LENGTH."""
long_tag = "C" * (MAX_TAG_LENGTH + 10)
response = {"name": "", "tags": [long_tag, "short"]}
config = StreamConfig.from_api_response(response)
assert len(config.tags[0]) == MAX_TAG_LENGTH
assert config.tags[0] == "C" * MAX_TAG_LENGTH
assert config.tags[1] == "short"
+126 -4
View File
@@ -14,6 +14,7 @@
"""Tests for the notification service.""" """Tests for the notification service."""
import asyncio
import logging import logging
import time import time
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -21,12 +22,15 @@ from typing import TYPE_CHECKING
import pytest import pytest
from owncastsentry.metrics import MetricsService from owncastsentry.metrics import MetricsService
from owncastsentry.notification_service import NotificationService from owncastsentry.notification_service import (
from owncastsentry.utils import SECONDS_BETWEEN_NOTIFICATIONS _SECONDS_BETWEEN_NOTIFICATIONS,
NotificationService,
_sanitize_for_plain_text,
)
from tests.conftest import _StubMatrixClient, generate_metrics_output from tests.conftest import _StubMatrixClient, generate_metrics_output
if TYPE_CHECKING: if TYPE_CHECKING:
from owncastsentry.database import StreamRepository, SubscriptionRepository from owncastsentry.repository import StreamRepository, SubscriptionRepository
def _make_service( def _make_service(
@@ -44,6 +48,27 @@ def _make_service(
) )
class TestSanitizeForPlainText:
"""Plain text sanitization for notifications."""
@pytest.mark.parametrize(
("input_text", "expected"),
[
pytest.param("hello world", "hello world", id="plain-text"),
pytest.param("line1\nline2", "line1 line2", id="newline-removed"),
pytest.param("line1\rline2", "line1 line2", id="carriage-return"),
pytest.param("line1\r\nline2", "line1 line2", id="crlf-removed"),
pytest.param(
"too many spaces", "too many spaces", id="spaces-collapsed"
),
pytest.param("", "", id="empty-string"),
],
)
def test_sanitizes(self, input_text: str, expected: str) -> None:
"""Sanitize the text for safe plain-text rendering."""
assert _sanitize_for_plain_text(input_text) == expected
class TestCanNotify: class TestCanNotify:
"""Rate-limiting logic for notification cooldowns.""" """Rate-limiting logic for notification cooldowns."""
@@ -75,7 +100,7 @@ class TestCanNotify:
) )
# Subtract an extra second to ensure the cooldown has fully elapsed # Subtract an extra second to ensure the cooldown has fully elapsed
service.notification_timers_cache["example.com"] = ( service.notification_timers_cache["example.com"] = (
time.monotonic() - SECONDS_BETWEEN_NOTIFICATIONS - 1 time.monotonic() - _SECONDS_BETWEEN_NOTIFICATIONS - 1
) )
assert service._can_notify("example.com") is True assert service._can_notify("example.com") is True
@@ -103,6 +128,35 @@ class TestGetLastNotificationTime:
assert service.get_last_notification_time("unknown.com") == 0 assert service.get_last_notification_time("unknown.com") == 0
class TestClearNotificationState:
"""Notification cache cleanup for deleted domains."""
def test_clears_cached_notification_time(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Remove cached notification state for a domain."""
service = _make_service(
client=_StubMatrixClient(), subscription_repo=subscription_repo
)
service.notification_timers_cache["example.com"] = 12345.0
service.clear_notification_state("example.com")
assert service.get_last_notification_time("example.com") == 0
def test_missing_domain_is_ignored(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Ignore cleanup for a domain with no cached state."""
service = _make_service(
client=_StubMatrixClient(), subscription_repo=subscription_repo
)
service.clear_notification_state("unknown.com")
assert service.get_last_notification_time("unknown.com") == 0
class TestFormatMessage: class TestFormatMessage:
"""Notification message formatting.""" """Notification message formatting."""
@@ -246,6 +300,74 @@ class TestNotifyStreamLive:
for msg in client.sent_messages: for msg in client.sent_messages:
assert msg.content.body == expected_body assert msg.content.body == expected_body
async def test_records_cooldown_after_success(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Record a cooldown after at least one room receives a notification."""
client = _StubMatrixClient()
service = _make_service(client=client, subscription_repo=subscription_repo)
await stream_repo.create("example.com")
await subscription_repo.add("example.com", "!room:matrix.org")
before_send = time.monotonic()
await service.notify_stream_live("example.com", "Stream", "Title", [])
assert service.get_last_notification_time("example.com") >= before_send
async def test_no_cooldown_when_no_subscribed_rooms(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Do not record a cooldown if no room receives the notification."""
client = _StubMatrixClient()
service = _make_service(client=client, subscription_repo=subscription_repo)
await service.notify_stream_live("example.com", "Stream", "Title", [])
assert len(client.sent_messages) == 0
assert service.get_last_notification_time("example.com") == 0
async def test_no_cooldown_when_all_deliveries_fail(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Do not record a cooldown if every room delivery fails."""
client = _StubMatrixClient()
client.should_fail_for_rooms.add("!bad:matrix.org")
service = _make_service(client=client, subscription_repo=subscription_repo)
await stream_repo.create("example.com")
await subscription_repo.add("example.com", "!bad:matrix.org")
await service.notify_stream_live("example.com", "Stream", "Title", [])
assert len(client.sent_messages) == 0
assert service.get_last_notification_time("example.com") == 0
async def test_counts_cancelled_delivery_as_failure(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Count a cancelled delivery result as a failure."""
client = _StubMatrixClient()
async def send_message(_room_id: str, _content: object) -> None:
raise asyncio.CancelledError
client.send_message = send_message
service = _make_service(client=client, subscription_repo=subscription_repo)
await stream_repo.create("example.com")
await subscription_repo.add("example.com", "!room:matrix.org")
await service.notify_stream_live("example.com", "Stream", "Title", [])
assert service.get_last_notification_time("example.com") == 0
async def test_skips_when_rate_limited( async def test_skips_when_rate_limited(
self, self,
stream_repo: StreamRepository, stream_repo: StreamRepository,
+231 -27
View File
@@ -22,7 +22,12 @@ import pytest
from aioresponses import aioresponses from aioresponses import aioresponses
from owncastsentry.metrics import MetricsService from owncastsentry.metrics import MetricsService
from owncastsentry.owncast_client import OwncastClient from owncastsentry.owncast_client import (
_MAX_JSON_RESPONSE_BYTES,
OwncastClient,
_read_limited_response_body,
_user_agent,
)
from tests.conftest import ( from tests.conftest import (
VALID_CONFIG_RESPONSE, VALID_CONFIG_RESPONSE,
VALID_STATUS_RESPONSE, VALID_STATUS_RESPONSE,
@@ -33,6 +38,32 @@ if TYPE_CHECKING:
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
class _ChunkedContent:
"""Fake aiohttp response content that yields predefined chunks."""
def __init__(self, chunks: tuple[bytes, ...]) -> None:
"""Store chunks to return from iter_chunked."""
self._chunks = chunks
async def iter_chunked(self, size: int) -> AsyncIterator[bytes]:
"""Yield chunks using the interface aiohttp exposes."""
for chunk in self._chunks:
yield chunk
class _ChunkedResponse:
"""Fake aiohttp response with chunked content."""
def __init__(
self,
chunks: tuple[bytes, ...],
content_length: int | None = None,
) -> None:
"""Store the content stream and optional Content-Length value."""
self.content = _ChunkedContent(chunks)
self.content_length = content_length
@pytest.fixture @pytest.fixture
async def owncast_client() -> AsyncIterator[OwncastClient]: async def owncast_client() -> AsyncIterator[OwncastClient]:
"""Create an OwncastClient and close it after the test.""" """Create an OwncastClient and close it after the test."""
@@ -45,6 +76,79 @@ async def owncast_client() -> AsyncIterator[OwncastClient]:
await client.close() await client.close()
class TestUserAgent:
"""User-Agent header construction."""
@pytest.mark.parametrize(
("version", "expected"),
[
pytest.param(
"1.2.3",
"OwncastSentry/1.2.3 (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)",
id="semver",
),
pytest.param(
"0.0.0",
"OwncastSentry/0.0.0 (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)",
id="zeroed",
),
pytest.param(
"1.1.1.dev10+gf0146d061.d20260313",
"OwncastSentry/1.1.1.dev10+gf0146d061.d20260313 (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)",
id="dev-version",
),
],
)
def test_user_agent(self, version: str, expected: str) -> None:
"""Build a correctly formatted User-Agent header."""
assert _user_agent(version) == expected
class TestReadLimitedResponseBody:
"""Bounded response body reading."""
async def test_reads_all_chunks_before_returning(self) -> None:
"""Return the full body when JSON arrives in multiple chunks."""
response = _ChunkedResponse(
(
b'{"streamTitle":',
b'"hello","online":true,',
b'"lastConnectTime":null,"lastDisconnectTime":null}',
)
)
result = await _read_limited_response_body(response)
assert result == bytearray(
b'{"streamTitle":"hello","online":true,'
b'"lastConnectTime":null,"lastDisconnectTime":null}'
)
async def test_returns_none_when_content_length_is_too_large(self) -> None:
"""Return None when Content-Length is already over the limit."""
response = _ChunkedResponse(
(),
content_length=_MAX_JSON_RESPONSE_BYTES + 1,
)
result = await _read_limited_response_body(response)
assert result is None
async def test_returns_none_when_streamed_body_is_too_large(self) -> None:
"""Return None when chunked content grows past the limit."""
response = _ChunkedResponse(
(
b"x" * _MAX_JSON_RESPONSE_BYTES,
b"x",
)
)
result = await _read_limited_response_body(response)
assert result is None
class TestGetStreamState: class TestGetStreamState:
"""Stream state retrieval from the status API.""" """Stream state retrieval from the status API."""
@@ -82,6 +186,23 @@ class TestGetStreamState:
assert result is None assert result is None
async def test_returns_none_on_invalid_field_type(
self, owncast_client: OwncastClient
) -> None:
"""Return None when the response has malformed field types."""
malformed = {
**VALID_STATUS_RESPONSE,
"streamTitle": 123,
}
with aioresponses() as mocked:
mocked.get(
"https://stream.logal.dev/api/status",
body=json.dumps(malformed).encode(),
)
result = await owncast_client.get_stream_state("stream.logal.dev")
assert result is None
async def test_returns_none_on_invalid_json( async def test_returns_none_on_invalid_json(
self, owncast_client: OwncastClient self, owncast_client: OwncastClient
) -> None: ) -> None:
@@ -95,6 +216,32 @@ class TestGetStreamState:
assert result is None assert result is None
async def test_returns_none_on_non_object_json(
self, owncast_client: OwncastClient
) -> None:
"""Return None when the response JSON is not an object."""
with aioresponses() as mocked:
mocked.get(
"https://stream.logal.dev/api/status",
body=json.dumps([]).encode(),
)
result = await owncast_client.get_stream_state("stream.logal.dev")
assert result is None
async def test_returns_none_on_oversized_json(
self, owncast_client: OwncastClient
) -> None:
"""Return None when the response body is too large."""
with aioresponses() as mocked:
mocked.get(
"https://stream.logal.dev/api/status",
body=b" " * (_MAX_JSON_RESPONSE_BYTES + 1),
)
result = await owncast_client.get_stream_state("stream.logal.dev")
assert result is None
async def test_returns_none_on_non_200(self, owncast_client: OwncastClient) -> None: async def test_returns_none_on_non_200(self, owncast_client: OwncastClient) -> None:
"""Return None when the response status is not 200.""" """Return None when the response status is not 200."""
with aioresponses() as mocked: with aioresponses() as mocked:
@@ -136,7 +283,7 @@ class TestGetStreamConfig:
assert result is not None assert result is not None
assert result.name == "LogalDeveloper's Live Stream" assert result.name == "LogalDeveloper's Live Stream"
assert result.tags == [ assert result.tags == (
"video games", "video games",
"chatting", "chatting",
"casual", "casual",
@@ -144,7 +291,7 @@ class TestGetStreamConfig:
"streaming", "streaming",
"owncast", "owncast",
"variety", "variety",
] )
async def test_returns_none_on_invalid_json( async def test_returns_none_on_invalid_json(
self, owncast_client: OwncastClient self, owncast_client: OwncastClient
@@ -159,6 +306,49 @@ class TestGetStreamConfig:
assert result is None assert result is None
async def test_returns_none_on_non_object_json(
self, owncast_client: OwncastClient
) -> None:
"""Return None when the response JSON is not an object."""
with aioresponses() as mocked:
mocked.get(
"https://stream.logal.dev/api/config",
body=json.dumps([]).encode(),
)
result = await owncast_client.get_stream_config("stream.logal.dev")
assert result is None
async def test_returns_none_on_invalid_field_type(
self, owncast_client: OwncastClient
) -> None:
"""Return None when the response has malformed field types."""
malformed = {
**VALID_CONFIG_RESPONSE,
"tags": "gaming",
}
with aioresponses() as mocked:
mocked.get(
"https://stream.logal.dev/api/config",
body=json.dumps(malformed).encode(),
)
result = await owncast_client.get_stream_config("stream.logal.dev")
assert result is None
async def test_returns_none_on_oversized_json(
self, owncast_client: OwncastClient
) -> None:
"""Return None when the response body is too large."""
with aioresponses() as mocked:
mocked.get(
"https://stream.logal.dev/api/config",
body=b" " * (_MAX_JSON_RESPONSE_BYTES + 1),
)
result = await owncast_client.get_stream_config("stream.logal.dev")
assert result is None
async def test_returns_none_on_non_200(self, owncast_client: OwncastClient) -> None: async def test_returns_none_on_non_200(self, owncast_client: OwncastClient) -> None:
"""Return None when the response status is not 200.""" """Return None when the response status is not 200."""
with aioresponses() as mocked: with aioresponses() as mocked:
@@ -225,15 +415,19 @@ class TestResponseTimeMetrics:
version="0.0.0", version="0.0.0",
metrics=metrics, metrics=metrics,
) )
with aioresponses() as mocked: try:
mocked.get( with aioresponses() as mocked:
"https://example.com/api/status", mocked.get(
body=json.dumps(VALID_STATUS_RESPONSE).encode(), "https://example.com/api/status",
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
)
await client.get_stream_state("example.com")
output = generate_metrics_output(metrics)
assert (
'owncastsentry_api_response_seconds{domain="example.com"}' in output
) )
await client.get_stream_state("example.com") finally:
output = generate_metrics_output(metrics) await client.close()
assert 'owncastsentry_api_response_seconds{domain="example.com"}' in output
await client.close()
async def test_no_observation_on_failure(self) -> None: async def test_no_observation_on_failure(self) -> None:
"""Do not record response time when request fails.""" """Do not record response time when request fails."""
@@ -243,15 +437,20 @@ class TestResponseTimeMetrics:
version="0.0.0", version="0.0.0",
metrics=metrics, metrics=metrics,
) )
with aioresponses() as mocked: try:
mocked.get( with aioresponses() as mocked:
"https://example.com/api/status", mocked.get(
status=500, "https://example.com/api/status",
status=500,
)
await client.get_stream_state("example.com")
output = generate_metrics_output(metrics)
assert (
'owncastsentry_api_response_seconds{domain="example.com"}'
not in output
) )
await client.get_stream_state("example.com") finally:
output = generate_metrics_output(metrics) await client.close()
assert 'owncastsentry_api_response_seconds{domain="example.com"}' not in output
await client.close()
async def test_no_observation_on_connection_error(self) -> None: async def test_no_observation_on_connection_error(self) -> None:
"""Do not record response time on connection error.""" """Do not record response time on connection error."""
@@ -261,15 +460,20 @@ class TestResponseTimeMetrics:
version="0.0.0", version="0.0.0",
metrics=metrics, metrics=metrics,
) )
with aioresponses() as mocked: try:
mocked.get( with aioresponses() as mocked:
"https://example.com/api/status", mocked.get(
exception=ConnectionError(), "https://example.com/api/status",
exception=ConnectionError(),
)
await client.get_stream_state("example.com")
output = generate_metrics_output(metrics)
assert (
'owncastsentry_api_response_seconds{domain="example.com"}'
not in output
) )
await client.get_stream_state("example.com") finally:
output = generate_metrics_output(metrics) await client.close()
assert 'owncastsentry_api_response_seconds{domain="example.com"}' not in output
await client.close()
class TestOpenConnectionCount: class TestOpenConnectionCount:
+370
View File
@@ -0,0 +1,370 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for database repository classes."""
from typing import TYPE_CHECKING
import pytest
from owncastsentry.types import (
UNKNOWN_STATUS_THRESHOLD,
AlreadySubscribedError,
NotSubscribedError,
StreamState,
)
if TYPE_CHECKING:
from owncastsentry.repository import StreamRepository, SubscriptionRepository
class TestStreamExists:
"""Stream existence checks."""
async def test_returns_true_for_existing_stream(
self, stream_repo: StreamRepository
) -> None:
"""Return True when the stream exists in the database."""
await stream_repo.create("example.com")
assert await stream_repo.exists("example.com") is True
async def test_returns_false_for_missing_stream(
self, stream_repo: StreamRepository
) -> None:
"""Return False when the stream does not exist in the database."""
assert await stream_repo.exists("missing.com") is False
class TestStreamCreate:
"""Stream creation behavior."""
async def test_returns_true_when_created(
self, stream_repo: StreamRepository
) -> None:
"""Return True when a stream row is inserted."""
assert await stream_repo.create("example.com") is True
async def test_returns_false_when_existing(
self, stream_repo: StreamRepository
) -> None:
"""Return False when a stream row already exists."""
await stream_repo.create("example.com")
assert await stream_repo.create("example.com") is False
class TestStreamDelete:
"""Stream record deletion."""
async def test_removes_stream_record(self, stream_repo: StreamRepository) -> None:
"""Remove the stream record so get_by_domain returns None."""
await stream_repo.create("example.com")
await stream_repo.delete("example.com")
assert await stream_repo.get_by_domain("example.com") is None
class TestGetSubscribedStreamsForRoom:
"""Subscribed stream lookup by room."""
async def test_returns_all_domains_for_room(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return all domains a room is subscribed to."""
await stream_repo.create("alpha.com")
await stream_repo.create("beta.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("beta.com", "!room1:example.com")
result = await subscription_repo.get_subscribed_streams_for_room(
"!room1:example.com"
)
assert sorted(result) == ["alpha.com", "beta.com"]
async def test_returns_empty_list_for_unsubscribed_room(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return an empty list when the room has no subscriptions."""
result = await subscription_repo.get_subscribed_streams_for_room(
"!nobody:example.com"
)
assert result == []
class TestHasRoomSubscriptions:
"""Room subscription existence checks."""
async def test_returns_true_for_subscribed_room(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return True when the room has at least one subscription."""
await subscription_repo.add("alpha.com", "!room1:example.com")
assert await subscription_repo.has_room_subscriptions("!room1:example.com")
async def test_returns_false_for_unsubscribed_room(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return False when the room has no subscriptions."""
assert not await subscription_repo.has_room_subscriptions(
"!nobody:example.com"
)
class TestGetRoomSubscriptions:
"""Resolved room subscription lookup."""
async def test_returns_sorted_stream_states_and_skips_missing_rows(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return sorted resolved subscriptions and skip missing stream rows."""
await stream_repo.create("beta.example")
await stream_repo.update(StreamState(domain="beta.example", name="Beta"))
await stream_repo.create("alpha.example")
await stream_repo.update(StreamState(domain="alpha.example", name="Alpha"))
await subscription_repo.add("beta.example", "!room:example.com")
await subscription_repo.add("missing.example", "!room:example.com")
await subscription_repo.add("alpha.example", "!room:example.com")
subscriptions = await subscription_repo.get_room_subscriptions(
"!room:example.com"
)
assert [subscription.domain for subscription in subscriptions] == [
"alpha.example",
"beta.example",
]
assert [subscription.stream_state.name for subscription in subscriptions] == [
"Alpha",
"Beta",
]
async def test_returns_empty_list_for_unsubscribed_room(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return an empty list when the room has no resolved subscriptions."""
subscriptions = await subscription_repo.get_room_subscriptions(
"!nobody:example.com"
)
assert subscriptions == []
class TestGetLiveRoomSubscriptions:
"""Resolved live room subscription lookup."""
async def test_returns_online_streams_and_skips_inactive_states(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return only streams whose derived status is ONLINE."""
await stream_repo.create("offline.example")
await stream_repo.update(
StreamState(
domain="offline.example",
name="Offline",
last_disconnect_time="2026-01-01T00:00:00Z",
)
)
await stream_repo.create("online.example")
await stream_repo.update(
StreamState(
domain="online.example",
name="Online",
last_connect_time="2026-01-01T00:00:00Z",
)
)
await stream_repo.create("unknown.example")
await stream_repo.update(
StreamState(
domain="unknown.example",
name="Unknown",
last_connect_time="2026-01-01T00:00:00Z",
)
)
for _ in range(UNKNOWN_STATUS_THRESHOLD + 1):
await stream_repo.increment_failure_counter("unknown.example")
await subscription_repo.add("offline.example", "!room:example.com")
await subscription_repo.add("online.example", "!room:example.com")
await subscription_repo.add("unknown.example", "!room:example.com")
subscriptions = await subscription_repo.get_live_room_subscriptions(
"!room:example.com"
)
assert [subscription.domain for subscription in subscriptions] == [
"online.example"
]
async def test_returns_empty_list_for_room_with_no_live_streams(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return an empty list when no subscribed streams are live."""
await stream_repo.create("offline.example")
await subscription_repo.add("offline.example", "!room:example.com")
subscriptions = await subscription_repo.get_live_room_subscriptions(
"!room:example.com"
)
assert subscriptions == []
class TestAddSubscription:
"""Subscription creation behavior."""
async def test_adds_subscription(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Add a subscription row."""
await subscription_repo.add("alpha.com", "!room1:example.com")
assert await subscription_repo.get_subscribed_rooms("alpha.com") == [
"!room1:example.com"
]
async def test_raises_when_existing(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Raise AlreadySubscribedError when a subscription already exists."""
await subscription_repo.add("alpha.com", "!room1:example.com")
with pytest.raises(AlreadySubscribedError):
await subscription_repo.add("alpha.com", "!room1:example.com")
class TestRemoveSubscription:
"""Subscription removal behavior."""
async def test_removes_subscription(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Remove an existing subscription row."""
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.remove("alpha.com", "!room1:example.com")
assert await subscription_repo.get_subscribed_rooms("alpha.com") == []
async def test_raises_when_missing(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Raise NotSubscribedError when no subscription exists."""
with pytest.raises(NotSubscribedError):
await subscription_repo.remove("alpha.com", "!room1:example.com")
class TestGetAllSubscribedDomains:
"""Unique subscribed domain retrieval."""
async def test_returns_each_domain_once(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return each domain once even with multiple subscriptions."""
await stream_repo.create("alpha.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
result = await subscription_repo.get_all_subscribed_domains()
assert result == ["alpha.com"]
async def test_returns_empty_list_with_no_subscriptions(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return an empty list when there are no subscriptions."""
result = await subscription_repo.get_all_subscribed_domains()
assert result == []
class TestCountByDomain:
"""Subscription count by domain."""
async def test_returns_correct_count(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return the correct subscription count for a domain."""
await stream_repo.create("alpha.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
assert await subscription_repo.count_by_domain("alpha.com") == 2
async def test_returns_zero_for_unknown_domain(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return 0 for a domain with no subscriptions."""
assert await subscription_repo.count_by_domain("unknown.com") == 0
class TestCountByDomains:
"""Bulk subscription counts by domain."""
async def test_returns_counts_for_requested_domains(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return counts for requested domains and zero for missing domains."""
await stream_repo.create("alpha.com")
await stream_repo.create("beta.com")
await stream_repo.create("ignored.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
await subscription_repo.add("beta.com", "!room3:example.com")
await subscription_repo.add("ignored.com", "!room4:example.com")
assert await subscription_repo.count_by_domains(
["beta.com", "missing.com", "alpha.com"]
) == {
"beta.com": 1,
"missing.com": 0,
"alpha.com": 2,
}
async def test_returns_empty_dict_for_empty_domain_list(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return an empty mapping when no domains are requested."""
assert await subscription_repo.count_by_domains([]) == {}
class TestDeleteAllForDomain:
"""Bulk subscription deletion by domain."""
async def test_deletes_all_subscriptions_and_returns_count(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Delete all subscriptions for the domain and return the count."""
await stream_repo.create("alpha.com")
await subscription_repo.add("alpha.com", "!room1:example.com")
await subscription_repo.add("alpha.com", "!room2:example.com")
deleted = await subscription_repo.delete_all_for_domain("alpha.com")
assert deleted == 2
rooms = await subscription_repo.get_subscribed_rooms("alpha.com")
assert rooms == []
async def test_returns_zero_for_unknown_domain(
self, subscription_repo: SubscriptionRepository
) -> None:
"""Return 0 when deleting subscriptions for an unknown domain."""
assert await subscription_repo.delete_all_for_domain("unknown.com") == 0
+99 -21
View File
@@ -18,15 +18,21 @@ import logging
import time import time
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import pytest
from owncastsentry.metrics import MetricsService from owncastsentry.metrics import MetricsService
from owncastsentry.models import StreamConfig, StreamState, StreamStatus from owncastsentry.notification_service import (
from owncastsentry.notification_service import NotificationService _SECONDS_BETWEEN_NOTIFICATIONS,
from owncastsentry.stream_monitor import StreamMonitor NotificationService,
from owncastsentry.utils import (
CLEANUP_DELETE_THRESHOLD,
CLEANUP_WARNING_THRESHOLD,
SECONDS_BETWEEN_NOTIFICATIONS,
) )
from owncastsentry.stream_monitor import (
_CLEANUP_DELETE_THRESHOLD,
_CLEANUP_WARNING_THRESHOLD,
_TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN,
StreamMonitor,
_should_query_stream,
)
from owncastsentry.types import StreamConfig, StreamState, StreamStatus
from tests.conftest import ( from tests.conftest import (
_StubMatrixClient, _StubMatrixClient,
_StubOwncastClient, _StubOwncastClient,
@@ -34,7 +40,7 @@ from tests.conftest import (
) )
if TYPE_CHECKING: if TYPE_CHECKING:
from owncastsentry.database import StreamRepository, SubscriptionRepository from owncastsentry.repository import StreamRepository, SubscriptionRepository
def _make_monitor( def _make_monitor(
@@ -108,6 +114,37 @@ def _make_monitor_with_metrics(
return monitor, notification_service, metrics return monitor, notification_service, metrics
class TestShouldQueryStream:
"""Progressive backoff logic for stream polling."""
@pytest.mark.parametrize(
("counter", "expected"),
[
pytest.param(0, True, id="counter-0-always-query"),
pytest.param(1, True, id="counter-1-always-query"),
pytest.param(4, True, id="counter-4-always-query"),
pytest.param(5, False, id="counter-5-skip-odd"),
pytest.param(6, True, id="counter-6-query-even"),
pytest.param(9, False, id="counter-9-skip-odd"),
pytest.param(10, False, id="counter-10-skip-not-mod-3"),
pytest.param(12, True, id="counter-12-query-mod-3"),
pytest.param(14, False, id="counter-14-skip-not-mod-3"),
pytest.param(15, True, id="counter-15-query-mod-5"),
pytest.param(16, False, id="counter-16-skip-not-mod-5"),
pytest.param(20, True, id="counter-20-query-mod-5"),
pytest.param(29, False, id="counter-29-skip-not-mod-5"),
pytest.param(30, True, id="counter-30-query-mod-15"),
pytest.param(31, False, id="counter-31-skip-not-mod-15"),
pytest.param(45, True, id="counter-45-query-mod-15"),
pytest.param(100, False, id="counter-100-skip-not-mod-15"),
pytest.param(105, True, id="counter-105-query-mod-15"),
],
)
def test_backoff_tiers(self, counter: int, expected: bool) -> None:
"""Return the expected query decision for each backoff tier."""
assert _should_query_stream(counter) == expected
class TestUpdateAllStreams: class TestUpdateAllStreams:
"""Parallel stream update orchestration.""" """Parallel stream update orchestration."""
@@ -252,7 +289,7 @@ class TestUpdateStreamGoesLive:
last_connect_time="2026-01-01T12:00:00Z", last_connect_time="2026-01-01T12:00:00Z",
last_disconnect_time="2026-01-01T10:00:00Z", last_disconnect_time="2026-01-01T10:00:00Z",
), ),
stream_config=StreamConfig(name="Live Stream", tags=["gaming"]), stream_config=StreamConfig(name="Live Stream", tags=("gaming",)),
) )
client = _StubMatrixClient() client = _StubMatrixClient()
monitor, _ = _make_monitor( monitor, _ = _make_monitor(
@@ -271,7 +308,9 @@ class TestUpdateStreamGoesLive:
) )
# Set offline timer to long ago so it's not a brief outage # Set offline timer to long ago so it's not a brief outage
monitor.offline_timer_cache["live.com"] = 0 monitor.offline_timer_cache["live.com"] = (
time.monotonic() - _TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN - 1
)
result = await monitor.update_stream("live.com") result = await monitor.update_stream("live.com")
assert result is True assert result is True
@@ -315,7 +354,9 @@ class TestUpdateStreamGoesLive:
last_disconnect_time="2026-01-01T10:00:00Z", last_disconnect_time="2026-01-01T10:00:00Z",
) )
monitor.offline_timer_cache["live.com"] = 0 monitor.offline_timer_cache["live.com"] = (
time.monotonic() - _TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN - 1
)
result = await monitor.update_stream("live.com") result = await monitor.update_stream("live.com")
assert result is True assert result is True
@@ -447,10 +488,13 @@ class TestUpdateStreamTitleChange:
last_connect_time="2026-01-01T12:00:00Z", last_connect_time="2026-01-01T12:00:00Z",
) )
monitor.offline_timer_cache["title.com"] = 0 now = time.monotonic()
monitor.offline_timer_cache["title.com"] = (
now - _SECONDS_BETWEEN_NOTIFICATIONS - 100
)
# Subtract an extra second to ensure the cooldown has fully elapsed # Subtract an extra second to ensure the cooldown has fully elapsed
notification_service.notification_timers_cache["title.com"] = ( notification_service.notification_timers_cache["title.com"] = (
time.monotonic() - SECONDS_BETWEEN_NOTIFICATIONS - 1 now - _SECONDS_BETWEEN_NOTIFICATIONS - 1
) )
result = await monitor.update_stream("title.com") result = await monitor.update_stream("title.com")
@@ -496,10 +540,13 @@ class TestUpdateStreamTitleChange:
# Last notification was long enough ago to pass rate limiting, # Last notification was long enough ago to pass rate limiting,
# but more recent than the offline timer (so title-change fires) # but more recent than the offline timer (so title-change fires)
monitor.offline_timer_cache["title.com"] = 0 now = time.monotonic()
monitor.offline_timer_cache["title.com"] = (
now - _SECONDS_BETWEEN_NOTIFICATIONS - 100
)
# Subtract an extra second to ensure the cooldown has fully elapsed # Subtract an extra second to ensure the cooldown has fully elapsed
notification_service.notification_timers_cache["title.com"] = ( notification_service.notification_timers_cache["title.com"] = (
time.monotonic() - SECONDS_BETWEEN_NOTIFICATIONS - 1 now - _SECONDS_BETWEEN_NOTIFICATIONS - 1
) )
result = await monitor.update_stream("title.com") result = await monitor.update_stream("title.com")
@@ -546,10 +593,10 @@ class TestUpdateStreamTitleChange:
# and both are old enough to pass rate limiting # and both are old enough to pass rate limiting
now = time.monotonic() now = time.monotonic()
monitor.offline_timer_cache["title.com"] = ( monitor.offline_timer_cache["title.com"] = (
now - SECONDS_BETWEEN_NOTIFICATIONS - 100 now - _SECONDS_BETWEEN_NOTIFICATIONS - 100
) )
notification_service.notification_timers_cache["title.com"] = ( notification_service.notification_timers_cache["title.com"] = (
now - SECONDS_BETWEEN_NOTIFICATIONS - 200 now - _SECONDS_BETWEEN_NOTIFICATIONS - 200
) )
result = await monitor.update_stream("title.com") result = await monitor.update_stream("title.com")
@@ -659,7 +706,7 @@ class TestCheckCleanupThresholds:
await _seed_stream(stream_repo, subscription_repo, domain="warn.com") await _seed_stream(stream_repo, subscription_repo, domain="warn.com")
await monitor._check_cleanup_thresholds("warn.com", CLEANUP_WARNING_THRESHOLD) await monitor._check_cleanup_thresholds("warn.com", _CLEANUP_WARNING_THRESHOLD)
assert len(client.sent_messages) == 1 assert len(client.sent_messages) == 1
assert client.sent_messages[0].content.body == ( assert client.sent_messages[0].content.body == (
@@ -680,7 +727,7 @@ class TestCheckCleanupThresholds:
"""Delete all subscriptions and the stream record at the 90-day threshold.""" """Delete all subscriptions and the stream record at the 90-day threshold."""
owncast = _StubOwncastClient() owncast = _StubOwncastClient()
client = _StubMatrixClient() client = _StubMatrixClient()
monitor, _ = _make_monitor( monitor, notification_service = _make_monitor(
owncast_client=owncast, owncast_client=owncast,
stream_repo=stream_repo, stream_repo=stream_repo,
subscription_repo=subscription_repo, subscription_repo=subscription_repo,
@@ -688,8 +735,10 @@ class TestCheckCleanupThresholds:
) )
await _seed_stream(stream_repo, subscription_repo, domain="delete.com") await _seed_stream(stream_repo, subscription_repo, domain="delete.com")
monitor.offline_timer_cache["delete.com"] = time.monotonic()
notification_service.notification_timers_cache["delete.com"] = time.monotonic()
await monitor._check_cleanup_thresholds("delete.com", CLEANUP_DELETE_THRESHOLD) await monitor._check_cleanup_thresholds("delete.com", _CLEANUP_DELETE_THRESHOLD)
# Deletion notification sent # Deletion notification sent
assert len(client.sent_messages) == 1 assert len(client.sent_messages) == 1
@@ -709,6 +758,8 @@ class TestCheckCleanupThresholds:
assert await stream_repo.get_by_domain("delete.com") is None assert await stream_repo.get_by_domain("delete.com") is None
rooms = await subscription_repo.get_subscribed_rooms("delete.com") rooms = await subscription_repo.get_subscribed_rooms("delete.com")
assert rooms == [] assert rooms == []
assert "delete.com" not in monitor.offline_timer_cache
assert notification_service.get_last_notification_time("delete.com") == 0
async def test_no_action_below_thresholds( async def test_no_action_below_thresholds(
self, self,
@@ -1061,7 +1112,34 @@ class TestStreamMonitorMetrics:
metrics.set_stream_status("delete.com", StreamStatus.OFFLINE) metrics.set_stream_status("delete.com", StreamStatus.OFFLINE)
assert 'domain="delete.com"' in generate_metrics_output(metrics) assert 'domain="delete.com"' in generate_metrics_output(metrics)
await monitor._check_cleanup_thresholds("delete.com", CLEANUP_DELETE_THRESHOLD) await monitor._check_cleanup_thresholds("delete.com", _CLEANUP_DELETE_THRESHOLD)
assert 'domain="delete.com"' not in generate_metrics_output(metrics)
async def test_update_stream_does_not_recreate_metrics_after_cleanup(
self,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Do not recreate per-domain metrics after update_stream deletes a stream."""
owncast = _StubOwncastClient(stream_state=None)
client = _StubMatrixClient()
monitor, _, metrics = _make_monitor_with_metrics(
owncast_client=owncast,
stream_repo=stream_repo,
subscription_repo=subscription_repo,
client=client,
)
await _seed_stream(stream_repo, subscription_repo, domain="delete.com")
metrics.set_stream_status("delete.com", StreamStatus.OFFLINE)
metrics.set_check_failures("delete.com", _CLEANUP_DELETE_THRESHOLD - 1)
for _ in range(_CLEANUP_DELETE_THRESHOLD - 1):
await stream_repo.increment_failure_counter("delete.com")
await monitor.update_stream("delete.com")
assert await stream_repo.get_by_domain("delete.com") is None
assert await subscription_repo.get_subscribed_rooms("delete.com") == []
assert 'domain="delete.com"' not in generate_metrics_output(metrics) assert 'domain="delete.com"' not in generate_metrics_output(metrics)
async def test_records_subscription_counts( async def test_records_subscription_counts(
+289
View File
@@ -0,0 +1,289 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for subscription manager business logic."""
import logging
from typing import TYPE_CHECKING
import pytest
from owncastsentry.subscription_manager import SubscriptionManager, _domainify
from owncastsentry.types import (
UNKNOWN_STATUS_THRESHOLD,
AlreadySubscribedError,
InvalidOwncastInstanceError,
NotSubscribedError,
StreamState,
)
if TYPE_CHECKING:
from owncastsentry.repository import StreamRepository, SubscriptionRepository
class TestDomainify:
"""Domain extraction and sanitization from user input."""
@pytest.mark.parametrize(
("input_url", "expected"),
[
pytest.param("example.com", "example.com", id="bare-domain"),
pytest.param(" example.com ", "example.com", id="surrounding-whitespace"),
pytest.param("https://example.com", "example.com", id="https-url"),
pytest.param("http://example.com", "example.com", id="http-url"),
pytest.param("https://example.com:8080", "example.com", id="url-with-port"),
pytest.param(
"https://example.com/path/to/page",
"example.com",
id="url-with-path",
),
pytest.param(
"user@stream.logal.dev",
"stream.logal.dev",
id="email-style",
),
pytest.param(
"matrix@notify@stream.logal.dev",
"stream.logal.dev",
id="last-at-sign-wins",
),
pytest.param("EXAMPLE.COM", "example.com", id="uppercase"),
pytest.param("exam!ple.com", "example.com", id="special-chars-stripped"),
pytest.param(".example.com.", "example.com", id="leading-trailing-dots"),
pytest.param("-example.com-", "example.com", id="leading-trailing-hyphens"),
pytest.param(
"sub.domain.example.com",
"sub.domain.example.com",
id="subdomain",
),
],
)
def test_extracts_domain(self, input_url: str, expected: str) -> None:
"""Extract and sanitize the domain from various input formats."""
assert _domainify(input_url) == expected
class _StubOwncastClient:
"""Owncast client stub for validation-only manager tests."""
def __init__(self, *, valid: bool = True) -> None:
"""Initialize the stub with a fixed validation result."""
self.valid = valid
self.validated_domains: list[str] = []
async def validate_instance(self, domain: str) -> bool:
"""Record the domain and return the configured validation result."""
self.validated_domains.append(domain)
return self.valid
@pytest.fixture
def owncast_client() -> _StubOwncastClient:
"""Return a validation-only Owncast client stub."""
return _StubOwncastClient()
@pytest.fixture
def manager(
owncast_client: _StubOwncastClient,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> SubscriptionManager:
"""SubscriptionManager built directly for unit tests."""
return SubscriptionManager(
owncast_client=owncast_client, # type: ignore[arg-type]
stream_repo=stream_repo,
subscription_repo=subscription_repo,
logger=logging.getLogger("test"),
)
class TestManagerSubscribe:
"""SubscriptionManager subscribe workflow."""
async def test_first_subscription_validates_and_creates_stream(
self,
manager: SubscriptionManager,
owncast_client: _StubOwncastClient,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""First subscription validates the instance and creates stream state."""
domain = await manager.subscribe(
"!room:example.com", "https://Stream.Example/foo"
)
assert domain == "stream.example"
assert owncast_client.validated_domains == ["stream.example"]
assert await stream_repo.exists("stream.example") is True
assert await subscription_repo.get_subscribed_rooms("stream.example") == [
"!room:example.com"
]
async def test_invalid_first_subscription_raises(
self,
manager: SubscriptionManager,
owncast_client: _StubOwncastClient,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Invalid first-time Owncast validation raises a domain error."""
owncast_client.valid = False
with pytest.raises(InvalidOwncastInstanceError) as exc_info:
await manager.subscribe("!room:example.com", "bad.example")
assert exc_info.value.domain == "bad.example"
assert owncast_client.validated_domains == ["bad.example"]
assert await stream_repo.exists("bad.example") is False
assert await subscription_repo.get_subscribed_rooms("bad.example") == []
async def test_duplicate_subscription_raises(
self,
manager: SubscriptionManager,
owncast_client: _StubOwncastClient,
subscription_repo: SubscriptionRepository,
) -> None:
"""Duplicate room subscription raises AlreadySubscribedError."""
await manager.subscribe("!room:example.com", "stream.example")
with pytest.raises(AlreadySubscribedError) as exc_info:
await manager.subscribe("!room:example.com", "stream.example")
assert exc_info.value.domain == "stream.example"
assert owncast_client.validated_domains == ["stream.example"]
assert await subscription_repo.get_subscribed_rooms("stream.example") == [
"!room:example.com"
]
async def test_existing_stream_new_room_skips_validation(
self,
manager: SubscriptionManager,
owncast_client: _StubOwncastClient,
subscription_repo: SubscriptionRepository,
) -> None:
"""Known domains skip remote validation for additional room subscriptions."""
await manager.subscribe("!room1:example.com", "stream.example")
owncast_client.valid = False
domain = await manager.subscribe("!room2:example.com", "stream.example")
assert domain == "stream.example"
assert owncast_client.validated_domains == ["stream.example"]
rooms = await subscription_repo.get_subscribed_rooms("stream.example")
assert sorted(rooms) == [
"!room1:example.com",
"!room2:example.com",
]
class TestManagerUnsubscribe:
"""SubscriptionManager unsubscribe workflow."""
async def test_removes_existing_subscription(
self,
manager: SubscriptionManager,
subscription_repo: SubscriptionRepository,
) -> None:
"""Existing room subscription is removed and its domain is returned."""
await manager.subscribe("!room:example.com", "stream.example")
domain = await manager.unsubscribe("!room:example.com", "stream.example")
assert domain == "stream.example"
assert await subscription_repo.get_subscribed_rooms("stream.example") == []
async def test_missing_subscription_raises(
self,
manager: SubscriptionManager,
) -> None:
"""Removing a non-existent subscription raises NotSubscribedError."""
with pytest.raises(NotSubscribedError) as exc_info:
await manager.unsubscribe("!room:example.com", "missing.example")
assert exc_info.value.domain == "missing.example"
class TestManagerListings:
"""SubscriptionManager room listing behavior."""
async def test_list_room_subscriptions_returns_sorted_stream_states(
self,
manager: SubscriptionManager,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Return sorted room subscriptions and skip missing stream rows."""
await stream_repo.create("beta.example")
await stream_repo.update(StreamState(domain="beta.example", name="Beta"))
await stream_repo.create("alpha.example")
await stream_repo.update(StreamState(domain="alpha.example", name="Alpha"))
await subscription_repo.add("beta.example", "!room:example.com")
await subscription_repo.add("missing.example", "!room:example.com")
await subscription_repo.add("alpha.example", "!room:example.com")
subscriptions = await manager.list_room_subscriptions("!room:example.com")
assert [subscription.domain for subscription in subscriptions] == [
"alpha.example",
"beta.example",
]
assert [subscription.stream_state.name for subscription in subscriptions] == [
"Alpha",
"Beta",
]
async def test_list_live_room_subscriptions_filters_online_streams(
self,
manager: SubscriptionManager,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
) -> None:
"""Live listing includes only subscriptions with ONLINE stream status."""
await stream_repo.create("offline.example")
await stream_repo.update(
StreamState(
domain="offline.example",
name="Offline",
last_disconnect_time="2026-01-01T00:00:00Z",
)
)
await stream_repo.create("online.example")
await stream_repo.update(
StreamState(
domain="online.example",
name="Online",
last_connect_time="2026-01-01T00:00:00Z",
)
)
await stream_repo.create("unknown.example")
await stream_repo.update(
StreamState(
domain="unknown.example",
name="Unknown",
last_connect_time="2026-01-01T00:00:00Z",
)
)
for _ in range(UNKNOWN_STATUS_THRESHOLD + 1):
await stream_repo.increment_failure_counter("unknown.example")
await subscription_repo.add("offline.example", "!room:example.com")
await subscription_repo.add("online.example", "!room:example.com")
await subscription_repo.add("unknown.example", "!room:example.com")
subscriptions = await manager.list_live_room_subscriptions("!room:example.com")
assert [subscription.domain for subscription in subscriptions] == [
"online.example"
]
+352
View File
@@ -0,0 +1,352 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for data models."""
from dataclasses import FrozenInstanceError
import pytest
from owncastsentry.types import (
_MAX_INSTANCE_TITLE_LENGTH,
_MAX_STREAM_TITLE_LENGTH,
_MAX_TAG_LENGTH,
UNKNOWN_STATUS_THRESHOLD,
AlreadySubscribedError,
InvalidApiResponseError,
InvalidOwncastInstanceError,
NotSubscribedError,
RoomSubscription,
StreamConfig,
StreamState,
StreamStatus,
SubscriptionError,
UpdateResult,
_truncate,
)
class TestTruncate:
"""Text truncation to a maximum length."""
@pytest.mark.parametrize(
("text", "max_length", "expected"),
[
pytest.param("hello", 10, "hello", id="under-limit"),
pytest.param("hello", 5, "hello", id="exact-limit"),
pytest.param("hello world", 5, "hello", id="over-limit"),
pytest.param("", 5, "", id="empty-string"),
],
)
def test_truncates(self, text: str, max_length: int, expected: str) -> None:
"""Truncate text that exceeds the maximum length."""
assert _truncate(text, max_length) == expected
class TestStreamStateStatus:
"""Stream status derivation from state fields."""
@pytest.mark.parametrize(
("failure_counter", "last_connect_time", "expected"),
[
pytest.param(
UNKNOWN_STATUS_THRESHOLD + 1,
None,
StreamStatus.UNKNOWN,
id="above-threshold-offline-returns-unknown",
),
pytest.param(
UNKNOWN_STATUS_THRESHOLD + 1,
"2026-01-01T00:00:00Z",
StreamStatus.UNKNOWN,
id="above-threshold-online-returns-unknown",
),
pytest.param(
0,
"2026-01-01T00:00:00Z",
StreamStatus.ONLINE,
id="zero-failures-with-connect-time-returns-online",
),
pytest.param(
0,
None,
StreamStatus.OFFLINE,
id="zero-failures-no-connect-time-returns-offline",
),
pytest.param(
UNKNOWN_STATUS_THRESHOLD,
"2026-01-01T00:00:00Z",
StreamStatus.ONLINE,
id="at-threshold-with-connect-time-returns-online",
),
pytest.param(
UNKNOWN_STATUS_THRESHOLD,
None,
StreamStatus.OFFLINE,
id="at-threshold-no-connect-time-returns-offline",
),
],
)
def test_status(
self,
failure_counter: int,
last_connect_time: str | None,
expected: StreamStatus,
) -> None:
"""Return the correct status based on failure counter and connect time."""
state = StreamState(
domain="example.com",
failure_counter=failure_counter,
last_connect_time=last_connect_time,
)
assert state.status is expected
class TestStreamStateFromApiResponse:
"""StreamState construction from an API response dictionary."""
def test_typical_response(self) -> None:
"""Populate all fields from a complete API response."""
response = {
"streamTitle": "My Stream",
"lastConnectTime": "2026-01-01T00:00:00Z",
"lastDisconnectTime": "2025-12-31T23:00:00Z",
"online": True,
}
state = StreamState.from_api_response(response, "example.com")
assert state.domain == "example.com"
assert state.title == "My Stream"
assert state.last_connect_time == "2026-01-01T00:00:00Z"
assert state.last_disconnect_time == "2025-12-31T23:00:00Z"
assert state.name is None
assert state.failure_counter == 0
def test_missing_required_field_raises(self) -> None:
"""Reject API responses without required stream state fields."""
with pytest.raises(InvalidApiResponseError):
StreamState.from_api_response({}, "bare.example.com")
def test_nullable_timestamp_fields(self) -> None:
"""Accept null values for Owncast timestamp fields."""
response = {
"streamTitle": "Offline Stream",
"lastConnectTime": None,
"lastDisconnectTime": None,
"online": False,
}
state = StreamState.from_api_response(response, "example.com")
assert state.last_connect_time is None
assert state.last_disconnect_time is None
def test_title_truncation(self) -> None:
"""Truncate the stream title to _MAX_STREAM_TITLE_LENGTH."""
long_title = "A" * (_MAX_STREAM_TITLE_LENGTH + 50)
response = {
"streamTitle": long_title,
"lastConnectTime": None,
"lastDisconnectTime": None,
"online": True,
}
state = StreamState.from_api_response(response, "example.com")
assert len(state.title) == _MAX_STREAM_TITLE_LENGTH
assert state.title == "A" * _MAX_STREAM_TITLE_LENGTH
@pytest.mark.parametrize(
("field", "value"),
[
pytest.param("streamTitle", 123, id="title-not-string"),
pytest.param("lastConnectTime", [], id="connect-time-not-string-or-null"),
pytest.param(
"lastDisconnectTime",
{},
id="disconnect-time-not-string-or-null",
),
pytest.param("online", "true", id="online-not-bool"),
],
)
def test_invalid_field_type_raises(self, field: str, value: object) -> None:
"""Reject stream state responses with malformed field types."""
response: dict[str, object] = {
"streamTitle": "My Stream",
"lastConnectTime": None,
"lastDisconnectTime": None,
"online": True,
}
response[field] = value
with pytest.raises(InvalidApiResponseError):
StreamState.from_api_response(response, "example.com")
class TestStreamStateFromDbRow:
"""StreamState construction from a database row dictionary."""
def test_typical_row(self) -> None:
"""Populate all fields from a complete database row."""
row = {
"domain": "example.com",
"name": "Test Instance",
"title": "Live Now",
"last_connect_time": "2026-01-01T00:00:00Z",
"last_disconnect_time": "2025-12-31T23:00:00Z",
"failure_counter": 3,
}
state = StreamState.from_db_row(row)
assert state.domain == "example.com"
assert state.name == "Test Instance"
assert state.title == "Live Now"
assert state.last_connect_time == "2026-01-01T00:00:00Z"
assert state.last_disconnect_time == "2025-12-31T23:00:00Z"
assert state.failure_counter == 3
def test_row_with_none_optional_fields(self) -> None:
"""Accept None for optional fields in a database row."""
row = {
"domain": "example.com",
"name": None,
"title": None,
"last_connect_time": None,
"last_disconnect_time": None,
"failure_counter": 0,
}
state = StreamState.from_db_row(row)
assert state.domain == "example.com"
assert state.name is None
assert state.title is None
assert state.last_connect_time is None
assert state.last_disconnect_time is None
assert state.failure_counter == 0
class TestStreamConfigFromApiResponse:
"""StreamConfig construction from an API response dictionary."""
def test_typical_response(self) -> None:
"""Populate name and tags from a complete API response."""
response = {"name": "My Instance", "tags": ["gaming", "music"]}
config = StreamConfig.from_api_response(response)
assert config.name == "My Instance"
assert config.tags == ("gaming", "music")
def test_missing_keys_defaults(self) -> None:
"""Use defaults when name and tags keys are missing."""
config = StreamConfig.from_api_response({})
assert config.name == ""
assert config.tags == ()
def test_name_truncation(self) -> None:
"""Truncate the instance name to _MAX_INSTANCE_TITLE_LENGTH."""
long_name = "B" * (_MAX_INSTANCE_TITLE_LENGTH + 50)
response = {"name": long_name, "tags": []}
config = StreamConfig.from_api_response(response)
assert len(config.name) == _MAX_INSTANCE_TITLE_LENGTH
assert config.name == "B" * _MAX_INSTANCE_TITLE_LENGTH
def test_tag_truncation(self) -> None:
"""Truncate each tag to _MAX_TAG_LENGTH."""
long_tag = "C" * (_MAX_TAG_LENGTH + 10)
response = {"name": "", "tags": [long_tag, "short"]}
config = StreamConfig.from_api_response(response)
assert len(config.tags[0]) == _MAX_TAG_LENGTH
assert config.tags[0] == "C" * _MAX_TAG_LENGTH
assert config.tags[1] == "short"
@pytest.mark.parametrize(
("field", "value"),
[
pytest.param("name", None, id="name-not-string"),
pytest.param("tags", "gaming", id="tags-not-list"),
pytest.param("tags", ["gaming", 123], id="tag-not-string"),
],
)
def test_invalid_field_type_raises(self, field: str, value: object) -> None:
"""Reject stream config responses with malformed field types."""
response: dict[str, object] = {"name": "My Instance", "tags": ["gaming"]}
response[field] = value
with pytest.raises(InvalidApiResponseError):
StreamConfig.from_api_response(response)
class TestValueTypeImmutability:
"""Dataclass value containers are immutable snapshots."""
def test_stream_state_is_immutable(self) -> None:
"""StreamState cannot be mutated in place."""
state = StreamState(domain="stream.example")
with pytest.raises(FrozenInstanceError):
state.title = "Changed" # type: ignore[misc]
def test_stream_config_is_immutable(self) -> None:
"""StreamConfig cannot be mutated in place."""
config = StreamConfig(name="Stream")
with pytest.raises(FrozenInstanceError):
config.name = "Changed" # type: ignore[misc]
def test_stream_config_tags_are_immutable(self) -> None:
"""StreamConfig tags are stored in an immutable tuple."""
config = StreamConfig(name="Stream", tags=("gaming",))
assert config.tags == ("gaming",)
def test_update_result_is_immutable(self) -> None:
"""UpdateResult cannot be mutated in place."""
result = UpdateResult(total_streams=1, successful_checks=1, failed_checks=0)
with pytest.raises(FrozenInstanceError):
result.failed_checks = 1 # type: ignore[misc]
class TestSubscriptionTypes:
"""Subscription display containers and domain error hierarchy."""
@pytest.mark.parametrize(
"error_cls",
[
pytest.param(InvalidOwncastInstanceError, id="invalid-instance"),
pytest.param(AlreadySubscribedError, id="already-subscribed"),
pytest.param(NotSubscribedError, id="not-subscribed"),
],
)
def test_errors_subclass_subscription_error(
self, error_cls: type[Exception]
) -> None:
"""Every subscription domain error subclasses SubscriptionError."""
assert issubclass(error_cls, SubscriptionError)
@pytest.mark.parametrize(
"error",
[
pytest.param(
InvalidOwncastInstanceError("bad.example"),
id="invalid-instance",
),
pytest.param(
AlreadySubscribedError("dupe.example"),
id="already-subscribed",
),
pytest.param(
NotSubscribedError("missing.example"),
id="not-subscribed",
),
],
)
def test_errors_store_domain(self, error: SubscriptionError) -> None:
"""Subscription domain errors expose the stream domain that failed."""
assert error.domain in str(error)
def test_room_subscription_is_immutable(self) -> None:
"""RoomSubscription is an immutable stream display snapshot."""
state = StreamState(domain="stream.example")
subscription = RoomSubscription(domain="stream.example", stream_state=state)
with pytest.raises(FrozenInstanceError):
subscription.domain = "other.example" # type: ignore[misc]
-197
View File
@@ -1,197 +0,0 @@
# Copyright 2026 Logan Fick
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for utility functions and constants."""
import pytest
from owncastsentry.utils import (
domainify,
escape_markdown,
sanitize_for_markdown,
sanitize_for_plain_text,
should_query_stream,
truncate,
user_agent,
)
class TestUserAgent:
"""User-Agent header construction."""
@pytest.mark.parametrize(
("version", "expected"),
[
pytest.param(
"1.2.3",
"OwncastSentry/1.2.3 (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)",
id="semver",
),
pytest.param(
"0.0.0",
"OwncastSentry/0.0.0 (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)",
id="zeroed",
),
pytest.param(
"1.1.1.dev10+gf0146d061.d20260313",
"OwncastSentry/1.1.1.dev10+gf0146d061.d20260313 (bot; +https://git.logal.dev/LogalDeveloper/OwncastSentry)",
id="dev-version",
),
],
)
def test_user_agent(self, version: str, expected: str) -> None:
"""Build a correctly formatted User-Agent header."""
assert user_agent(version) == expected
class TestShouldQueryStream:
"""Progressive backoff logic for stream polling."""
@pytest.mark.parametrize(
("counter", "expected"),
[
pytest.param(0, True, id="counter-0-always-query"),
pytest.param(1, True, id="counter-1-always-query"),
pytest.param(4, True, id="counter-4-always-query"),
pytest.param(5, False, id="counter-5-skip-odd"),
pytest.param(6, True, id="counter-6-query-even"),
pytest.param(9, False, id="counter-9-skip-odd"),
pytest.param(10, False, id="counter-10-skip-not-mod-3"),
pytest.param(12, True, id="counter-12-query-mod-3"),
pytest.param(14, False, id="counter-14-skip-not-mod-3"),
pytest.param(15, True, id="counter-15-query-mod-5"),
pytest.param(16, False, id="counter-16-skip-not-mod-5"),
pytest.param(20, True, id="counter-20-query-mod-5"),
pytest.param(29, False, id="counter-29-skip-not-mod-5"),
pytest.param(30, True, id="counter-30-query-mod-15"),
pytest.param(31, False, id="counter-31-skip-not-mod-15"),
pytest.param(45, True, id="counter-45-query-mod-15"),
pytest.param(100, False, id="counter-100-skip-not-mod-15"),
pytest.param(105, True, id="counter-105-query-mod-15"),
],
)
def test_backoff_tiers(self, counter: int, expected: bool) -> None:
"""Return the expected query decision for each backoff tier."""
assert should_query_stream(counter) == expected
class TestDomainify:
"""Domain extraction and sanitization from user input."""
@pytest.mark.parametrize(
("input_url", "expected"),
[
pytest.param("example.com", "example.com", id="bare-domain"),
pytest.param("https://example.com", "example.com", id="https-url"),
pytest.param("http://example.com", "example.com", id="http-url"),
pytest.param("https://example.com:8080", "example.com", id="url-with-port"),
pytest.param(
"https://example.com/path/to/page",
"example.com",
id="url-with-path",
),
pytest.param(
"user@stream.logal.dev",
"stream.logal.dev",
id="email-style",
),
pytest.param("EXAMPLE.COM", "example.com", id="uppercase"),
pytest.param("exam!ple.com", "example.com", id="special-chars-stripped"),
pytest.param(".example.com.", "example.com", id="leading-trailing-dots"),
pytest.param("-example.com-", "example.com", id="leading-trailing-hyphens"),
pytest.param(
"sub.domain.example.com",
"sub.domain.example.com",
id="subdomain",
),
],
)
def test_extracts_domain(self, input_url: str, expected: str) -> None:
"""Extract and sanitize the domain from various input formats."""
assert domainify(input_url) == expected
class TestTruncate:
"""Text truncation to a maximum length."""
@pytest.mark.parametrize(
("text", "max_length", "expected"),
[
pytest.param("hello", 10, "hello", id="under-limit"),
pytest.param("hello", 5, "hello", id="exact-limit"),
pytest.param("hello world", 5, "hello", id="over-limit"),
pytest.param("", 5, "", id="empty-string"),
],
)
def test_truncates(self, text: str, max_length: int, expected: str) -> None:
"""Truncate text that exceeds the maximum length."""
assert truncate(text, max_length) == expected
class TestEscapeMarkdown:
"""Markdown special character escaping."""
@pytest.mark.parametrize(
("input_text", "expected"),
[
pytest.param("hello", "hello", id="plain-text-unchanged"),
pytest.param("*bold*", "\\*bold\\*", id="asterisks"),
pytest.param("_italic_", "\\_italic\\_", id="underscores"),
pytest.param("[link](url)", "\\[link\\]\\(url\\)", id="link-syntax"),
pytest.param("`code`", "\\`code\\`", id="backticks"),
pytest.param("# heading", "\\# heading", id="heading"),
pytest.param("> quote", "\\> quote", id="blockquote"),
pytest.param("<html>", "\\<html\\>", id="angle-brackets"),
pytest.param("a & b", "a \\& b", id="ampersand"),
pytest.param("a\\b", "a\\\\b", id="backslash"),
pytest.param("", "", id="empty-string"),
],
)
def test_escapes_special_chars(self, input_text: str, expected: str) -> None:
"""Escape the given Markdown special character."""
assert escape_markdown(input_text) == expected
class TestSanitizeForPlainText:
"""Plain text sanitization for notifications."""
@pytest.mark.parametrize(
("input_text", "expected"),
[
pytest.param("hello world", "hello world", id="plain-text"),
pytest.param("line1\nline2", "line1 line2", id="newline-removed"),
pytest.param("line1\rline2", "line1 line2", id="carriage-return"),
pytest.param("line1\r\nline2", "line1 line2", id="crlf-removed"),
pytest.param(
"too many spaces", "too many spaces", id="spaces-collapsed"
),
pytest.param("", "", id="empty-string"),
],
)
def test_sanitizes(self, input_text: str, expected: str) -> None:
"""Sanitize the text for safe plain-text rendering."""
assert sanitize_for_plain_text(input_text) == expected
class TestSanitizeForMarkdown:
"""Markdown sanitization combining newline removal and escaping."""
def test_removes_newlines_and_escapes(self) -> None:
"""Remove newlines and escape Markdown special characters."""
result = sanitize_for_markdown("*bold*\nnew line")
assert result == "\\*bold\\* new line"
def test_empty_string(self) -> None:
"""Return empty string unchanged."""
assert sanitize_for_markdown("") == ""