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 .config import Config
from .database import StreamRepository, SubscriptionRepository
from .metrics import ErrorSource, MetricsService
from .migrations import get_upgrade_table
from .notification_service import NotificationService
from .owncast_client import OwncastClient
from .repository import StreamRepository, SubscriptionRepository, get_upgrade_table
from .stream_monitor import StreamMonitor
from .subscription_manager import SubscriptionManager
if TYPE_CHECKING:
from mautrix.util.async_db import Database, UpgradeTable
@@ -97,14 +97,19 @@ class OwncastSentry(Plugin):
metrics=self.metrics_service,
)
# Initialize command handler
self.command_handler = CommandHandler(
# Initialize subscription manager
self.subscription_manager = SubscriptionManager(
self.owncast_client,
self.stream_repo,
self.subscription_repo,
self.log,
)
# Initialize command handler
self.command_handler = CommandHandler(
self.subscription_manager,
)
# Schedule periodic stream state updates every 60 seconds
self.sched.run_periodically(60, self._update_all_stream_states)
+112 -139
View File
@@ -14,20 +14,71 @@
"""Command handlers for OwncastSentry bot commands."""
import sqlite3
from datetime import UTC, datetime
from typing import TYPE_CHECKING
from .models import StreamStatus
from .utils import domainify, sanitize_for_markdown
from .types import (
AlreadySubscribedError,
InvalidOwncastInstanceError,
NotSubscribedError,
StreamStatus,
)
if TYPE_CHECKING:
import logging
from maubot import MessageEvent # type: ignore[attr-defined]
from .database import StreamRepository, SubscriptionRepository
from .owncast_client import OwncastClient
from .subscription_manager import SubscriptionManager
_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:
@@ -35,22 +86,13 @@ class CommandHandler:
def __init__(
self,
owncast_client: OwncastClient,
stream_repo: StreamRepository,
subscription_repo: SubscriptionRepository,
logger: logging.Logger,
subscription_manager: SubscriptionManager,
) -> None:
"""Initialize the command handler.
:param owncast_client: Client for making API calls to Owncast instances.
:param stream_repo: Repository for stream data.
:param subscription_repo: Repository for subscription data.
:param logger: Logger instance for debugging.
:param subscription_manager: Subscription domain workflow coordinator.
"""
self.owncast_client = owncast_client
self.stream_repo = stream_repo
self.subscription_repo = subscription_repo
self.log = logger
self.subscription_manager = subscription_manager
async def subscribe(self, evt: MessageEvent, url: str) -> None:
"""Subscribe a room to a stream's notifications.
@@ -58,46 +100,22 @@ class CommandHandler:
:param evt: MessageEvent of the message calling the command.
: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:
await self.subscription_repo.add(stream_domain, evt.room_id)
except sqlite3.IntegrityError:
# Room is already subscribed.
stream_domain = await self.subscription_manager.subscribe(evt.room_id, url)
except InvalidOwncastInstanceError:
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
# 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(
f"Subscription added! This room will receive "
f"notifications when {stream_domain} goes live."
@@ -109,66 +127,32 @@ class CommandHandler:
:param evt: MessageEvent of the message calling the command.
:param url: User supplied URL to a stream to unsubscribe from.
"""
# Convert the user input to only a domain
stream_domain = domainify(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}."
try:
stream_domain = await self.subscription_manager.unsubscribe(
evt.room_id, url
)
await evt.reply(
f"Subscription removed! This room will no "
f"longer receive notifications for {stream_domain}."
)
else:
# No, nothing changed. Tell the user.
except NotSubscribedError as e:
await evt.reply(
"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:
"""Calculate and format the duration from a timestamp to now.
: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"
await evt.reply(
f"Subscription removed! This room will no "
f"longer receive notifications for {stream_domain}."
)
async def subscriptions(self, evt: MessageEvent) -> None:
"""List all stream subscriptions in the current room.
:param evt: MessageEvent of the message calling the command.
"""
# Get all stream domains this room is subscribed to
subscribed_domains = (
await self.subscription_repo.get_subscribed_streams_for_room(evt.room_id)
subscriptions = await self.subscription_manager.list_room_subscriptions(
evt.room_id
)
# Check if there are no subscriptions
if not subscribed_domains:
if not subscriptions:
await evt.reply(
"This room is not subscribed to any Owncast "
"instances.\n\nTo subscribe to an Owncast "
@@ -178,36 +162,33 @@ class CommandHandler:
return
# Build the response message body as Markdown
count = len(subscribed_domains)
count = len(subscriptions)
parts = [f"**Subscriptions for this room ({count}):**\n\n"]
now = datetime.now(UTC)
for domain in subscribed_domains:
# Get the stream state from the database
stream_state = await self.stream_repo.get_by_domain(domain)
if stream_state is None:
continue
# Determine stream name (use domain as fallback)
for subscription in subscriptions:
domain = subscription.domain
stream_state = subscription.stream_state
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
parts.append(f"- **{safe_stream_name}** \n")
# Add title if stream is online (as a sub-bullet)
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")
# Determine status and duration (as a sub-bullet)
match stream_state.status:
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")
case StreamStatus.UNKNOWN:
parts.append(" - Status: Unknown (instance unreachable) \n")
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")
case StreamStatus.OFFLINE:
parts.append(" - Status: Offline \n")
@@ -229,30 +210,20 @@ class CommandHandler:
:param evt: MessageEvent of the message calling the command.
"""
# Get all stream domains this room is subscribed to
subscribed_domains = (
await self.subscription_repo.get_subscribed_streams_for_room(evt.room_id)
live_streams = await self.subscription_manager.list_live_room_subscriptions(
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 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(
"No subscribed Owncast instances are currently "
"live.\n\nUse `!subscriptions` to list all "
@@ -264,23 +235,25 @@ class CommandHandler:
# Build the response message body as Markdown
count = len(live_streams)
parts = [f"**Live Owncast instances ({count}):**\n\n"]
now = datetime.now(UTC)
for domain, stream_state in live_streams:
# Determine stream name (use domain as fallback)
for subscription in live_streams:
domain = subscription.domain
stream_state = subscription.stream_state
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
parts.append(f"- **{safe_stream_name}** \n")
# Add title (should be present for live streams)
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")
# Add status with duration
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")
# 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 .models import StreamStatus
from .types import StreamStatus
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 .metrics import NotificationType
from .utils import (
CLEANUP_DELETE_DAYS,
CLEANUP_WARNING_DAYS,
SECONDS_BETWEEN_NOTIFICATIONS,
sanitize_for_plain_text,
)
if TYPE_CHECKING:
import logging
from collections.abc import Sequence
from .database import SubscriptionRepository
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:
@@ -65,7 +74,7 @@ class NotificationService:
domain: str,
name: str,
title: str,
tags: list[str],
tags: Sequence[str],
*,
title_change: bool = False,
) -> None:
@@ -83,28 +92,32 @@ class NotificationService:
time.monotonic() - self.notification_timers_cache[domain]
)
self.log.info(
f"[{domain}] Not sending notifications. Only "
f"{seconds_since_last} of required "
f"{SECONDS_BETWEEN_NOTIFICATIONS} seconds have "
f"passed since last notification."
"[%s] Not sending notifications. Only %s of required "
"%s seconds have passed since last notification.",
domain,
seconds_since_last,
_SECONDS_BETWEEN_NOTIFICATIONS,
)
return
# Record that we're sending a notification now
self._record_notification(domain)
# Build the notification message
body_text = self._format_message(name, title, domain, tags, title_change)
# Send notifications to all subscribed rooms in parallel
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
notification_type = "title change" if title_change else "going live"
self.log.info(
f"[{domain}] Completed sending {notification_type} "
f"notifications! {successful} succeeded, "
f"{failed} failed."
"[%s] Completed sending %s notifications! %s succeeded, %s failed.",
domain,
notification_type,
successful,
failed,
)
self.metrics.record_delivery(
@@ -113,6 +126,75 @@ class NotificationService:
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(
self, room_id: str, body_text: str, domain: str
) -> None:
@@ -128,13 +210,20 @@ class NotificationService:
await self.client.send_message(room_id, content)
except Exception as exception:
self.log.warning(
f"[{domain}] Failed to send notification "
f"message to room [{room_id}]: {exception}"
"[%s] Failed to send notification message to room [%s]: %s",
domain,
room_id,
exception,
)
raise
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:
"""Format the notification message body.
@@ -147,7 +236,7 @@ class NotificationService:
"""
# Use name if available, fallback to 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
if title_change:
@@ -157,7 +246,7 @@ class NotificationService:
# Add title if present
if title:
safe_title = sanitize_for_plain_text(title)
safe_title = _sanitize_for_plain_text(title)
parts.append(f"\nStream Title: {safe_title}")
# Add stream URL
@@ -165,39 +254,30 @@ class NotificationService:
# Add tags if present
if tags:
safe_tags = [
safe_tag
tag_text = " ".join(
f"#{safe_tag}"
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(".")
]
)
if safe_tags:
parts.append(f"\n\n{' '.join(f'#{tag}' for tag in safe_tags)}")
if tag_text:
parts.append(f"\n\n{tag_text}")
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:
"""Check if enough time has passed to send another notification.
:param domain: The stream domain.
:return: True if notification can be sent, False otherwise.
"""
if domain not in self.notification_timers_cache:
return True
seconds_since_last = round(
time.monotonic() - self.notification_timers_cache[domain]
last_notification_time = self.notification_timers_cache.get(domain)
return (
last_notification_time is None
or time.monotonic() - last_notification_time
>= _SECONDS_BETWEEN_NOTIFICATIONS
)
return seconds_since_last >= SECONDS_BETWEEN_NOTIFICATIONS
def _record_notification(self, domain: str) -> None:
"""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
]
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
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."""
import json
from typing import TYPE_CHECKING, Any
import aiohttp
from .models import StreamConfig, StreamState
from .utils import (
OWNCAST_CONFIG_PATH,
OWNCAST_STATUS_PATH,
REQUIRED_STATUS_FIELDS,
user_agent,
)
from .types import InvalidApiResponseError, StreamConfig, StreamState
if TYPE_CHECKING:
import logging
@@ -32,6 +27,48 @@ if TYPE_CHECKING:
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:
"""HTTP client for communicating with Owncast instances."""
@@ -51,15 +88,18 @@ class OwncastClient:
self.metrics = metrics
# Set up HTTP session configuration
headers = {"User-Agent": user_agent(version)}
headers = {"User-Agent": _user_agent(version)}
cookie_jar = aiohttp.DummyCookieJar()
connector = aiohttp.TCPConnector(
use_dns_cache=False,
limit=1000,
limit_per_host=1,
keepalive_timeout=120,
limit=_HTTP_CONNECTION_LIMIT,
limit_per_host=_HTTP_CONNECTION_LIMIT_PER_HOST,
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(
headers=headers,
@@ -80,23 +120,48 @@ class OwncastClient:
async with self.session.get(url, allow_redirects=False) as response:
if response.status != 200:
self.log.warning(
f"[{domain}] Response to request on "
f"{path} was not 200, "
f"got {response.status} instead."
"[%s] Response to request on %s was not 200, "
"got %s instead.",
domain,
path,
response.status,
)
return None
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
except (ValueError, aiohttp.ContentTypeError) as e:
except ValueError as e:
self.log.warning(
f"[{domain}] Rejecting response to request on "
f"{path} as could not be "
f"interpreted as JSON: {e}"
"[%s] Rejecting response to request on %s as could not "
"be interpreted as JSON: %s",
domain,
path,
e,
)
return None
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
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.
: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:
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:
return None
# Validate the response contains all basic info needed
missing = REQUIRED_STATUS_FIELDS - new_state.keys()
if missing:
try:
stream_state = StreamState.from_api_response(new_state, domain)
except InvalidApiResponseError as e:
self.log.warning(
f"[{domain}] Rejecting response to request on "
f"{OWNCAST_STATUS_PATH} as it is missing "
f"fields: {', '.join(sorted(missing))}"
"[%s] Rejecting response to request on %s as response "
"shape is invalid: %s",
domain,
_OWNCAST_STATUS_PATH,
e,
)
return None
timer.success()
return StreamState.from_api_response(new_state, domain)
return stream_state
async def get_stream_config(self, domain: str) -> StreamConfig | None:
"""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.
: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:
config = await self._fetch_json(domain, OWNCAST_CONFIG_PATH)
config = await self._fetch_json(domain, _OWNCAST_CONFIG_PATH)
if config is 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()
return StreamConfig.from_api_response(config)
return stream_config
async def validate_instance(self, domain: str) -> bool:
"""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
from typing import TYPE_CHECKING
from .models import StreamState, StreamStatus, UpdateResult
from .utils import (
CLEANUP_DELETE_THRESHOLD,
CLEANUP_WARNING_THRESHOLD,
TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN,
should_query_stream,
)
from .types import StreamState, StreamStatus, UpdateResult
if TYPE_CHECKING:
import logging
from .database import StreamRepository, SubscriptionRepository
from .metrics import MetricsService
from .notification_service import NotificationService
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:
@@ -92,7 +105,8 @@ class StreamMonitor:
for domain, result in zip(subscribed_domains, results, strict=True):
if isinstance(result, BaseException):
self.log.exception(
f"[{domain}] Unhandled exception during stream update.",
"[%s] Unhandled exception during stream update.",
domain,
exc_info=result,
)
failed_checks += 1
@@ -102,13 +116,19 @@ class StreamMonitor:
failed_checks += 1
self.log.debug(
f"Update complete. {successful_checks}/{total_streams} succeeded, "
f"{failed_checks} failed."
"Update complete. %s/%s succeeded, %s failed.",
successful_checks,
total_streams,
failed_checks,
)
subscription_counts = await self.subscription_repo.count_by_domains(
subscribed_domains
)
for domain in subscribed_domains:
count = await self.subscription_repo.count_by_domain(domain)
self.metrics.set_subscription_count(domain, count)
self.metrics.set_subscription_count(
domain, subscription_counts.get(domain, 0)
)
return UpdateResult(
total_streams=total_streams,
@@ -134,19 +154,20 @@ class StreamMonitor:
return True
# 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
await self.stream_repo.increment_failure_counter(domain)
self.log.debug(
f"[{domain}] Skipping query due to backoff "
f"(counter={failure_counter + 1})"
"[%s] Skipping query due to backoff (counter=%s)",
domain,
failure_counter + 1,
)
# Check cleanup thresholds even when skipping query
await self._check_cleanup_thresholds(domain, failure_counter + 1)
updated_state = await self.stream_repo.get_by_domain(domain)
if updated_state is not None:
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
return True
@@ -168,14 +189,16 @@ class StreamMonitor:
if new_state is None:
await self.stream_repo.increment_failure_counter(domain)
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
await self._check_cleanup_thresholds(domain, failure_counter + 1)
updated_state = await self.stream_repo.get_by_domain(domain)
if updated_state is not None:
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
return False
@@ -204,7 +227,7 @@ class StreamMonitor:
update_database = True
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
seconds_since_last_offline = round(
@@ -215,10 +238,13 @@ class StreamMonitor:
if not first_update:
# Use fallback values if config fetch failed
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?
if seconds_since_last_offline < TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN:
if (
seconds_since_last_offline
< _TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN
):
# Did the stream title change?
if old_state.title != new_state.title:
# Stream was briefly down; send title
@@ -233,13 +259,12 @@ class StreamMonitor:
else:
# Briefly offline, no title change. Skip.
self.log.info(
f"[{domain}] Not sending "
f"notifications. Stream was only "
f"offline for "
f"{seconds_since_last_offline} of "
f"{TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN}"
f" seconds and did not change its "
f"title."
"[%s] Not sending notifications. Stream was only "
"offline for %s of %s seconds and did not change "
"its title.",
domain,
seconds_since_last_offline,
_TEMPORARY_OFFLINE_NOTIFICATION_COOLDOWN,
)
else:
# Offline for a while. Send a normal notification.
@@ -253,9 +278,9 @@ class StreamMonitor:
else:
# No, this is the first time we're querying
self.log.info(
f"[{domain}] Not sending notifications. "
f"This is the first state update for "
f"this stream."
"[%s] Not sending notifications. This is the first state "
"update for this stream.",
domain,
)
if (
@@ -264,13 +289,13 @@ class StreamMonitor:
):
# Did the stream title change mid-session?
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
stream_config = await self.owncast_client.get_stream_config(domain)
# Use fallback values if config fetch failed
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
# last went offline? If so, send a regular go-live
@@ -303,7 +328,7 @@ class StreamMonitor:
# Yep. This stream is now offline. Log it.
update_database = True
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.
if update_database:
@@ -314,7 +339,7 @@ class StreamMonitor:
# Use fallback value if config fetch failed
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)
updated_state = StreamState(
@@ -328,7 +353,7 @@ class StreamMonitor:
await self.stream_repo.update(updated_state)
# 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:
self.metrics.set_stream_status(domain, StreamStatus.ONLINE)
else:
@@ -342,17 +367,19 @@ class StreamMonitor:
:param counter: The current failure counter value.
"""
# Check for 83-day warning threshold
if counter == CLEANUP_WARNING_THRESHOLD:
if counter == _CLEANUP_WARNING_THRESHOLD:
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)
# Check for 90-day deletion threshold
if counter >= CLEANUP_DELETE_THRESHOLD:
if counter >= _CLEANUP_DELETE_THRESHOLD:
self.log.warning(
f"[{domain}] Reached 90-day deletion threshold."
f" Removing all subscriptions."
"[%s] Reached 90-day deletion threshold. "
"Removing all subscriptions.",
domain,
)
# Send deletion notification
await self.notification_service.send_cleanup_deletion(domain)
@@ -362,10 +389,13 @@ class StreamMonitor:
# Delete the stream record
await self.stream_repo.delete(domain)
self.offline_timer_cache.pop(domain, None)
self.notification_service.clear_notification_state(domain)
self.log.info(
f"[{domain}] Cleanup complete. "
f"Deleted {deleted_count} subscriptions "
f"and stream record."
"[%s] Cleanup complete. Deleted %s subscriptions "
"and stream record.",
domain,
deleted_count,
)
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.config import Config
from owncastsentry.database import StreamRepository, SubscriptionRepository
from owncastsentry.migrations import get_upgrade_table
from owncastsentry.repository import (
StreamRepository,
SubscriptionRepository,
get_upgrade_table,
)
if TYPE_CHECKING:
from collections.abc import AsyncIterator
from pathlib import Path
from owncastsentry.metrics import MetricsService
from owncastsentry.models import StreamConfig, StreamState
from owncastsentry.types import StreamConfig, StreamState
def generate_metrics_output(metrics: MetricsService) -> str:
+102 -32
View File
@@ -15,28 +15,57 @@
"""Tests for bot command handlers."""
import json
import logging
from datetime import UTC, datetime, timedelta
from unittest.mock import MagicMock
import pytest
import time_machine
from aioresponses import aioresponses
from owncastsentry.commands import CommandHandler
from owncastsentry.models import StreamState
from owncastsentry.utils import OWNCAST_STATUS_PATH, UNKNOWN_STATUS_THRESHOLD
from owncastsentry.commands import (
_escape_markdown,
_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
def _make_command_handler() -> CommandHandler:
"""Build a CommandHandler with dummy dependencies for pure logic tests."""
return CommandHandler(
owncast_client=MagicMock(),
stream_repo=MagicMock(),
subscription_repo=MagicMock(),
logger=logging.getLogger("test"),
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 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:
@@ -57,18 +86,24 @@ class TestFormatDuration:
pytest.param(172800, "2 days", id="plural-days"),
],
)
@time_machine.travel(_NOW)
def test_formats_duration(self, seconds_ago: int, expected: str) -> None:
"""Format a timestamp into a human-readable duration."""
handler = _make_command_handler()
timestamp = (self._NOW - timedelta(seconds=seconds_ago)).isoformat()
result = handler._format_duration(timestamp)
result = _format_duration(timestamp, self._NOW)
assert result == expected
def test_invalid_timestamp(self) -> None:
"""Return 'unknown duration' for unparsable timestamps."""
handler = _make_command_handler()
assert handler._format_duration("not-a-timestamp") == "unknown duration"
assert _format_duration("not-a-timestamp", self._NOW) == "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:
@@ -76,7 +111,7 @@ class TestSubscribeCommand:
async def test_subscribe_valid_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""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:
mocked.get(
status_url,
@@ -94,7 +129,7 @@ class TestSubscribeCommand:
self, maubot_test_bot, maubot_plugin
) -> None:
"""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:
mocked.get(status_url, status=404)
await maubot_test_bot.send("!subscribe invalid.com")
@@ -111,7 +146,7 @@ class TestSubscribeCommand:
self, maubot_test_bot, maubot_plugin
) -> None:
"""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:
mocked.get(
status_url,
@@ -131,7 +166,7 @@ class TestSubscribeCommand:
self, maubot_test_bot, maubot_plugin
) -> None:
"""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:
mocked.get(
status_url,
@@ -159,7 +194,7 @@ class TestUnsubscribeCommand:
async def test_unsubscribe_existing(self, maubot_test_bot, maubot_plugin) -> None:
"""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:
mocked.get(
status_url,
@@ -205,7 +240,7 @@ class TestSubscriptionsCommand:
async def test_shows_online_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show stream details including title and duration."""
# 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:
mocked.get(
status_url,
@@ -236,11 +271,46 @@ class TestSubscriptionsCommand:
"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))
async def test_shows_offline_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show offline status for non-live streams."""
# 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:
mocked.get(
status_url,
@@ -273,7 +343,7 @@ class TestSubscriptionsCommand:
self, maubot_test_bot, maubot_plugin
) -> None:
"""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:
mocked.get(
status_url,
@@ -296,7 +366,7 @@ class TestSubscriptionsCommand:
async def test_shows_unknown_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""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:
mocked.get(
status_url,
@@ -330,11 +400,11 @@ class TestSubscriptionsCommand:
# Subscribe in reverse alphabetical order to verify sorted output
with aioresponses() as mocked:
mocked.get(
f"https://beta.com{OWNCAST_STATUS_PATH}",
f"https://beta.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
)
mocked.get(
f"https://alpha.com{OWNCAST_STATUS_PATH}",
f"https://alpha.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
)
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:
"""Show 'no live' message when all streams are offline."""
# 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:
mocked.get(
status_url,
@@ -423,7 +493,7 @@ class TestLiveCommand:
async def test_shows_live_stream(self, maubot_test_bot, maubot_plugin) -> None:
"""Show live stream with title and duration."""
# 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:
mocked.get(
status_url,
@@ -460,11 +530,11 @@ class TestLiveCommand:
# Subscribe in reverse alphabetical order to verify sorted output
with aioresponses() as mocked:
mocked.get(
f"https://beta.com{OWNCAST_STATUS_PATH}",
f"https://beta.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
)
mocked.get(
f"https://alpha.com{OWNCAST_STATUS_PATH}",
f"https://alpha.com{_OWNCAST_STATUS_PATH}",
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
)
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
from owncastsentry.metrics import ErrorSource, MetricsService, NotificationType
from owncastsentry.models import StreamStatus
from owncastsentry.types import StreamStatus
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."""
import asyncio
import logging
import time
from typing import TYPE_CHECKING
@@ -21,12 +22,15 @@ from typing import TYPE_CHECKING
import pytest
from owncastsentry.metrics import MetricsService
from owncastsentry.notification_service import NotificationService
from owncastsentry.utils import SECONDS_BETWEEN_NOTIFICATIONS
from owncastsentry.notification_service import (
_SECONDS_BETWEEN_NOTIFICATIONS,
NotificationService,
_sanitize_for_plain_text,
)
from tests.conftest import _StubMatrixClient, generate_metrics_output
if TYPE_CHECKING:
from owncastsentry.database import StreamRepository, SubscriptionRepository
from owncastsentry.repository import StreamRepository, SubscriptionRepository
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:
"""Rate-limiting logic for notification cooldowns."""
@@ -75,7 +100,7 @@ class TestCanNotify:
)
# Subtract an extra second to ensure the cooldown has fully elapsed
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
@@ -103,6 +128,35 @@ class TestGetLastNotificationTime:
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:
"""Notification message formatting."""
@@ -246,6 +300,74 @@ class TestNotifyStreamLive:
for msg in client.sent_messages:
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(
self,
stream_repo: StreamRepository,
+231 -27
View File
@@ -22,7 +22,12 @@ import pytest
from aioresponses import aioresponses
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 (
VALID_CONFIG_RESPONSE,
VALID_STATUS_RESPONSE,
@@ -33,6 +38,32 @@ if TYPE_CHECKING:
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
async def owncast_client() -> AsyncIterator[OwncastClient]:
"""Create an OwncastClient and close it after the test."""
@@ -45,6 +76,79 @@ async def owncast_client() -> AsyncIterator[OwncastClient]:
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:
"""Stream state retrieval from the status API."""
@@ -82,6 +186,23 @@ class TestGetStreamState:
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(
self, owncast_client: OwncastClient
) -> None:
@@ -95,6 +216,32 @@ class TestGetStreamState:
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:
"""Return None when the response status is not 200."""
with aioresponses() as mocked:
@@ -136,7 +283,7 @@ class TestGetStreamConfig:
assert result is not None
assert result.name == "LogalDeveloper's Live Stream"
assert result.tags == [
assert result.tags == (
"video games",
"chatting",
"casual",
@@ -144,7 +291,7 @@ class TestGetStreamConfig:
"streaming",
"owncast",
"variety",
]
)
async def test_returns_none_on_invalid_json(
self, owncast_client: OwncastClient
@@ -159,6 +306,49 @@ class TestGetStreamConfig:
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:
"""Return None when the response status is not 200."""
with aioresponses() as mocked:
@@ -225,15 +415,19 @@ class TestResponseTimeMetrics:
version="0.0.0",
metrics=metrics,
)
with aioresponses() as mocked:
mocked.get(
"https://example.com/api/status",
body=json.dumps(VALID_STATUS_RESPONSE).encode(),
try:
with aioresponses() as mocked:
mocked.get(
"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")
output = generate_metrics_output(metrics)
assert 'owncastsentry_api_response_seconds{domain="example.com"}' in output
await client.close()
finally:
await client.close()
async def test_no_observation_on_failure(self) -> None:
"""Do not record response time when request fails."""
@@ -243,15 +437,20 @@ class TestResponseTimeMetrics:
version="0.0.0",
metrics=metrics,
)
with aioresponses() as mocked:
mocked.get(
"https://example.com/api/status",
status=500,
try:
with aioresponses() as mocked:
mocked.get(
"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")
output = generate_metrics_output(metrics)
assert 'owncastsentry_api_response_seconds{domain="example.com"}' not in output
await client.close()
finally:
await client.close()
async def test_no_observation_on_connection_error(self) -> None:
"""Do not record response time on connection error."""
@@ -261,15 +460,20 @@ class TestResponseTimeMetrics:
version="0.0.0",
metrics=metrics,
)
with aioresponses() as mocked:
mocked.get(
"https://example.com/api/status",
exception=ConnectionError(),
try:
with aioresponses() as mocked:
mocked.get(
"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")
output = generate_metrics_output(metrics)
assert 'owncastsentry_api_response_seconds{domain="example.com"}' not in output
await client.close()
finally:
await client.close()
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
from typing import TYPE_CHECKING
import pytest
from owncastsentry.metrics import MetricsService
from owncastsentry.models import StreamConfig, StreamState, StreamStatus
from owncastsentry.notification_service import NotificationService
from owncastsentry.stream_monitor import StreamMonitor
from owncastsentry.utils import (
CLEANUP_DELETE_THRESHOLD,
CLEANUP_WARNING_THRESHOLD,
SECONDS_BETWEEN_NOTIFICATIONS,
from owncastsentry.notification_service import (
_SECONDS_BETWEEN_NOTIFICATIONS,
NotificationService,
)
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 (
_StubMatrixClient,
_StubOwncastClient,
@@ -34,7 +40,7 @@ from tests.conftest import (
)
if TYPE_CHECKING:
from owncastsentry.database import StreamRepository, SubscriptionRepository
from owncastsentry.repository import StreamRepository, SubscriptionRepository
def _make_monitor(
@@ -108,6 +114,37 @@ def _make_monitor_with_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:
"""Parallel stream update orchestration."""
@@ -252,7 +289,7 @@ class TestUpdateStreamGoesLive:
last_connect_time="2026-01-01T12: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()
monitor, _ = _make_monitor(
@@ -271,7 +308,9 @@ class TestUpdateStreamGoesLive:
)
# 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")
assert result is True
@@ -315,7 +354,9 @@ class TestUpdateStreamGoesLive:
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")
assert result is True
@@ -447,10 +488,13 @@ class TestUpdateStreamTitleChange:
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
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")
@@ -496,10 +540,13 @@ class TestUpdateStreamTitleChange:
# Last notification was long enough ago to pass rate limiting,
# 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
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")
@@ -546,10 +593,10 @@ class TestUpdateStreamTitleChange:
# and both are old enough to pass rate limiting
now = time.monotonic()
monitor.offline_timer_cache["title.com"] = (
now - SECONDS_BETWEEN_NOTIFICATIONS - 100
now - _SECONDS_BETWEEN_NOTIFICATIONS - 100
)
notification_service.notification_timers_cache["title.com"] = (
now - SECONDS_BETWEEN_NOTIFICATIONS - 200
now - _SECONDS_BETWEEN_NOTIFICATIONS - 200
)
result = await monitor.update_stream("title.com")
@@ -659,7 +706,7 @@ class TestCheckCleanupThresholds:
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 client.sent_messages[0].content.body == (
@@ -680,7 +727,7 @@ class TestCheckCleanupThresholds:
"""Delete all subscriptions and the stream record at the 90-day threshold."""
owncast = _StubOwncastClient()
client = _StubMatrixClient()
monitor, _ = _make_monitor(
monitor, notification_service = _make_monitor(
owncast_client=owncast,
stream_repo=stream_repo,
subscription_repo=subscription_repo,
@@ -688,8 +735,10 @@ class TestCheckCleanupThresholds:
)
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
assert len(client.sent_messages) == 1
@@ -709,6 +758,8 @@ class TestCheckCleanupThresholds:
assert await stream_repo.get_by_domain("delete.com") is None
rooms = await subscription_repo.get_subscribed_rooms("delete.com")
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(
self,
@@ -1061,7 +1112,34 @@ class TestStreamMonitorMetrics:
metrics.set_stream_status("delete.com", StreamStatus.OFFLINE)
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)
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("") == ""