Proyecto LCX Dispatcharr multicuenta
Base Image Build / prepare (push) Has been cancelled
Build and Push Multi-Arch Docker Image / build-and-push (push) Has been cancelled
Frontend Tests / test (push) Has been cancelled
Base Image Build / docker (amd64, ubuntu-24.04) (push) Has been cancelled
Base Image Build / docker (arm64, ubuntu-24.04-arm) (push) Has been cancelled
Base Image Build / create-manifest (push) Has been cancelled

This commit is contained in:
root
2026-05-09 21:24:50 +02:00
commit f56b088643
721 changed files with 177870 additions and 0 deletions
View File
+211
View File
@@ -0,0 +1,211 @@
from django.test import TestCase
from django.contrib.auth import get_user_model
from rest_framework.test import APIClient
from rest_framework import status
from apps.channels.models import Channel, ChannelGroup
User = get_user_model()
class ChannelBulkEditAPITests(TestCase):
def setUp(self):
# Create a test admin user (user_level >= 10) and authenticate
self.user = User.objects.create_user(username="testuser", password="testpass123")
self.user.user_level = 10 # Set admin level
self.user.save()
self.client = APIClient()
self.client.force_authenticate(user=self.user)
self.bulk_edit_url = "/api/channels/channels/edit/bulk/"
# Create test channel group
self.group1 = ChannelGroup.objects.create(name="Test Group 1")
self.group2 = ChannelGroup.objects.create(name="Test Group 2")
# Create test channels
self.channel1 = Channel.objects.create(
channel_number=1.0,
name="Channel 1",
tvg_id="channel1",
channel_group=self.group1
)
self.channel2 = Channel.objects.create(
channel_number=2.0,
name="Channel 2",
tvg_id="channel2",
channel_group=self.group1
)
self.channel3 = Channel.objects.create(
channel_number=3.0,
name="Channel 3",
tvg_id="channel3"
)
def test_bulk_edit_success(self):
"""Test successful bulk update of multiple channels"""
data = [
{"id": self.channel1.id, "name": "Updated Channel 1"},
{"id": self.channel2.id, "name": "Updated Channel 2", "channel_number": 22.0},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
self.assertEqual(len(response.data["channels"]), 2)
# Verify database changes
self.channel1.refresh_from_db()
self.channel2.refresh_from_db()
self.assertEqual(self.channel1.name, "Updated Channel 1")
self.assertEqual(self.channel2.name, "Updated Channel 2")
self.assertEqual(self.channel2.channel_number, 22.0)
def test_bulk_edit_with_empty_validated_data_first(self):
"""
Test the bug fix: when first channel has empty validated_data.
This was causing: ValueError: Field names must be given to bulk_update()
"""
# Create a channel with data that will be "unchanged" (empty validated_data)
# We'll send the same data it already has
data = [
# First channel: no actual changes (this would create empty validated_data)
{"id": self.channel1.id},
# Second channel: has changes
{"id": self.channel2.id, "name": "Updated Channel 2"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should not crash with ValueError
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
# Verify the channel with changes was updated
self.channel2.refresh_from_db()
self.assertEqual(self.channel2.name, "Updated Channel 2")
def test_bulk_edit_all_empty_updates(self):
"""Test when all channels have empty updates (no actual changes)"""
data = [
{"id": self.channel1.id},
{"id": self.channel2.id},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should succeed without calling bulk_update
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
def test_bulk_edit_mixed_fields(self):
"""Test bulk update where different channels update different fields"""
data = [
{"id": self.channel1.id, "name": "New Name 1"},
{"id": self.channel2.id, "channel_number": 99.0},
{"id": self.channel3.id, "tvg_id": "new_tvg_id", "name": "New Name 3"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 3 channels")
# Verify all updates
self.channel1.refresh_from_db()
self.channel2.refresh_from_db()
self.channel3.refresh_from_db()
self.assertEqual(self.channel1.name, "New Name 1")
self.assertEqual(self.channel2.channel_number, 99.0)
self.assertEqual(self.channel3.tvg_id, "new_tvg_id")
self.assertEqual(self.channel3.name, "New Name 3")
def test_bulk_edit_with_channel_group(self):
"""Test bulk update with channel_group_id changes"""
data = [
{"id": self.channel1.id, "channel_group_id": self.group2.id},
{"id": self.channel3.id, "channel_group_id": self.group1.id},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
# Verify group changes
self.channel1.refresh_from_db()
self.channel3.refresh_from_db()
self.assertEqual(self.channel1.channel_group, self.group2)
self.assertEqual(self.channel3.channel_group, self.group1)
def test_bulk_edit_nonexistent_channel(self):
"""Test bulk update with a channel that doesn't exist"""
nonexistent_id = 99999
data = [
{"id": nonexistent_id, "name": "Should Fail"},
{"id": self.channel1.id, "name": "Should Still Update"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should return 400 with errors
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn("errors", response.data)
self.assertEqual(len(response.data["errors"]), 1)
self.assertEqual(response.data["errors"][0]["channel_id"], nonexistent_id)
self.assertEqual(response.data["errors"][0]["error"], "Channel not found")
# The valid channel should still be updated
self.assertEqual(response.data["updated_count"], 1)
def test_bulk_edit_validation_error(self):
"""Test bulk update with invalid data (validation error)"""
data = [
{"id": self.channel1.id, "channel_number": "invalid_number"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should return 400 with validation errors
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn("errors", response.data)
self.assertEqual(len(response.data["errors"]), 1)
self.assertIn("channel_number", response.data["errors"][0]["errors"])
def test_bulk_edit_empty_channel_updates(self):
"""Test bulk update with empty list"""
data = []
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Empty list is accepted and returns success with 0 updates
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 0 channels")
def test_bulk_edit_missing_channel_updates(self):
"""Test bulk update without proper format (dict instead of list)"""
data = {"channel_updates": {}}
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertEqual(response.data["error"], "Expected a list of channel updates")
def test_bulk_edit_preserves_other_fields(self):
"""Test that bulk update only changes specified fields"""
original_channel_number = self.channel1.channel_number
original_tvg_id = self.channel1.tvg_id
data = [
{"id": self.channel1.id, "name": "Only Name Changed"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
# Verify only name changed, other fields preserved
self.channel1.refresh_from_db()
self.assertEqual(self.channel1.name, "Only Name Changed")
self.assertEqual(self.channel1.channel_number, original_channel_number)
self.assertEqual(self.channel1.tvg_id, original_tvg_id)
+278
View File
@@ -0,0 +1,278 @@
"""Tests for DVR retry logic.
Covers:
- _db_retry(): exponential backoff, max retries, connection reset
- Final metadata save retry in run_recording post-processing
- Initial TS proxy connection retry (per-base retry on retriable errors)
- recover_recordings_on_startup DB retry wrappers
"""
from datetime import timedelta
from unittest.mock import MagicMock, patch, call
from django.db import OperationalError
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel, Recording
from apps.channels.tasks import _db_retry
# ---------------------------------------------------------------------------
# _db_retry unit tests
# ---------------------------------------------------------------------------
class DbRetryTests(TestCase):
"""Tests for the _db_retry() exponential backoff helper."""
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_succeeds_on_first_attempt(self, _close, _sleep):
"""No retry needed when fn succeeds immediately."""
result = _db_retry(lambda: "ok", max_retries=3)
self.assertEqual(result, "ok")
_sleep.assert_not_called()
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_retries_on_operational_error_then_succeeds(self, mock_close, mock_sleep):
"""Retry succeeds on second attempt after OperationalError."""
call_count = {"n": 0}
def flaky():
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("connection reset")
return "recovered"
result = _db_retry(flaky, max_retries=3, base_interval=1)
self.assertEqual(result, "recovered")
self.assertEqual(call_count["n"], 2)
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_raises_after_max_retries_exhausted(self, mock_close, mock_sleep):
"""Raises OperationalError after all retries fail."""
def always_fail():
raise OperationalError("db gone")
with self.assertRaises(OperationalError):
_db_retry(always_fail, max_retries=3, base_interval=1)
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_exponential_backoff_timing(self, mock_close, mock_sleep):
"""Sleep durations follow exponential backoff: 1s, 2s, 4s."""
call_count = {"n": 0}
def fail_twice():
call_count["n"] += 1
if call_count["n"] <= 2:
raise OperationalError("retry me")
return "done"
_db_retry(fail_twice, max_retries=3, base_interval=1)
mock_sleep.assert_has_calls([call(1), call(2)])
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_close_old_connections_called_between_retries(self, mock_close, mock_sleep):
"""Stale DB connections are reset before each retry attempt."""
call_count = {"n": 0}
def fail_once():
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("stale conn")
return "ok"
_db_retry(fail_once, max_retries=3)
mock_close.assert_called_once()
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_non_operational_error_not_retried(self, mock_close, mock_sleep):
"""Non-OperationalError exceptions propagate immediately."""
def raise_value_error():
raise ValueError("not a DB error")
with self.assertRaises(ValueError):
_db_retry(raise_value_error, max_retries=3)
mock_sleep.assert_not_called()
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_returns_fn_return_value(self, mock_close, mock_sleep):
"""Return value of fn() is passed through."""
result = _db_retry(lambda: {"key": "value"}, max_retries=3)
self.assertEqual(result, {"key": "value"})
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_single_retry_allowed(self, mock_close, mock_sleep):
"""max_retries=1 means no retry — fail immediately."""
with self.assertRaises(OperationalError):
_db_retry(
lambda: (_ for _ in ()).throw(OperationalError("fail")),
max_retries=1,
)
mock_sleep.assert_not_called()
# ---------------------------------------------------------------------------
# Final metadata save retry integration tests
# ---------------------------------------------------------------------------
class FinalMetadataSaveRetryTests(TestCase):
"""The final recording metadata save must retry on transient DB errors."""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=95, name="Retry Test Channel"
)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_metadata_save_uses_db_retry(self, _ws):
"""Verify recording metadata is saved via _db_retry (retries on OperationalError)."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
# Directly call _db_retry to save metadata as run_recording does
cp = rec.custom_properties.copy()
cp["status"] = "completed"
cp["ended_at"] = str(now)
cp["bytes_written"] = 1024
def _save():
rec.custom_properties = cp
rec.save(update_fields=["custom_properties"])
_db_retry(_save, max_retries=3, base_interval=1, label="test save")
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], "completed")
self.assertEqual(rec.custom_properties["bytes_written"], 1024)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_metadata_survives_transient_save_failure(self, _ws):
"""Simulate OperationalError on first save, success on retry."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
cp = {"status": "completed", "bytes_written": 2048}
call_count = {"n": 0}
_real_save = rec.save
def patched_save(**kwargs):
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("connection reset by peer")
return _real_save(**kwargs)
with patch.object(rec, "save", side_effect=patched_save):
with patch("apps.channels.tasks.time.sleep"):
with patch("apps.channels.tasks.close_old_connections"):
def _save():
rec.custom_properties = cp
rec.save(update_fields=["custom_properties"])
_db_retry(_save, max_retries=3, base_interval=1, label="test")
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], "completed")
# ---------------------------------------------------------------------------
# Initial connection retry tests
# ---------------------------------------------------------------------------
class InitialConnectionRetryTests(TestCase):
"""Verify that the DVR task's reconnection logic retries the same
base URL before falling back to the next candidate."""
def test_reconnect_max_constant_exists_in_run_recording(self):
"""run_recording must define a max-reconnect limit to prevent
infinite retries on the same broken base URL."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
# The reconnection counter pattern must be present
self.assertIn("reconnect", source.lower(),
"run_recording must contain reconnection logic")
# ---------------------------------------------------------------------------
# recover_recordings_on_startup retry tests
# ---------------------------------------------------------------------------
class RecoveryRetryTests(TestCase):
"""DB operations in recover_recordings_on_startup must use _db_retry."""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=97, name="Recovery Retry Channel"
)
@patch("apps.channels.tasks.run_recording.apply_async")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_recovery_save_retries_on_operational_error(self, _ws, mock_async):
"""Recovery status update uses _db_retry — survives one OperationalError."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=30),
end_time=now + timedelta(minutes=30),
custom_properties={},
)
# Simulate what recovery does: mark interrupted, then save with retry
cp = rec.custom_properties or {}
cp["status"] = "interrupted"
cp["interrupted_reason"] = "server_restarted"
rec.custom_properties = cp
call_count = {"n": 0}
_real_save = Recording.save
def patched_save(self_rec, **kwargs):
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("db temporarily unavailable")
return _real_save(self_rec, **kwargs)
with patch.object(Recording, "save", patched_save):
with patch("apps.channels.tasks.time.sleep"):
with patch("apps.channels.tasks.close_old_connections"):
_db_retry(
lambda: rec.save(update_fields=["custom_properties"]),
max_retries=3,
label="test recovery",
)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("status"), "interrupted")
self.assertEqual(rec.custom_properties.get("interrupted_reason"), "server_restarted")
def test_db_retry_fetches_recording_list(self):
"""_db_retry correctly returns query results for recording list fetch."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=30),
end_time=now + timedelta(minutes=30),
custom_properties={},
)
result = _db_retry(
lambda: list(Recording.objects.filter(
start_time__lte=now, end_time__gt=now
)),
label="test query",
)
self.assertGreaterEqual(len(result), 1)
ids = [r.id for r in result]
self.assertIn(rec.id, ids)
@@ -0,0 +1,59 @@
import os
from django.test import SimpleTestCase
from unittest.mock import patch
from apps.channels.tasks import build_dvr_candidates
class DVRPortResolutionTests(SimpleTestCase):
"""
Tests that DVR recording candidate URLs respect the DISPATCHARR_PORT
environment variable instead of hardcoding port 9191.
"""
@patch.dict(os.environ, {'REDIS_HOST': 'redis'}, clear=True)
def test_default_port_uses_9191(self):
"""Without DISPATCHARR_PORT set, candidates default to 9191."""
candidates = build_dvr_candidates()
self.assertIn('http://web:9191', candidates)
self.assertIn('http://localhost:9191', candidates)
@patch.dict(os.environ, {'DISPATCHARR_PORT': '8080', 'REDIS_HOST': 'redis'}, clear=True)
def test_custom_port_reflected_in_candidates(self):
"""DISPATCHARR_PORT=8080 replaces all hardcoded 9191 references."""
candidates = build_dvr_candidates()
self.assertIn('http://web:8080', candidates)
self.assertIn('http://localhost:8080', candidates)
self.assertNotIn('http://web:9191', candidates)
self.assertNotIn('http://localhost:9191', candidates)
@patch.dict(os.environ, {
'DISPATCHARR_PORT': '7777',
'DISPATCHARR_ENV': 'dev',
'REDIS_HOST': 'redis',
}, clear=True)
def test_dev_mode_includes_5656_and_custom_port(self):
"""Dev mode includes both uwsgi internal port (5656) and custom port."""
candidates = build_dvr_candidates()
self.assertIn('http://127.0.0.1:5656', candidates)
self.assertIn('http://127.0.0.1:7777', candidates)
@patch.dict(os.environ, {
'DISPATCHARR_INTERNAL_TS_BASE_URL': 'http://custom:1234',
'REDIS_HOST': 'redis',
}, clear=True)
def test_explicit_override_is_first(self):
"""DISPATCHARR_INTERNAL_TS_BASE_URL should be the first candidate."""
candidates = build_dvr_candidates()
self.assertEqual(candidates[0], 'http://custom:1234')
@patch.dict(os.environ, {
'DISPATCHARR_PORT': '3000',
'DISPATCHARR_INTERNAL_API_BASE': 'http://myhost:4000',
'REDIS_HOST': 'redis',
}, clear=True)
def test_internal_api_base_overrides_web_fallback(self):
"""DISPATCHARR_INTERNAL_API_BASE replaces the http://web:{port} default."""
candidates = build_dvr_candidates()
self.assertIn('http://myhost:4000', candidates)
self.assertNotIn('http://web:3000', candidates)
+180
View File
@@ -0,0 +1,180 @@
"""Tests for the _match_epg_program_by_timeslot() helper in tasks.py.
Covers:
- Exact time-slot match returns program dict
- 80% overlap threshold: at boundary, above, and below
- Multiple overlapping programs: dominant vs. evenly split
- Edge cases: None inputs, zero-duration recording, no EPG data
- Returned dict structure (id, title, sub_title, description)
"""
from datetime import timedelta
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel
from apps.epg.models import EPGSource, EPGData, ProgramData
from apps.channels.tasks import _match_epg_program_by_timeslot
class EpgMatchingSetupMixin:
"""Shared setup for EPG matching tests."""
def setUp(self):
self.source = EPGSource.objects.create(name="Test Source")
self.epg = EPGData.objects.create(
tvg_id="test.channel", name="Test Channel EPG", epg_source=self.source,
)
self.channel = Channel.objects.create(
channel_number=50, name="EPG Match Channel", epg_data=self.epg,
)
self.base = timezone.now().replace(second=0, microsecond=0)
def _prog(self, offset_min, duration_min, title="Test Show", **kwargs):
"""Create a ProgramData starting offset_min from self.base."""
start = self.base + timedelta(minutes=offset_min)
end = start + timedelta(minutes=duration_min)
return ProgramData.objects.create(
epg=self.epg, start_time=start, end_time=end, title=title, **kwargs,
)
class ExactMatchTests(EpgMatchingSetupMixin, TestCase):
"""Recording window exactly matches an EPG program."""
def test_exact_match_returns_program_dict(self):
prog = self._prog(0, 60, title="News at 9", sub_title="Top Stories",
description="Evening news broadcast")
result = _match_epg_program_by_timeslot(
self.epg, prog.start_time, prog.end_time,
)
self.assertIsNotNone(result)
self.assertEqual(result["id"], prog.id)
self.assertEqual(result["title"], "News at 9")
self.assertEqual(result["sub_title"], "Top Stories")
self.assertEqual(result["description"], "Evening news broadcast")
def test_missing_optional_fields_returned_as_empty_strings(self):
prog = self._prog(0, 30, title="Minimal Show")
result = _match_epg_program_by_timeslot(
self.epg, prog.start_time, prog.end_time,
)
self.assertIsNotNone(result)
self.assertEqual(result["sub_title"], "")
self.assertEqual(result["description"], "")
class OverlapThresholdTests(EpgMatchingSetupMixin, TestCase):
"""80% overlap threshold boundary tests."""
def test_exactly_80_percent_overlap_returns_match(self):
"""Program covers exactly 80% of the recording window."""
# Program: 0-60min, Recording: 0-75min → overlap = 60/75 = 80%
prog = self._prog(0, 60, title="Borderline Show")
rec_start = self.base
rec_end = self.base + timedelta(minutes=75)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Borderline Show")
def test_below_80_percent_returns_none(self):
"""Program covers 79% of the recording — below threshold."""
# Program: 0-60min, Recording: 0-76min → overlap = 60/76 ≈ 78.9%
prog = self._prog(0, 60, title="Too Short")
rec_start = self.base
rec_end = self.base + timedelta(minutes=76)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNone(result)
def test_above_80_percent_returns_match(self):
"""Program covers 90% of the recording."""
# Program: 0-60min, Recording: 0-66min → overlap = 60/66 ≈ 90.9%
prog = self._prog(0, 60, title="Good Match")
rec_start = self.base
rec_end = self.base + timedelta(minutes=66)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Good Match")
class MultipleProgramTests(EpgMatchingSetupMixin, TestCase):
"""Recording spans multiple EPG programs."""
def test_dominant_program_returned(self):
"""Recording spans 2 programs; one covers 85%, the other 15%."""
# Show A: 0-60min, Show B: 60-120min
# Recording: 9-69min → A overlap=51/60=85%, B overlap=9/60=15%
self._prog(0, 60, title="Show A")
self._prog(60, 60, title="Show B")
rec_start = self.base + timedelta(minutes=9)
rec_end = self.base + timedelta(minutes=69)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Show A")
def test_evenly_split_returns_none(self):
"""Recording spans 2 equal programs — neither reaches 80%."""
# Show A: 0-60min, Show B: 60-120min
# Recording: 30-90min → each covers 50%
self._prog(0, 60, title="Show A")
self._prog(60, 60, title="Show B")
rec_start = self.base + timedelta(minutes=30)
rec_end = self.base + timedelta(minutes=90)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNone(result)
def test_three_programs_one_dominant(self):
"""Recording spans 3 programs; middle one is dominant."""
# A: 0-30min, B: 30-90min, C: 90-120min
# Recording: 25-95min (70min window) → B overlap=60/70≈85.7%
self._prog(0, 30, title="Show A")
self._prog(30, 60, title="Show B")
self._prog(90, 30, title="Show C")
rec_start = self.base + timedelta(minutes=25)
rec_end = self.base + timedelta(minutes=95)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Show B")
class EdgeCaseTests(EpgMatchingSetupMixin, TestCase):
"""Edge cases and error handling."""
def test_none_epg_data_returns_none(self):
result = _match_epg_program_by_timeslot(None, self.base, self.base + timedelta(hours=1))
self.assertIsNone(result)
def test_none_start_time_returns_none(self):
result = _match_epg_program_by_timeslot(self.epg, None, self.base + timedelta(hours=1))
self.assertIsNone(result)
def test_none_end_time_returns_none(self):
result = _match_epg_program_by_timeslot(self.epg, self.base, None)
self.assertIsNone(result)
def test_zero_duration_returns_none(self):
"""Recording with start == end should return None."""
result = _match_epg_program_by_timeslot(self.epg, self.base, self.base)
self.assertIsNone(result)
def test_negative_duration_returns_none(self):
"""Recording with end before start should return None."""
result = _match_epg_program_by_timeslot(
self.epg, self.base + timedelta(hours=1), self.base,
)
self.assertIsNone(result)
def test_no_overlapping_programs_returns_none(self):
"""No EPG programs in the recording window."""
self._prog(0, 60, title="Earlier Show")
rec_start = self.base + timedelta(hours=5)
rec_end = rec_start + timedelta(hours=1)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNone(result)
def test_empty_epg_no_programs_returns_none(self):
"""EPGData exists but has no programs."""
result = _match_epg_program_by_timeslot(
self.epg, self.base, self.base + timedelta(hours=1),
)
self.assertIsNone(result)
@@ -0,0 +1,235 @@
"""Tests for the Extend In-Progress Recording feature.
Covers:
- extend() API endpoint (happy path and validation)
- pre_save signal guard: end_time change must NOT revoke a live recording
- pre_save signal guard: end_time change MUST still revoke an upcoming recording
- TOCTOU edge cases (extend on a completed/stopped/nonexistent recording)
"""
from datetime import timedelta
from unittest.mock import MagicMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording
from apps.channels.api_views import RecordingViewSet
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="extend_test_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
# ---------------------------------------------------------------------------
# Extend endpoint tests
# ---------------------------------------------------------------------------
class ExtendEndpointTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/extend/"""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=88, name="Extend Test Channel"
)
self.user = _make_admin()
self.factory = APIRequestFactory()
def _extend(self, rec, extra_minutes):
request = self.factory.post(
f"/api/channels/recordings/{rec.id}/extend/",
{"extra_minutes": extra_minutes},
format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "extend"})
return view(request, pk=rec.id)
def _make_rec(self, status="recording"):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": status},
)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_updates_end_time_in_db(self, _ws):
rec = self._make_rec()
original_end = rec.end_time
response = self._extend(rec, 30)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data.get("success"))
rec.refresh_from_db()
expected = original_end + timedelta(minutes=30)
delta = abs((rec.end_time - expected).total_seconds())
self.assertLess(delta, 1, "end_time was not extended by the correct amount")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_stacks_multiple_extensions(self, _ws):
"""Calling extend() twice adds both increments."""
rec = self._make_rec()
original_end = rec.end_time
self._extend(rec, 15)
self._extend(rec, 30)
rec.refresh_from_db()
expected = original_end + timedelta(minutes=45)
delta = abs((rec.end_time - expected).total_seconds())
self.assertLess(delta, 1)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_does_not_clear_task_id(self, _ws):
"""The running Celery task must survive the DB save."""
rec = self._make_rec()
rec.task_id = "dvr-recording-999"
rec.save(update_fields=["task_id"])
self._extend(rec, 30)
rec.refresh_from_db()
self.assertEqual(rec.task_id, "dvr-recording-999")
def test_extend_returns_400_if_finished(self):
"""Cannot extend a completed, stopped, or interrupted recording."""
for bad_status in ("completed", "stopped", "interrupted"):
with self.subTest(status=bad_status):
rec = self._make_rec(status=bad_status)
response = self._extend(rec, 30)
self.assertEqual(response.status_code, 400)
self.assertFalse(response.data.get("success"))
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_succeeds_before_task_sets_status(self, _ws):
"""Extend must work when status is empty (task hasn't started yet)."""
rec = self._make_rec(status="")
response = self._extend(rec, 15)
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
expected = rec.end_time # already extended
self.assertTrue(response.data.get("success"))
@patch("apps.channels.signals.revoke_task")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_bypasses_signals_no_revoke(self, _ws, mock_revoke):
"""Extend uses .update() to bypass pre_save — revoke_task must never fire."""
rec = self._make_rec(status="")
rec.task_id = "dvr-recording-500"
rec.save(update_fields=["task_id"])
self._extend(rec, 15)
self._extend(rec, 30)
mock_revoke.assert_not_called()
rec.refresh_from_db()
self.assertEqual(rec.task_id, "dvr-recording-500")
def test_extend_returns_400_for_zero_minutes(self):
response = self._extend(self._make_rec(), 0)
self.assertEqual(response.status_code, 400)
def test_extend_returns_400_for_negative_minutes(self):
response = self._extend(self._make_rec(), -15)
self.assertEqual(response.status_code, 400)
def test_extend_returns_400_for_non_numeric_minutes(self):
rec = self._make_rec()
request = self.factory.post(
f"/api/channels/recordings/{rec.id}/extend/",
{"extra_minutes": "lots"},
format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "extend"})
response = view(request, pk=rec.id)
self.assertEqual(response.status_code, 400)
def test_extend_returns_404_for_nonexistent_recording(self):
request = self.factory.post(
"/api/channels/recordings/999999/extend/",
{"extra_minutes": 30},
format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "extend"})
response = view(request, pk=999999)
self.assertEqual(response.status_code, 404)
# ---------------------------------------------------------------------------
# pre_save signal guard tests
# ---------------------------------------------------------------------------
class PreSaveExtendGuardTests(TestCase):
"""The pre_save signal must NOT revoke a live recording when end_time changes,
but MUST still revoke a scheduled (upcoming) recording as before."""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=77, name="Signal Guard Channel"
)
def _make_rec(self, status="", task_id="dvr-recording-42"):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now + timedelta(hours=1),
end_time=now + timedelta(hours=2),
task_id=task_id,
custom_properties={"status": status} if status else {},
)
@patch("apps.channels.signals.revoke_task")
def test_end_time_change_does_not_revoke_live_recording(self, mock_revoke):
"""When status='recording', extending end_time must not call revoke_task."""
rec = self._make_rec(status="recording", task_id="dvr-recording-42")
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
mock_revoke.assert_not_called()
@patch("apps.channels.signals.revoke_task")
def test_task_id_preserved_after_extend_on_live_recording(self, mock_revoke):
"""task_id must not be cleared for a live recording's end_time change."""
rec = self._make_rec(status="recording", task_id="dvr-recording-42")
original_task_id = rec.task_id
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
rec.refresh_from_db()
self.assertEqual(rec.task_id, original_task_id)
@patch("apps.channels.signals.revoke_task")
def test_end_time_change_still_revokes_upcoming_recording(self, mock_revoke):
"""The guard must NOT apply to upcoming recordings — existing behavior preserved."""
rec = self._make_rec(status="", task_id="dvr-recording-77")
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
mock_revoke.assert_called_once_with("dvr-recording-77")
@patch("apps.channels.signals.revoke_task")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_pre_save_guard_reads_db_status_not_memory_status(self, _ws, mock_revoke):
"""pre_save reads status from DB (old object), not from the instance being saved."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
task_id="dvr-recording-66",
custom_properties={"status": "recording"},
)
# Simulate: DB status changes to 'completed' behind the instance's back
Recording.objects.filter(pk=rec.pk).update(
custom_properties={"status": "completed"}
)
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
# revoke_task should be called because DB status is "completed", not "recording"
mock_revoke.assert_called_once_with("dvr-recording-66")
@@ -0,0 +1,379 @@
"""Tests for recording metadata endpoints and logo proxy negative cache.
Covers:
- update_metadata endpoint: title/description, user_edited flag, validation
- refresh_artwork endpoint: returns immediately, background thread behavior
- Logo proxy negative cache: cache hit/miss, expiry, eviction, success clears
"""
import time as time_mod
from datetime import timedelta
from unittest.mock import MagicMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording, Logo
from apps.channels.api_views import RecordingViewSet, LogoViewSet
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="metadata_test_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
# ---------------------------------------------------------------------------
# update_metadata endpoint
# ---------------------------------------------------------------------------
class UpdateMetadataTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/update-metadata/"""
def setUp(self):
self.channel = Channel.objects.create(channel_number=70, name="Meta Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _update(self, rec, data):
request = self.factory.post(
f"/api/channels/recordings/{rec.id}/update-metadata/",
data, format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "update_metadata"})
return view(request, pk=rec.id)
def _make_rec(self, custom_properties=None):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties=custom_properties or {},
)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_update_title_only(self, _ws):
rec = self._make_rec()
response = self._update(rec, {"title": "My Show"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
program = rec.custom_properties["program"]
self.assertEqual(program["title"], "My Show")
self.assertTrue(program["user_edited"])
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_update_description_only(self, _ws):
rec = self._make_rec({"program": {"title": "Existing Title"}})
response = self._update(rec, {"description": "A great episode"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
program = rec.custom_properties["program"]
self.assertEqual(program["description"], "A great episode")
self.assertEqual(program["title"], "Existing Title")
self.assertTrue(program["user_edited"])
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_update_both_fields(self, _ws):
rec = self._make_rec()
response = self._update(rec, {"title": "New Title", "description": "New Desc"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
program = rec.custom_properties["program"]
self.assertEqual(program["title"], "New Title")
self.assertEqual(program["description"], "New Desc")
self.assertTrue(program["user_edited"])
def test_no_fields_returns_400(self):
rec = self._make_rec()
response = self._update(rec, {})
self.assertEqual(response.status_code, 400)
self.assertFalse(response.data.get("success"))
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_whitespace_trimmed(self, _ws):
rec = self._make_rec()
response = self._update(rec, {"title": " Padded Title "})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["program"]["title"], "Padded Title")
def test_whitespace_only_title_returns_400(self):
"""Whitespace-only title and description should be rejected."""
rec = self._make_rec()
response = self._update(rec, {"title": " ", "description": " "})
self.assertEqual(response.status_code, 400)
self.assertFalse(response.data.get("success"))
def test_whitespace_only_title_with_valid_description(self):
"""Whitespace-only title is ignored; valid description is accepted."""
rec = self._make_rec({"program": {"title": "Original"}})
with patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None):
response = self._update(rec, {"title": " ", "description": "Valid desc"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
# Title should remain unchanged since the whitespace-only value is not applied
self.assertEqual(rec.custom_properties["program"]["title"], "Original")
self.assertEqual(rec.custom_properties["program"]["description"], "Valid desc")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_creates_program_dict_when_absent(self, _ws):
"""Recording with no program dict gets one created."""
rec = self._make_rec({"status": "completed"})
response = self._update(rec, {"title": "Brand New"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
self.assertIn("program", rec.custom_properties)
self.assertEqual(rec.custom_properties["program"]["title"], "Brand New")
def test_returns_404_for_nonexistent(self):
request = self.factory.post(
"/api/channels/recordings/99999/update-metadata/",
{"title": "Ghost"}, format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "update_metadata"})
self.assertEqual(view(request, pk=99999).status_code, 404)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_sends_websocket_event(self, mock_ws):
rec = self._make_rec()
self._update(rec, {"title": "WS Test"})
mock_ws.assert_called_once()
payload = mock_ws.call_args[0][2]
self.assertEqual(payload["type"], "recording_updated")
self.assertEqual(payload["recording_id"], rec.id)
@patch("core.utils.send_websocket_update", side_effect=Exception("WS down"))
def test_ws_failure_does_not_fail_request(self, _ws):
"""WebSocket errors are silenced — the save still succeeds."""
rec = self._make_rec()
response = self._update(rec, {"title": "Resilient"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["program"]["title"], "Resilient")
# ---------------------------------------------------------------------------
# refresh_artwork endpoint
# ---------------------------------------------------------------------------
class RefreshArtworkTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/refresh-artwork/"""
def setUp(self):
self.channel = Channel.objects.create(channel_number=71, name="Artwork Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _refresh(self, rec):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
return view(request, pk=rec.id)
def _make_rec(self, custom_properties=None):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties=custom_properties or {},
)
@patch("threading.Thread")
def test_returns_200_immediately(self, mock_thread):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
response = self._refresh(rec)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data.get("success"))
@patch("threading.Thread")
def test_spawns_background_thread(self, mock_thread):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
self._refresh(rec)
mock_thread.assert_called_once()
self.assertTrue(mock_thread.call_args[1].get("daemon", False))
mock_thread.return_value.start.assert_called_once()
def test_returns_404_for_nonexistent(self):
request = self.factory.post("/api/channels/recordings/99999/refresh-artwork/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
self.assertEqual(view(request, pk=99999).status_code, 404)
@patch("django.db.close_old_connections")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_no_downgrade_to_channel_logo(self, _ws, _close):
"""When the pipeline returns the channel's own logo, existing poster is preserved."""
logo = Logo.objects.create(name="Channel Logo", url="https://example.com/ch.png")
self.channel.logo = logo
self.channel.save()
rec = self._make_rec({
"poster_logo_id": 999, # existing real poster
"poster_url": "https://tmdb.com/real-poster.jpg",
})
with patch("apps.channels.tasks._resolve_poster_for_program",
return_value=(logo.id, None)):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
force_authenticate(request, user=self.user)
# Run synchronously by intercepting the thread
captured_fn = None
def capture_thread(*args, **kwargs):
nonlocal captured_fn
captured_fn = kwargs.get("target") or args[0]
mock = MagicMock()
mock.start = lambda: captured_fn(*kwargs.get("args", ()))
return mock
with patch("threading.Thread", side_effect=capture_thread):
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
view(request, pk=rec.id)
rec.refresh_from_db()
# Existing poster should be preserved — not downgraded to channel logo
self.assertEqual(rec.custom_properties.get("poster_logo_id"), 999)
self.assertEqual(rec.custom_properties.get("poster_url"), "https://tmdb.com/real-poster.jpg")
@patch("django.db.close_old_connections")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_upgrade_from_no_poster(self, _ws, _close):
"""When a recording has no poster and the pipeline finds one, it gets updated."""
rec = self._make_rec({"program": {"title": "Some Show", "id": 42}})
with patch("apps.channels.tasks._resolve_poster_for_program",
return_value=(555, "https://tmdb.com/new-poster.jpg")):
captured_fn = None
def capture_thread(*args, **kwargs):
nonlocal captured_fn
captured_fn = kwargs.get("target") or args[0]
mock = MagicMock()
mock.start = lambda: captured_fn(*kwargs.get("args", ()))
return mock
with patch("threading.Thread", side_effect=capture_thread):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
view(request, pk=rec.id)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("poster_logo_id"), 555)
self.assertEqual(rec.custom_properties.get("poster_url"), "https://tmdb.com/new-poster.jpg")
# ---------------------------------------------------------------------------
# Logo proxy negative cache
# ---------------------------------------------------------------------------
class LogoNegativeCacheTests(TestCase):
"""Tests for the _logo_fetch_failures negative cache in LogoViewSet.cache()."""
def setUp(self):
from apps.channels import api_views
self._failures = api_views._logo_fetch_failures
self._failures.clear()
self.factory = APIRequestFactory()
self.user = _make_admin()
def _fetch_logo(self, logo):
request = self.factory.get(f"/api/channels/logos/{logo.id}/cache/")
force_authenticate(request, user=self.user)
view = LogoViewSet.as_view({"get": "cache"})
return view(request, pk=logo.id)
def test_failed_url_cached_on_non_200(self):
"""Non-200 response adds URL to negative cache."""
logo = Logo.objects.create(name="Dead Logo", url="https://dead-cdn.com/logo.png")
mock_resp = MagicMock(status_code=404)
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 404)
self.assertIn("https://dead-cdn.com/logo.png", self._failures)
def test_cached_failure_returns_404_immediately(self):
"""Subsequent request for a cached-failed URL returns 404 without making a request."""
logo = Logo.objects.create(name="Cached Fail", url="https://cached-fail.com/logo.png")
self._failures["https://cached-fail.com/logo.png"] = time_mod.monotonic() + 300
with patch("apps.channels.api_views.requests.get") as mock_get:
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 404)
mock_get.assert_not_called()
def test_expired_cache_entry_allows_retry(self):
"""After TTL expires, a new request is made."""
logo = Logo.objects.create(name="Expired", url="https://expired.com/logo.png")
self._failures["https://expired.com/logo.png"] = time_mod.monotonic() - 1 # already expired
mock_resp = MagicMock(status_code=200)
mock_resp.headers = {"Content-Type": "image/png"}
mock_resp.iter_content = MagicMock(return_value=[b"img"])
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 200)
def test_success_clears_previous_failure(self):
"""A successful fetch removes the URL from the failure cache."""
url = "https://recovered.com/logo.png"
logo = Logo.objects.create(name="Recovered", url=url)
self._failures[url] = time_mod.monotonic() - 1 # expired
mock_resp = MagicMock(status_code=200)
mock_resp.headers = {"Content-Type": "image/png"}
mock_resp.iter_content = MagicMock(return_value=[b"img"])
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
self._fetch_logo(logo)
self.assertNotIn(url, self._failures)
def test_request_exception_cached(self):
"""Network errors are cached the same as non-200 responses."""
import requests
logo = Logo.objects.create(name="Timeout", url="https://timeout.com/logo.png")
with patch("apps.channels.api_views.requests.get", side_effect=requests.Timeout("timed out")), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 404)
self.assertIn("https://timeout.com/logo.png", self._failures)
def test_eviction_when_cache_exceeds_256(self):
"""Stale entries are evicted when the cache grows past 256."""
now = time_mod.monotonic()
# Fill with 257 expired entries
for i in range(257):
self._failures[f"https://old-{i}.com/x.png"] = now - 1 # already expired
logo = Logo.objects.create(name="Trigger", url="https://trigger-evict.com/logo.png")
import requests
with patch("apps.channels.api_views.requests.get", side_effect=requests.ConnectionError("fail")), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
self._fetch_logo(logo)
# Expired entries should be evicted
old_entries = [k for k in self._failures if k.startswith("https://old-")]
self.assertEqual(len(old_entries), 0)
# New failure entry should exist
self.assertIn("https://trigger-evict.com/logo.png", self._failures)
@@ -0,0 +1,524 @@
"""Tests for recent DVR fixes.
Covers:
1. Collision avoidance: _build_output_paths checks both .mkv and .ts files
2. Logo guard: _resolve_poster_for_program skips external APIs when title ≈ channel name
3. Recording status lifecycle: status transitions visible via API
4. Concat flags: error-tolerant ffmpeg flags used for segment concatenation
5. Recovery skip-list: "recording" status NOT in terminal skip list
"""
import os
from datetime import timedelta
from unittest.mock import MagicMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="dvr_fixes_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
def _make_channel(name="Test Channel", number=100):
return Channel.objects.create(channel_number=number, name=name)
def _make_recording(channel, **overrides):
now = timezone.now()
defaults = {
"channel": channel,
"start_time": now - timedelta(hours=1),
"end_time": now + timedelta(hours=1),
"custom_properties": {},
}
defaults.update(overrides)
return Recording.objects.create(**defaults)
# =========================================================================
# 1. Collision avoidance — _build_output_paths
# =========================================================================
class CollisionAvoidanceTests(TestCase):
"""_build_output_paths must increment the filename counter when
EITHER the .mkv OR the .ts file already exists with size > 0."""
def _call(self, channel, program, start, end):
from apps.channels.tasks import _build_output_paths
return _build_output_paths(channel, program, start, end)
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_no_collision_when_nothing_exists(self, _tv, _fb):
"""Fresh path — no files exist, counter stays at 1."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
raise OSError("No such file")
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
# Should NOT have a _2 suffix
self.assertNotIn("_2", final)
self.assertTrue(final.endswith(".mkv"))
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_collision_when_ts_exists_but_mkv_is_zero_bytes(self, _tv, _fb):
"""Pre-restart scenario: MKV is 0-byte placeholder, TS has real data.
The old code only checked MKV size, so it would reuse the path.
The fix also checks TS, so it must increment."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
if "_2" in path:
raise OSError("No such file")
result = MagicMock()
if path.endswith('.mkv'):
result.st_size = 0 # MKV is 0-byte placeholder
elif path.endswith('.ts'):
result.st_size = 5000000 # TS has real data from pre-restart
else:
result.st_size = 0
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
# Must have incremented to _2
self.assertIn("_2", final, "Should increment counter when TS file has data")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_collision_when_mkv_has_data(self, _tv, _fb):
"""Standard collision: MKV file has data, should increment."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
if "_2" in path:
raise OSError("No such file")
result = MagicMock()
if path.endswith('.mkv'):
result.st_size = 1000000 # MKV has data
else:
result.st_size = 0
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
self.assertIn("_2", final, "Should increment counter when MKV file has data")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_no_collision_when_both_zero_bytes(self, _tv, _fb):
"""Both MKV and TS exist but are 0 bytes — no collision."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
result = MagicMock()
result.st_size = 0 # All files empty
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
self.assertNotIn("_2", final, "Should NOT increment when all files are empty")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_collision_increments_to_3_when_2_also_occupied(self, _tv, _fb):
"""When both base and _2 are occupied, should go to _3."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
if "_3" in path:
raise OSError("No such file")
result = MagicMock()
if path.endswith('.ts'):
result.st_size = 5000000
else:
result.st_size = 0
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
self.assertIn("_3", final, "Should increment to _3 when base and _2 are occupied")
# =========================================================================
# 2. Logo guard — _resolve_poster_for_program
# =========================================================================
class LogoGuardTests(TestCase):
"""When the program title matches the channel name, external API
searches (VOD, TMDB, OMDb, TVMaze, iTunes) must be skipped."""
def _call(self, channel_name, program, channel_logo_id=None):
from apps.channels.tasks import _resolve_poster_for_program
return _resolve_poster_for_program(channel_name, program, channel_logo_id)
@patch("apps.channels.tasks.requests.get")
def test_channel_name_as_title_skips_external_apis(self, mock_get):
"""Title = 'USA A&E SD*', channel = 'USA A&E SD*' → no external calls."""
program = {"title": "USA A&E SD*"}
logo_id, url = self._call("USA A&E SD*", program, channel_logo_id=42)
# Should NOT have called any external APIs
mock_get.assert_not_called()
# Should fall back to channel logo
self.assertEqual(logo_id, 42)
self.assertIsNone(url)
@patch("apps.channels.tasks.requests.get")
def test_channel_name_normalized_match(self, mock_get):
"""Title = 'fox news', channel = 'FOX-News*' → normalized match, skip APIs."""
program = {"title": "fox news"}
logo_id, url = self._call("FOX-News*", program, channel_logo_id=99)
mock_get.assert_not_called()
self.assertEqual(logo_id, 99)
@patch("apps.channels.tasks.requests.get")
def test_real_title_still_searched(self, mock_get):
"""Title = 'Breaking Bad' on channel 'AMC' → should try external APIs."""
# Mock TVMaze returning a result
mock_resp = MagicMock(ok=True, status_code=200)
mock_resp.json.return_value = {
"image": {"original": "https://tvmaze.com/breaking-bad.jpg"}
}
mock_get.return_value = mock_resp
program = {"title": "Breaking Bad"}
logo_id, url = self._call("AMC", program)
# Should have made at least one external API call
self.assertTrue(mock_get.called, "Should search external APIs for real titles")
self.assertIsNotNone(url)
@patch("apps.channels.tasks.requests.get")
def test_no_title_skips_to_channel_logo(self, mock_get):
"""No title at all → falls through to channel logo, no API calls."""
program = {}
logo_id, url = self._call("SomeChannel", program, channel_logo_id=55)
mock_get.assert_not_called()
self.assertEqual(logo_id, 55)
@patch("apps.channels.tasks.requests.get")
def test_epg_image_still_used_even_when_title_is_channel_name(self, mock_get):
"""Even when title = channel name, Stage 1 (EPG images) should still work."""
from apps.epg.models import ProgramData, EPGSource, EPGData
# Create an EPG source + EPGData entry + program with an icon URL
epg_source = EPGSource.objects.create(source_type="xmltv", name="Test EPG")
epg_data = EPGData.objects.create(tvg_id="test.ch", epg_source=epg_source)
prog = ProgramData.objects.create(
epg=epg_data,
title="Test Channel HD",
start_time=timezone.now() - timedelta(hours=1),
end_time=timezone.now() + timedelta(hours=1),
custom_properties={"icon": "https://epg-cdn.com/test-icon.png"},
)
program = {"title": "Test Channel HD", "id": prog.id}
# Mock _validate_url to return True for the icon URL
with patch("apps.channels.tasks._validate_url", return_value=True):
logo_id, url = self._call("Test Channel HD", program, channel_logo_id=10)
# EPG icon should still be used (Stage 1 doesn't depend on title guard)
self.assertEqual(url, "https://epg-cdn.com/test-icon.png")
mock_get.assert_not_called()
# =========================================================================
# 3. Recording status lifecycle via API
# =========================================================================
class RecordingStatusLifecycleTests(TestCase):
"""Verify recording status transitions and that terminal recordings
are properly filterable (supports the red-dot fix in guideUtils)."""
def setUp(self):
self.channel = _make_channel("Status Test Channel", 200)
self.user = _make_admin()
self.factory = APIRequestFactory()
def _list_recordings(self):
from apps.channels.api_views import RecordingViewSet
request = self.factory.get("/api/channels/recordings/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"get": "list"})
return view(request)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_stopped_recording_has_terminal_status(self, _ws):
"""After stop, custom_properties.status = 'stopped'."""
from apps.channels.api_views import RecordingViewSet
rec = _make_recording(self.channel, custom_properties={
"status": "recording",
"program": {"id": 1, "title": "Live Show"},
})
request = self.factory.post(f"/api/channels/recordings/{rec.id}/stop/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "stop"})
with patch("apps.channels.signals.revoke_task"):
response = view(request, pk=rec.id)
self.assertIn(response.status_code, [200, 204])
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("status"), "stopped")
def test_listing_includes_status_in_custom_properties(self):
"""API listing returns custom_properties with status field."""
_make_recording(self.channel, custom_properties={
"status": "recording",
"program": {"id": 1, "title": "Recording Show"},
})
_make_recording(self.channel, custom_properties={
"status": "stopped",
"program": {"id": 2, "title": "Stopped Show"},
})
response = self._list_recordings()
self.assertEqual(response.status_code, 200)
statuses = [r["custom_properties"].get("status") for r in response.data]
self.assertIn("recording", statuses)
self.assertIn("stopped", statuses)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_delete_recording_removes_from_listing(self, _ws):
"""Deleting a recording removes it from the listing entirely."""
from apps.channels.api_views import RecordingViewSet
rec = _make_recording(self.channel, custom_properties={
"status": "stopped",
"program": {"id": 3, "title": "To Delete"},
})
rec_id = rec.id
request = self.factory.delete(f"/api/channels/recordings/{rec_id}/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"delete": "destroy"})
with patch("apps.channels.signals.revoke_task"):
response = view(request, pk=rec_id)
self.assertIn(response.status_code, [200, 204])
self.assertFalse(Recording.objects.filter(id=rec_id).exists())
# =========================================================================
# 4. Concat flags — error-tolerant ffmpeg
# =========================================================================
class ConcatFlagsTests(TestCase):
"""Verify that the finalize phase uses error-tolerant ffmpeg flags
when concatenating pre-restart segments."""
def test_concat_command_includes_error_tolerant_flags(self):
"""Inspect the source code to confirm error-tolerant flags are present.
This is a static analysis test — no ffmpeg execution needed."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
# The concat subprocess.run call must include these flags
self.assertIn("+genpts+igndts+discardcorrupt", source,
"Concat must use +genpts+igndts+discardcorrupt fflags")
self.assertIn("ignore_err", source,
"Concat must use -err_detect ignore_err")
self.assertIn("-f", source)
self.assertIn("concat", source)
def test_concat_goes_directly_to_mkv(self):
"""Concat must produce MKV directly (not intermediate .ts) to
preserve timestamp boundaries and avoid playback freeze at splice."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
# Must contain reset_timestamps for proper segment boundary handling
self.assertIn("reset_timestamps", source,
"Concat must use -reset_timestamps 1 for seamless seeking")
# Must write directly to final_path (MKV), not an intermediate .ts
self.assertIn("_concat_did_remux", source,
"Concat path must set flag to skip separate remux step")
def test_segment_time_metadata_present(self):
"""Verify concat uses -segment_time_metadata for boundary awareness."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
self.assertIn("segment_time_metadata", source,
"Concat must use -segment_time_metadata 1 for segment boundary handling")
# =========================================================================
# 5. Recovery skip-list
# =========================================================================
class RecoverySkipListTests(TestCase):
"""Verify that the recovery function does NOT skip 'recording' status,
since that's the exact status recordings have when the server crashes."""
def test_recording_status_not_in_skip_list(self):
"""Inspect recover_recordings_on_startup to ensure 'recording' is
NOT treated as a terminal/skip state."""
import inspect
from apps.channels.tasks import recover_recordings_on_startup
source = inspect.getsource(recover_recordings_on_startup)
# Find the skip condition line
# It should be: if current_status in ("completed", "stopped"):
# NOT: if current_status in ("completed", "stopped", "recording"):
lines = source.split('\n')
skip_line = None
for line in lines:
if 'current_status in' in line and ('completed' in line or 'stopped' in line):
skip_line = line.strip()
break
self.assertIsNotNone(skip_line, "Should find the skip-list condition")
self.assertNotIn('"recording"', skip_line,
"Skip list must NOT contain 'recording' — "
"that's the status of crashed mid-stream recordings that need recovery")
@patch("core.utils.RedisClient")
@patch("apps.channels.tasks.run_recording")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_recovery_processes_recording_status(self, _ws, mock_run, mock_redis_cls):
"""A recording with status='recording' should be recovered, not skipped."""
mock_redis_conn = MagicMock()
mock_redis_conn.set.return_value = True # Acquire lock
mock_redis_cls.get_client.return_value = mock_redis_conn
channel = _make_channel("Recovery Test", 300)
now = timezone.now()
rec = _make_recording(channel, custom_properties={
"status": "recording",
"program": {"title": "Crashed Show"},
}, end_time=now + timedelta(hours=2))
from apps.channels.tasks import recover_recordings_on_startup
with patch("apps.channels.signals.revoke_task"):
result = recover_recordings_on_startup()
# The recording should have been dispatched for recovery
self.assertTrue(mock_run.apply_async.called,
"Recording with status='recording' should be dispatched for recovery")
@patch("core.utils.RedisClient")
@patch("apps.channels.tasks.run_recording")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_recovery_skips_stopped_recordings(self, _ws, mock_run, mock_redis_cls):
"""A recording with status='stopped' should be skipped by recovery."""
mock_redis_conn = MagicMock()
mock_redis_conn.set.return_value = True
mock_redis_cls.get_client.return_value = mock_redis_conn
channel = _make_channel("Recovery Skip Test", 301)
now = timezone.now()
rec = _make_recording(channel, custom_properties={
"status": "stopped",
"program": {"title": "Finished Show"},
}, end_time=now + timedelta(hours=2))
from apps.channels.tasks import recover_recordings_on_startup
with patch("apps.channels.signals.revoke_task"):
recover_recordings_on_startup()
# Should NOT have dispatched a recovery task
mock_run.apply_async.assert_not_called()
# =========================================================================
# 6. Frontend red-dot filter (guideUtils.mapRecordingsByProgramId)
# =========================================================================
class MapRecordingsByProgramIdTests(TestCase):
"""These test the BACKEND side — confirming that recording status
is preserved in the API response so the frontend can filter on it.
The actual frontend filtering is covered by frontend/src/pages/__tests__/DVR.test.jsx
and the guideUtils code, but we verify the data contract here."""
def test_recording_custom_properties_status_persisted(self):
"""Recording status in custom_properties survives save/load cycle."""
channel = _make_channel("Red Dot Test", 400)
rec = _make_recording(channel, custom_properties={
"status": "stopped",
"program": {"id": 42, "title": "A Show"},
})
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], "stopped")
def test_terminal_statuses_are_well_defined(self):
"""Verify the terminal status set matches what the frontend uses."""
# These are the statuses that should NOT show a red dot in the Guide
terminal = {"stopped", "completed", "interrupted", "failed"}
channel = _make_channel("Terminal Status Test", 410)
# Verify each status is a valid recording status
for status in terminal:
rec = _make_recording(channel, custom_properties={
"status": status,
"program": {"id": 100, "title": "Test"},
})
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], status)
@@ -0,0 +1,585 @@
"""Tests for DVR recording scheduling with ClockedSchedule.
Uses ClockedSchedule instead of apply_async with countdown because Redis
visibility_timeout (default 3600s) causes task redelivery for long countdowns,
leading to duplicate recordings.
"""
from datetime import timedelta
from unittest.mock import patch, MagicMock
from django.test import TestCase
from django.utils import timezone
from django_celery_beat.models import ClockedSchedule, PeriodicTask
from apps.channels.models import Channel, Recording
from apps.channels.signals import (
schedule_recording_task,
revoke_task,
_dvr_task_name,
)
class ScheduleRecordingTaskTests(TestCase):
"""Tests for schedule_recording_task()."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.run_recording")
def test_future_recording_creates_periodic_task(self, mock_run_recording):
"""Recordings in the future create a ClockedSchedule + PeriodicTask."""
future_time = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future_time,
end_time=future_time + timedelta(hours=1),
)
task_id = schedule_recording_task(rec, eta=future_time)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
pt = PeriodicTask.objects.get(name=expected_name)
self.assertTrue(pt.one_off)
self.assertTrue(pt.enabled)
self.assertEqual(pt.task, "apps.channels.tasks.run_recording")
self.assertIsNotNone(pt.clocked)
# apply_async should not have been called
mock_run_recording.apply_async.assert_not_called()
@patch("apps.channels.signals.run_recording")
def test_immediate_recording_creates_periodic_task(self, mock_run_recording):
"""Recordings starting now also use ClockedSchedule for consistency."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now,
end_time=now + timedelta(hours=1),
)
task_id = schedule_recording_task(rec, eta=now)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
self.assertTrue(PeriodicTask.objects.filter(name=expected_name).exists())
@patch("apps.channels.signals.run_recording")
def test_past_start_time_clamps_to_now(self, mock_run_recording):
"""Recordings with past start_time get clamped to now."""
past_time = timezone.now() - timedelta(minutes=5)
rec = Recording.objects.create(
channel=self.channel,
start_time=past_time,
end_time=timezone.now() + timedelta(hours=1),
)
task_id = schedule_recording_task(rec, eta=past_time)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
pt = PeriodicTask.objects.get(name=expected_name)
# Clocked time should be >= now
self.assertGreaterEqual(pt.clocked.clocked_time, past_time)
@patch("apps.channels.signals.run_recording")
def test_reschedule_updates_existing_periodic_task(self, mock_run_recording):
"""Calling schedule_recording_task twice updates the existing PeriodicTask."""
future_time = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future_time,
end_time=future_time + timedelta(hours=1),
)
schedule_recording_task(rec, eta=future_time)
# Reschedule with a different time
new_eta = future_time + timedelta(hours=1)
schedule_recording_task(rec, eta=new_eta)
# Should still be exactly one PeriodicTask
task_name = f"dvr-recording-{rec.id}"
self.assertEqual(PeriodicTask.objects.filter(name=task_name).count(), 1)
@patch("apps.channels.signals.run_recording")
def test_naive_eta_is_made_aware(self, mock_run_recording):
"""A naive (timezone-unaware) eta is made timezone-aware."""
from datetime import datetime
naive_eta = datetime(2030, 6, 15, 14, 0, 0)
rec = Recording.objects.create(
channel=self.channel,
start_time=timezone.now() + timedelta(hours=1),
end_time=timezone.now() + timedelta(hours=2),
)
task_id = schedule_recording_task(rec, eta=naive_eta)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
pt = PeriodicTask.objects.get(name=expected_name)
self.assertTrue(timezone.is_aware(pt.clocked.clocked_time))
class RevokeTaskTests(TestCase):
"""Tests for revoke_task()."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
def test_revoke_deletes_periodic_task_and_clocked_schedule(self):
"""revoke_task deletes the PeriodicTask and orphaned ClockedSchedule."""
eta = timezone.now() + timedelta(hours=5)
clocked = ClockedSchedule.objects.create(clocked_time=eta)
PeriodicTask.objects.create(
name="dvr-recording-10",
task="apps.channels.tasks.run_recording",
clocked=clocked,
one_off=True,
enabled=True,
)
revoke_task("dvr-recording-10")
self.assertFalse(PeriodicTask.objects.filter(name="dvr-recording-10").exists())
self.assertFalse(ClockedSchedule.objects.filter(id=clocked.id).exists())
def test_revoke_keeps_shared_clocked_schedule(self):
"""ClockedSchedule is kept if another PeriodicTask still references it."""
eta = timezone.now() + timedelta(hours=5)
clocked = ClockedSchedule.objects.create(clocked_time=eta)
PeriodicTask.objects.create(
name="dvr-recording-10",
task="apps.channels.tasks.run_recording",
clocked=clocked,
one_off=True,
)
PeriodicTask.objects.create(
name="dvr-recording-11",
task="apps.channels.tasks.run_recording",
clocked=clocked,
one_off=True,
)
revoke_task("dvr-recording-10")
self.assertFalse(PeriodicTask.objects.filter(name="dvr-recording-10").exists())
self.assertTrue(ClockedSchedule.objects.filter(id=clocked.id).exists())
@patch("apps.channels.signals.AsyncResult")
def test_revoke_falls_back_to_async_result_for_legacy_ids(self, mock_async_result):
"""revoke_task falls back to AsyncResult.revoke() for old-style UUIDs."""
revoke_task("550e8400-e29b-41d4-a716-446655440000")
mock_async_result.assert_called_once_with("550e8400-e29b-41d4-a716-446655440000")
mock_async_result.return_value.revoke.assert_called_once()
def test_revoke_none_is_noop(self):
"""revoke_task(None) does nothing."""
revoke_task(None) # Should not raise
def test_revoke_empty_string_is_noop(self):
"""revoke_task('') does nothing."""
revoke_task("") # Should not raise
class DvrTaskNameTests(TestCase):
"""Tests for the naming convention helper."""
def test_task_name_format(self):
self.assertEqual(_dvr_task_name(42), "dvr-recording-42")
def test_task_name_fits_in_charfield(self):
name = _dvr_task_name(999999999)
self.assertLessEqual(len(name), 255)
class SignalIntegrationTests(TestCase):
"""Integration tests for the post_save / post_delete signal handlers."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_save_creates_periodic_task_for_future_recording(self, mock_artwork):
"""Saving a future Recording creates a PeriodicTask via post_save signal."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
)
rec.refresh_from_db()
task_name = f"dvr-recording-{rec.id}"
self.assertEqual(rec.task_id, task_name)
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_delete_removes_periodic_task(self, mock_artwork):
"""Deleting a Recording removes its PeriodicTask."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
)
rec.refresh_from_db()
task_name = rec.task_id
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
rec.delete()
self.assertFalse(PeriodicTask.objects.filter(name=task_name).exists())
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_bulk_delete_cleans_up_all_periodic_tasks(self, mock_artwork):
"""Bulk deleting recordings cleans up all their PeriodicTasks."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec_ids = []
for i in range(5):
rec = Recording.objects.create(
channel=self.channel,
start_time=future + timedelta(hours=i),
end_time=future + timedelta(hours=i + 1),
)
rec_ids.append(rec.id)
for rid in rec_ids:
self.assertTrue(
PeriodicTask.objects.filter(name=f"dvr-recording-{rid}").exists()
)
Recording.objects.filter(channel=self.channel).delete()
self.assertEqual(
PeriodicTask.objects.filter(name__startswith="dvr-recording-").count(), 0
)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_save_schedules_currently_playing_recording(self, mock_artwork):
"""A recording with past start_time but future end_time schedules immediately."""
mock_artwork.apply_async.return_value = MagicMock()
past_start = timezone.now() - timedelta(minutes=30)
future_end = timezone.now() + timedelta(minutes=30)
rec = Recording.objects.create(
channel=self.channel,
start_time=past_start,
end_time=future_end,
)
rec.refresh_from_db()
task_name = f"dvr-recording-{rec.id}"
self.assertEqual(rec.task_id, task_name)
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_save_skips_fully_past_recording(self, mock_artwork):
"""A recording with both start_time and end_time in the past is not scheduled."""
mock_artwork.apply_async.return_value = MagicMock()
past_start = timezone.now() - timedelta(hours=2)
past_end = timezone.now() - timedelta(hours=1)
rec = Recording.objects.create(
channel=self.channel,
start_time=past_start,
end_time=past_end,
)
rec.refresh_from_db()
self.assertIsNone(rec.task_id)
self.assertFalse(
PeriodicTask.objects.filter(name=f"dvr-recording-{rec.id}").exists()
)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_pre_save_revokes_on_time_change(self, mock_artwork):
"""Changing a recording's start_time revokes the old task and creates a new one."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
)
rec.refresh_from_db()
old_task_name = rec.task_id
self.assertTrue(PeriodicTask.objects.filter(name=old_task_name).exists())
# Change the start time — pre_save clears task_id, post_save reschedules
new_future = future + timedelta(hours=3)
rec.start_time = new_future
rec.end_time = new_future + timedelta(hours=1)
rec.save()
rec.refresh_from_db()
# Old PeriodicTask should be deleted; new one should exist
self.assertIsNotNone(rec.task_id)
self.assertTrue(
PeriodicTask.objects.filter(name=f"dvr-recording-{rec.id}").exists()
)
class IdempotencyGuardTests(TestCase):
"""Tests for the idempotency guard in run_recording()."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
@patch("apps.channels.tasks.get_channel_layer")
def test_skips_if_already_recording(self, mock_layer):
"""run_recording returns early if status is already 'recording'."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now,
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording", "started_at": str(now)},
)
from apps.channels.tasks import run_recording as run_rec_task
result = run_rec_task(rec.id, self.channel.id, str(now), str(now + timedelta(hours=1)))
self.assertIsNone(result)
# get_channel_layer should not have been called (returned before)
mock_layer.assert_not_called()
@patch("apps.channels.tasks.get_channel_layer")
def test_skips_if_already_completed(self, mock_layer):
"""run_recording returns early if status is already 'completed'."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=2),
end_time=now - timedelta(hours=1),
custom_properties={"status": "completed"},
)
from apps.channels.tasks import run_recording as run_rec_task
result = run_rec_task(rec.id, self.channel.id, str(rec.start_time), str(rec.end_time))
self.assertIsNone(result)
mock_layer.assert_not_called()
@patch("apps.channels.tasks.get_channel_layer")
def test_skips_if_already_stopped(self, mock_layer):
"""run_recording returns early if status is already 'stopped' (user stopped it early)."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": "stopped", "stopped_at": str(now)},
)
from apps.channels.tasks import run_recording as run_rec_task
result = run_rec_task(rec.id, self.channel.id, str(rec.start_time), str(rec.end_time))
self.assertIsNone(result)
mock_layer.assert_not_called()
class ArtworkPrefetchSignalGuardTests(TestCase):
"""Tests that the post_save signal does not schedule artwork prefetch when
the recording is in an active or terminal state."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_not_scheduled_when_status_recording(self, mock_artwork):
"""post_save must NOT schedule artwork prefetch when status='recording'
to prevent a race that overwrites the running task's status updates."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={"status": "recording"},
)
# Simulate a save that run_recording itself might do mid-recording
rec.custom_properties = {"status": "recording", "file_path": "/data/recordings/test.mkv"}
rec.save(update_fields=["custom_properties"])
# apply_async was not called for the "recording" save
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_not_scheduled_when_status_completed(self, mock_artwork):
"""post_save must NOT schedule artwork prefetch when status='completed'."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={"status": "completed"},
)
rec.custom_properties = {"status": "completed"}
rec.save(update_fields=["custom_properties"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_not_scheduled_when_status_stopped(self, mock_artwork):
"""post_save must NOT schedule artwork prefetch when status='stopped'."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={"status": "stopped"},
)
rec.custom_properties = {"status": "stopped"}
rec.save(update_fields=["custom_properties"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_scheduled_for_new_upcoming_recording(self, mock_artwork):
"""post_save SHOULD schedule artwork prefetch for a newly created upcoming recording."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={}, # no status yet — should trigger prefetch
)
self.assertTrue(mock_artwork.apply_async.called)
class DestroyDvrClientIsolationTests(TestCase):
"""Tests that deleting a recording only stops DVR clients when the
recording is actively streaming — never for completed/upcoming recordings
that could share a channel with an unrelated in-progress recording."""
def setUp(self):
from django.contrib.auth import get_user_model
from rest_framework.test import APIRequestFactory, force_authenticate
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
User = get_user_model()
self.user = User.objects.create_user(
username="dvr_test_admin", password="pass",
user_level=User.UserLevel.ADMIN,
)
self.factory = APIRequestFactory()
self.force_authenticate = force_authenticate
def _delete_recording(self, rec):
from apps.channels.api_views import RecordingViewSet
request = self.factory.delete(f"/api/channels/recordings/{rec.id}/")
self.force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"delete": "destroy"})
return view(request, pk=rec.id)
@patch("apps.channels.api_views._stop_dvr_clients")
def test_destroy_completed_recording_does_not_stop_dvr_clients(self, mock_stop):
"""Deleting a completed recording must NOT call _stop_dvr_clients."""
rec = Recording.objects.create(
channel=self.channel,
start_time=timezone.now() - timedelta(hours=2),
end_time=timezone.now() - timedelta(hours=1),
custom_properties={"status": "completed", "file_path": "/data/recordings/test.mkv"},
)
self._delete_recording(rec)
mock_stop.assert_not_called()
@patch("apps.channels.api_views._stop_dvr_clients")
def test_destroy_upcoming_recording_does_not_stop_dvr_clients(self, mock_stop):
"""Deleting an upcoming (scheduled) recording must NOT call _stop_dvr_clients."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={},
)
self._delete_recording(rec)
mock_stop.assert_not_called()
@patch("apps.channels.api_views._stop_dvr_clients")
def test_destroy_active_recording_does_stop_dvr_clients(self, mock_stop):
"""Deleting an in-progress recording MUST call _stop_dvr_clients."""
mock_stop.return_value = 1
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=5),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
self._delete_recording(rec)
mock_stop.assert_called_once_with(str(self.channel.uuid), recording_id=rec.id)
class PeriodicTaskCleanupOnExecutionTests(TestCase):
"""Tests for PeriodicTask cleanup when run_recording starts."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.prefetch_recording_artwork")
@patch("apps.channels.tasks.get_channel_layer")
def test_periodic_task_cleaned_up_on_execution(self, mock_layer, mock_artwork):
"""When run_recording executes, it deletes its own PeriodicTask."""
mock_layer.return_value = MagicMock()
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={},
)
# post_save signal should have created the PeriodicTask
task_name = f"dvr-recording-{rec.id}"
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
pt = PeriodicTask.objects.get(name=task_name)
clocked_id = pt.clocked_id
from apps.channels.tasks import run_recording as run_rec_task
# This will proceed past guards, clean up the PeriodicTask, then
# eventually fail on the actual stream connection (expected)
try:
run_rec_task(rec.id, self.channel.id, str(future), str(future + timedelta(hours=1)))
except Exception:
pass
self.assertFalse(PeriodicTask.objects.filter(name=task_name).exists())
self.assertFalse(ClockedSchedule.objects.filter(id=clocked_id).exists())
@@ -0,0 +1,356 @@
"""Tests for the DVR Stop/Cancel feature set.
Covers:
- stop() endpoint
- destroy() was_in_progress field in recording_cancelled WebSocket event
- signals.py update_fields re-entrancy guard
- run_recording race guard before status write
- _stop_dvr_clients() DVR client isolation
"""
from datetime import timedelta
from unittest.mock import MagicMock, AsyncMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording
from apps.channels.api_views import RecordingViewSet, _stop_dvr_clients
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="stop_test_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
def _async_channel_layer_mock():
layer = MagicMock()
layer.group_send = AsyncMock()
return layer
class StopEndpointTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/stop/"""
def setUp(self):
self.channel = Channel.objects.create(channel_number=99, name="Stop Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _stop(self, rec):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/stop/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "stop"})
return view(request, pk=rec.id)
def _make_rec(self, status="recording"):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": status},
)
@patch("core.utils.send_websocket_update")
@patch("threading.Thread")
def test_stop_writes_status_to_db_before_returning(self, mock_thread, mock_ws):
"""DB write is synchronous — run_recording polls for this."""
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
response = self._stop(rec)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data.get("success"))
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("status"), "stopped")
@patch("core.utils.send_websocket_update")
@patch("threading.Thread")
def test_stop_writes_stopped_at_timestamp(self, mock_thread, mock_ws):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
self._stop(rec)
rec.refresh_from_db()
self.assertIn("stopped_at", rec.custom_properties)
def test_stop_calls_stop_dvr_clients_in_background(self):
"""stop() spawns a background thread whose target calls _stop_dvr_clients."""
rec = self._make_rec()
with patch("apps.channels.api_views._stop_dvr_clients", return_value=1) as mock_stop, \
patch("core.utils.send_websocket_update"), \
patch("threading.Thread") as mock_thread:
mock_thread.return_value.start = MagicMock()
self._stop(rec)
# Verify a daemon thread was spawned
mock_thread.assert_called_once()
thread_kwargs = mock_thread.call_args[1]
self.assertTrue(thread_kwargs.get("daemon"), "Thread must be daemon")
# Execute the captured target with DB connection close patched out
target = thread_kwargs["target"]
with patch("apps.channels.api_views._stop_dvr_clients", return_value=1) as mock_stop2, \
patch("apps.channels.signals.revoke_task", side_effect=Exception("skip")), \
patch("django.db.connection") as mock_conn:
target()
self.assertTrue(mock_stop2.called)
args, kwargs = mock_stop2.call_args
actual_rec_id = kwargs.get("recording_id") or (args[1] if len(args) > 1 else None)
self.assertEqual(actual_rec_id, rec.id)
def test_stop_returns_404_for_nonexistent(self):
request = self.factory.post("/api/channels/recordings/99999/stop/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "stop"})
self.assertEqual(view(request, pk=99999).status_code, 404)
@patch("core.utils.send_websocket_update")
@patch("threading.Thread")
def test_stop_idempotent_on_already_stopped(self, mock_thread, mock_ws):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec(status="stopped")
self.assertEqual(self._stop(rec).status_code, 200)
class CancelDestroyWasInProgressTests(TestCase):
"""was_in_progress field in the recording_cancelled WebSocket event."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=98, name="Cancel Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _delete(self, rec):
request = self.factory.delete(f"/api/channels/recordings/{rec.id}/")
force_authenticate(request, user=self.user)
return RecordingViewSet.as_view({"delete": "destroy"})(request, pk=rec.id)
@patch("apps.channels.api_views._stop_dvr_clients", return_value=1)
@patch("core.utils.send_websocket_update")
def test_in_progress_sends_was_in_progress_true(self, mock_ws, _):
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=10),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
self._delete(rec)
payload = mock_ws.call_args[0][2]
self.assertEqual(payload["type"], "recording_cancelled")
self.assertTrue(payload["was_in_progress"])
@patch("core.utils.send_websocket_update")
def test_completed_sends_was_in_progress_false(self, mock_ws):
rec = Recording.objects.create(
channel=self.channel,
start_time=timezone.now() - timedelta(hours=2),
end_time=timezone.now() - timedelta(hours=1),
custom_properties={"status": "completed"},
)
self._delete(rec)
self.assertFalse(mock_ws.call_args[0][2]["was_in_progress"])
class SignalUpdateFieldsReentrancyGuardTests(TestCase):
"""update_fields guard in schedule_task_on_save prevents redundant WS events."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=97, name="Signal Guard Channel")
def _create_upcoming(self):
future = timezone.now() + timedelta(hours=2)
return Recording.objects.create(
channel=self.channel, start_time=future,
end_time=future + timedelta(hours=1), custom_properties={},
)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_custom_properties_save_skips_artwork(self, mock_artwork):
rec = self._create_upcoming()
mock_artwork.reset_mock()
rec.custom_properties = {"poster_url": "https://example.com/p.jpg"}
rec.save(update_fields=["custom_properties"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_task_id_save_skips_artwork(self, mock_artwork):
rec = self._create_upcoming()
mock_artwork.reset_mock()
rec.task_id = "dvr-recording-999"
rec.save(update_fields=["task_id"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_combined_metadata_save_skips_artwork(self, mock_artwork):
rec = self._create_upcoming()
mock_artwork.reset_mock()
rec.task_id = "dvr-recording-1000"
rec.custom_properties = {"poster_url": "x"}
rec.save(update_fields=["custom_properties", "task_id"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_creation_dispatches_artwork(self, mock_artwork):
mock_artwork.apply_async.return_value = MagicMock()
self._create_upcoming()
self.assertTrue(mock_artwork.apply_async.called)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_scheduling_field_update_dispatches_artwork(self, mock_artwork):
"""save(update_fields=['start_time']) is not a metadata save — dispatch runs."""
mock_artwork.apply_async.return_value = MagicMock()
rec = self._create_upcoming()
mock_artwork.reset_mock()
future = timezone.now() + timedelta(hours=3)
rec.start_time = future
rec.end_time = future + timedelta(hours=1)
rec.save(update_fields=["start_time", "end_time"])
mock_artwork.apply_async.assert_called()
class RunRecordingRaceGuardTests(TestCase):
"""Race guard: stop() fires between idempotency check and status write."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=96, name="Race Guard Channel")
def test_race_guard_exits_when_stopped_at_db_read(self):
"""If Recording.objects.get() shows 'stopped', the task must exit
without writing 'recording' to the DB."""
from apps.channels.tasks import run_recording as run_rec
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=1),
end_time=now + timedelta(hours=1),
custom_properties={},
)
mock_layer = _async_channel_layer_mock()
original_get = Recording.objects.get
def patched_get(*args, **kwargs):
obj = original_get(*args, **kwargs)
if kwargs.get("id") == rec.id or (args and args[0] == rec.id):
obj.custom_properties = {"status": "stopped"}
return obj
with patch("apps.channels.tasks.get_channel_layer", return_value=mock_layer), \
patch("core.utils.log_system_event", side_effect=Exception("skip")), \
patch.object(Recording.objects, "get", side_effect=patched_get):
result = run_rec(
rec.id, self.channel.id, str(rec.start_time), str(rec.end_time),
)
self.assertIsNone(result)
rec.refresh_from_db()
self.assertNotEqual(
rec.custom_properties.get("status"), "recording",
"Race guard failed: task overwrote 'stopped' with 'recording'",
)
def test_idempotency_guard_catches_stopped_before_channel_layer(self):
"""When status='stopped' at the idempotency check, get_channel_layer is never called."""
from apps.channels.tasks import run_recording as run_rec
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=5),
end_time=now + timedelta(hours=1),
custom_properties={"status": "stopped"},
)
with patch("apps.channels.tasks.get_channel_layer") as mock_get_layer:
result = run_rec(
rec.id, self.channel.id, str(rec.start_time), str(rec.end_time),
)
self.assertIsNone(result)
mock_get_layer.assert_not_called()
class StopDvrClientsTests(TestCase):
"""_stop_dvr_clients() DVR client isolation."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=95, name="DVR Clients Channel")
self._redis = "core.utils.RedisClient"
self._sc = "apps.proxy.ts_proxy.services.channel_service.ChannelService.stop_client"
self._sch = "apps.proxy.ts_proxy.services.channel_service.ChannelService.stop_channel"
def _mock_redis(self, client_ids, ua_map):
r = MagicMock()
r.smembers.return_value = {c.encode() for c in client_ids}
def hget_side(key, field):
ks = key if isinstance(key, str) else key.decode("utf-8", errors="replace")
for cid, ua in ua_map.items():
if cid in ks:
return ua.encode() if isinstance(ua, str) else ua
return b""
r.hget.side_effect = hget_side
return r
def test_returns_zero_when_redis_none(self):
with patch(self._redis) as rc:
rc.get_client.return_value = None
self.assertEqual(_stop_dvr_clients(str(self.channel.uuid)), 0)
def test_stops_only_matching_client_when_recording_id_given(self):
r = self._mock_redis(
["client-a", "client-b"],
{"client-a": "Dispatcharr-DVR/recording-42",
"client-b": "Dispatcharr-DVR/recording-99"},
)
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
result = _stop_dvr_clients(str(self.channel.uuid), recording_id=42)
self.assertEqual(result, 1)
stopped = [c[0][1] for c in sc.call_args_list]
self.assertIn("client-a", stopped)
self.assertNotIn("client-b", stopped)
def test_stops_all_dvr_clients_without_recording_id(self):
r = self._mock_redis(
["client-a", "client-b"],
{"client-a": "Dispatcharr-DVR/recording-42",
"client-b": "Dispatcharr-DVR/recording-99"},
)
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
result = _stop_dvr_clients(str(self.channel.uuid))
self.assertEqual(result, 2)
def test_skips_non_dvr_clients(self):
r = self._mock_redis(
["viewer", "dvr-client"],
{"viewer": "Mozilla/5.0", "dvr-client": "Dispatcharr-DVR/recording-1"},
)
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
result = _stop_dvr_clients(str(self.channel.uuid))
self.assertEqual(result, 1)
stopped = [c[0][1] for c in sc.call_args_list]
self.assertNotIn("viewer", stopped)
def test_returns_zero_for_empty_channel(self):
r = MagicMock()
r.smembers.return_value = set()
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
self.assertEqual(_stop_dvr_clients(str(self.channel.uuid)), 0)
sc.assert_not_called()
def test_never_calls_stop_channel(self):
"""Must not stop the whole channel proxy — only individual clients."""
r = self._mock_redis(["dvr-1"], {"dvr-1": "Dispatcharr-DVR/recording-1"})
with patch(self._redis) as rc, patch(self._sc), patch(self._sch) as sch:
rc.get_client.return_value = r
_stop_dvr_clients(str(self.channel.uuid))
sch.assert_not_called()
@@ -0,0 +1,40 @@
from datetime import datetime, timedelta
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel, RecurringRecordingRule, Recording
from apps.channels.tasks import sync_recurring_rule_impl, purge_recurring_rule_impl
class RecurringRecordingRuleTasksTests(TestCase):
def test_sync_recurring_rule_creates_and_purges_recordings(self):
now = timezone.now()
channel = Channel.objects.create(channel_number=1, name='Test Channel')
start_time = (now + timedelta(minutes=15)).time().replace(second=0, microsecond=0)
end_time = (now + timedelta(minutes=75)).time().replace(second=0, microsecond=0)
rule = RecurringRecordingRule.objects.create(
channel=channel,
days_of_week=[now.weekday()],
start_time=start_time,
end_time=end_time,
)
created = sync_recurring_rule_impl(rule.id, drop_existing=True, horizon_days=1)
self.assertEqual(created, 1)
recording = Recording.objects.filter(custom_properties__rule__id=rule.id).first()
self.assertIsNotNone(recording)
self.assertEqual(recording.channel, channel)
self.assertEqual(recording.custom_properties.get('rule', {}).get('id'), rule.id)
expected_start = timezone.make_aware(
datetime.combine(recording.start_time.date(), start_time),
timezone.get_current_timezone(),
)
self.assertLess(abs((recording.start_time - expected_start).total_seconds()), 60)
removed = purge_recurring_rule_impl(rule.id)
self.assertEqual(removed, 1)
self.assertFalse(Recording.objects.filter(custom_properties__rule__id=rule.id).exists())
@@ -0,0 +1,718 @@
"""Tests for series rule evaluation deduplication.
Unit tests verify the dedup logic in evaluate_series_rules_impl.
Integration tests exercise the full path: EPG refresh → series rule
evaluation → Recording creation → post_save signal chain.
"""
from datetime import timedelta
from unittest.mock import patch, MagicMock
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel, Recording
from apps.epg.models import EPGSource, EPGData, ProgramData
from core.models import CoreSettings
def _set_series_rules(rules):
"""Helper to store series rules in CoreSettings."""
CoreSettings.set_dvr_series_rules(rules)
def _set_dvr_offsets(pre_min=0, post_min=0):
"""Helper to store DVR pre/post offsets."""
CoreSettings._update_group("dvr_settings", "DVR Settings", {
"pre_offset_minutes": pre_min,
"post_offset_minutes": post_min,
})
class SeriesRuleDedupBaseTestCase(TestCase):
"""Shared setup for series rule dedup tests."""
def setUp(self):
self.now = timezone.now()
self.epg_source = EPGSource.objects.create(
name="Test EPG", source_type="xmltv"
)
self.epg = EPGData.objects.create(
tvg_id="test.channel.1",
name="Test Channel EPG",
epg_source=self.epg_source,
)
self.channel = Channel.objects.create(
channel_number=1, name="Test Channel", epg_data=self.epg
)
_set_series_rules([{
"tvg_id": "test.channel.1",
"mode": "all",
"title": "Test Show",
}])
_set_dvr_offsets(pre_min=0, post_min=0)
def _create_program(self, hours_from_now=1, title="Test Show",
sub_title="Episode 1", tvg_id="test.channel.1"):
"""Create a ProgramData at the given offset."""
start = self.now + timedelta(hours=hours_from_now)
end = start + timedelta(hours=1)
return ProgramData.objects.create(
epg=self.epg,
tvg_id=tvg_id,
start_time=start,
end_time=end,
title=title,
sub_title=sub_title,
)
def _simulate_epg_refresh(self, programs_data):
"""Delete all ProgramData and recreate with new IDs (simulates EPG refresh)."""
ProgramData.objects.filter(epg=self.epg).delete()
new_programs = []
for data in programs_data:
prog = ProgramData.objects.create(epg=self.epg, **data)
new_programs.append(prog)
return new_programs
def _program_data_for_refresh(self, prog):
"""Build the dict needed by _simulate_epg_refresh from a ProgramData."""
return {
"tvg_id": prog.tvg_id,
"start_time": prog.start_time,
"end_time": prog.end_time,
"title": prog.title,
"sub_title": prog.sub_title,
}
# ---------------------------------------------------------------------------
# Unit tests: dedup logic in evaluate_series_rules_impl
# ---------------------------------------------------------------------------
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class ProgramIdStabilityTests(SeriesRuleDedupBaseTestCase):
"""Verify dedup works after EPG refresh changes ProgramData IDs."""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_no_duplicate_after_epg_refresh(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Same program should not be recorded twice after EPG refresh."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
old_id = prog.id
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
new_programs = self._simulate_epg_refresh(
[self._program_data_for_refresh(prog)]
)
self.assertNotEqual(old_id, new_programs[0].id)
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self.assertEqual(result2["scheduled"], 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_no_duplicate_with_offsets_after_refresh(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Dedup works when DVR offsets shift Recording times away from program times."""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=5, post_min=5)
prog = self._create_program(hours_from_now=2)
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
rec = Recording.objects.first()
self.assertEqual(rec.start_time, prog.start_time - timedelta(minutes=5))
self.assertEqual(rec.end_time, prog.end_time + timedelta(minutes=5))
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_different_episodes_still_recorded(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Different episodes on the same channel should each get a recording."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2, sub_title="Episode 1")
self._create_program(hours_from_now=4, sub_title="Episode 2")
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 2)
self.assertEqual(Recording.objects.count(), 2)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_new_episode_after_refresh_is_recorded(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""A genuinely new episode appearing after EPG refresh should be recorded."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2, sub_title="Episode 1")
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self._simulate_epg_refresh([
self._program_data_for_refresh(prog),
{
"tvg_id": "test.channel.1",
"start_time": prog.end_time,
"end_time": prog.end_time + timedelta(hours=1),
"title": "Test Show",
"sub_title": "Episode 2",
},
])
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2)
self.assertEqual(result2["scheduled"], 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_multiple_epg_refreshes_no_duplicates(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Multiple consecutive EPG refreshes should not accumulate duplicates."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
for _ in range(5):
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class ConcurrencyGuardTests(SeriesRuleDedupBaseTestCase):
"""Verify the task lock prevents concurrent evaluation."""
def test_lock_acquired_and_released(self, mock_schedule, mock_artwork):
"""evaluate_series_rules_impl acquires and releases the task lock."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock", return_value=True) as mock_lock, \
patch("apps.channels.tasks.release_task_lock") as mock_release:
evaluate_series_rules_impl()
mock_lock.assert_called_once_with('evaluate_series_rules', 'all')
mock_release.assert_called_once_with('evaluate_series_rules', 'all')
def test_skips_when_lock_held(self, mock_schedule, mock_artwork):
"""Returns early with skip reason when lock is already held."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock", return_value=False):
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 0)
self.assertTrue(
any(d.get("reason") == "concurrent evaluation in progress"
for d in result["details"]),
)
self.assertEqual(Recording.objects.count(), 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_lock_released_on_exception(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Lock is released even if the inner implementation raises."""
from apps.channels.tasks import evaluate_series_rules_impl
with patch("apps.channels.tasks._evaluate_series_rules_locked",
side_effect=RuntimeError("test error")):
with self.assertRaises(RuntimeError):
evaluate_series_rules_impl()
mock_release.assert_called_once_with('evaluate_series_rules', 'all')
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class SecondaryGuardTests(SeriesRuleDedupBaseTestCase):
"""Verify the secondary DB guard uses stable program attributes."""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_secondary_guard_catches_duplicate_with_offsets(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Secondary guard works with stale program IDs and DVR offsets."""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=10, post_min=10)
prog = self._create_program(hours_from_now=2)
# Pre-existing recording with a stale program ID (from previous EPG refresh)
Recording.objects.create(
channel=self.channel,
start_time=prog.start_time - timedelta(minutes=10),
end_time=prog.end_time + timedelta(minutes=10),
custom_properties={
"program": {
"id": 99999,
"tvg_id": prog.tvg_id,
"title": prog.title,
"start_time": prog.start_time.isoformat(),
"end_time": prog.end_time.isoformat(),
}
},
)
result = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self.assertEqual(result["scheduled"], 0)
# ---------------------------------------------------------------------------
# Integration tests: full path from EPG refresh through recording creation
# ---------------------------------------------------------------------------
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class IntegrationEPGRefreshTests(SeriesRuleDedupBaseTestCase):
"""End-to-end tests simulating the EPG refresh → evaluate → record flow.
These exercise the full signal chain: evaluate_series_rules_impl creates
a Recording, the post_save signal fires schedule_recording_task, and
subsequent evaluations (after EPG refresh) must not create duplicates.
"""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_single_episode_no_duplicates(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Simulate: create rule → evaluate → EPG refresh → re-evaluate.
The full recording lifecycle must result in exactly 1 recording.
"""
from apps.channels.tasks import evaluate_series_rules_impl
# Initial EPG data
prog = self._create_program(hours_from_now=2, sub_title="Pilot")
# First evaluation creates the recording
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
# Verify the recording was created with correct program metadata
rec = Recording.objects.first()
self.assertEqual(rec.custom_properties["program"]["tvg_id"], "test.channel.1")
self.assertEqual(rec.custom_properties["program"]["title"], "Test Show")
self.assertEqual(
rec.custom_properties["program"]["start_time"],
prog.start_time.isoformat()
)
# Verify the post_save signal scheduled a task
mock_schedule.assert_called()
initial_schedule_count = mock_schedule.call_count
# Simulate EPG refresh (programs get new DB IDs)
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
# Re-evaluate after refresh (this is what EPG refresh triggers)
result2 = evaluate_series_rules_impl()
self.assertEqual(result2["scheduled"], 0)
self.assertEqual(Recording.objects.count(), 1)
# No additional task scheduling should have occurred
self.assertEqual(mock_schedule.call_count, initial_schedule_count)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_with_offsets_no_duplicates(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Full flow with DVR offsets: recording times differ from program times."""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=5, post_min=10)
prog = self._create_program(hours_from_now=3, sub_title="Episode 1")
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
rec = Recording.objects.first()
# Verify offset-adjusted recording times
self.assertEqual(rec.start_time, prog.start_time - timedelta(minutes=5))
self.assertEqual(rec.end_time, prog.end_time + timedelta(minutes=10))
# Verify original (unadjusted) program times in custom_properties
self.assertEqual(
rec.custom_properties["program"]["start_time"],
prog.start_time.isoformat()
)
self.assertEqual(
rec.custom_properties["program"]["end_time"],
prog.end_time.isoformat()
)
# EPG refresh + re-evaluate
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self.assertEqual(result2["scheduled"], 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_multiple_episodes_across_refreshes(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""New episodes appear across multiple EPG refreshes; each recorded once."""
from apps.channels.tasks import evaluate_series_rules_impl
ep1 = self._create_program(hours_from_now=2, sub_title="Episode 1")
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
# EPG refresh adds episode 2 alongside episode 1
ep1_data = self._program_data_for_refresh(ep1)
ep2_start = ep1.end_time
ep2_data = {
"tvg_id": "test.channel.1",
"start_time": ep2_start,
"end_time": ep2_start + timedelta(hours=1),
"title": "Test Show",
"sub_title": "Episode 2",
}
self._simulate_epg_refresh([ep1_data, ep2_data])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2)
# Another EPG refresh adds episode 3
ep3_start = ep2_start + timedelta(hours=1)
ep3_data = {
"tvg_id": "test.channel.1",
"start_time": ep3_start,
"end_time": ep3_start + timedelta(hours=1),
"title": "Test Show",
"sub_title": "Episode 3",
}
self._simulate_epg_refresh([ep1_data, ep2_data, ep3_data])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 3)
# Final EPG refresh with no new episodes — count must stay at 3
self._simulate_epg_refresh([ep1_data, ep2_data, ep3_data])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 3)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_multiple_series_rules(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Multiple series rules on different channels, each evaluated correctly."""
from apps.channels.tasks import evaluate_series_rules_impl
# Second channel with its own EPG
epg2 = EPGData.objects.create(
tvg_id="test.channel.2",
name="Channel 2 EPG",
epg_source=self.epg_source,
)
channel2 = Channel.objects.create(
channel_number=2, name="Test Channel 2", epg_data=epg2
)
_set_series_rules([
{"tvg_id": "test.channel.1", "mode": "all", "title": "Show A"},
{"tvg_id": "test.channel.2", "mode": "all", "title": "Show B"},
])
# Programs on both channels
start1 = self.now + timedelta(hours=2)
prog1 = ProgramData.objects.create(
epg=self.epg, tvg_id="test.channel.1",
start_time=start1, end_time=start1 + timedelta(hours=1),
title="Show A", sub_title="Episode 1",
)
start2 = self.now + timedelta(hours=3)
prog2 = ProgramData.objects.create(
epg=epg2, tvg_id="test.channel.2",
start_time=start2, end_time=start2 + timedelta(hours=1),
title="Show B", sub_title="Episode 1",
)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2)
self.assertEqual(Recording.objects.filter(channel=self.channel).count(), 1)
self.assertEqual(Recording.objects.filter(channel=channel2).count(), 1)
# EPG refresh for both channels
ProgramData.objects.filter(epg=self.epg).delete()
ProgramData.objects.filter(epg=epg2).delete()
ProgramData.objects.create(
epg=self.epg, tvg_id="test.channel.1",
start_time=start1, end_time=start1 + timedelta(hours=1),
title="Show A", sub_title="Episode 1",
)
ProgramData.objects.create(
epg=epg2, tvg_id="test.channel.2",
start_time=start2, end_time=start2 + timedelta(hours=1),
title="Show B", sub_title="Episode 1",
)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2,
"No duplicates across multiple series rules after EPG refresh")
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_rapid_epg_refreshes_simulate_user_report(
self, mock_release, mock_lock, mock_schedule, mock_artwork
):
"""Reproduce the user-reported scenario: series rule + multiple EPG refreshes
causing count to balloon from 6 to 25 and 5 simultaneous recordings.
Simulates 6 episodes with 5 EPG refreshes (each assigning new ProgramData IDs).
"""
from apps.channels.tasks import evaluate_series_rules_impl
# Create 6 episodes (the user had "next of 6")
episodes = []
for i in range(6):
start = self.now + timedelta(hours=2 + i * 2)
episodes.append({
"tvg_id": "test.channel.1",
"start_time": start,
"end_time": start + timedelta(hours=1),
"title": "Test Show",
"sub_title": f"Episode {i + 1}",
})
# Create initial ProgramData
for ep in episodes:
ProgramData.objects.create(epg=self.epg, **ep)
# First evaluation: should create exactly 6 recordings
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 6)
# Simulate 5 EPG refreshes (the user saw count balloon to 25)
for refresh_num in range(5):
self._simulate_epg_refresh(episodes)
result = evaluate_series_rules_impl()
self.assertEqual(
Recording.objects.count(), 6,
f"After EPG refresh #{refresh_num + 1}, expected 6 recordings "
f"but got {Recording.objects.count()}"
)
self.assertEqual(result["scheduled"], 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_recording_survives_program_removal_and_readd(
self, mock_release, mock_lock, mock_schedule, mock_artwork
):
"""Program temporarily disappears from EPG then reappears — no duplicate."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2, sub_title="Episode 1")
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
# EPG refresh removes the program entirely
self._simulate_epg_refresh([])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"Existing recording preserved when program disappears from EPG")
# EPG refresh adds the program back (new ID)
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"No duplicate when program reappears with new ID")
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_celery_task_wrapper_calls_impl(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""The @shared_task evaluate_series_rules delegates to _impl correctly."""
from apps.channels.tasks import evaluate_series_rules
self._create_program(hours_from_now=2)
result = evaluate_series_rules()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
# Call again (simulating a second EPG refresh trigger)
result2 = evaluate_series_rules()
self.assertEqual(result2["scheduled"], 0)
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_tvg_id_scoped_evaluation(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Scoped evaluation (tvg_id parameter) still prevents duplicates."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
result1 = evaluate_series_rules_impl(tvg_id="test.channel.1")
self.assertEqual(result1["scheduled"], 1)
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result2 = evaluate_series_rules_impl(tvg_id="test.channel.1")
self.assertEqual(result2["scheduled"], 0)
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_offset_change_between_refreshes(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Changing DVR offsets between EPG refreshes doesn't create duplicates.
Even though Recording.start_time/end_time change when offsets change,
the dedup key uses the original program times from custom_properties.
"""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=5, post_min=5)
prog = self._create_program(hours_from_now=2)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
rec = Recording.objects.first()
original_start = rec.start_time
original_end = rec.end_time
# Change offsets
_set_dvr_offsets(pre_min=10, post_min=15)
# EPG refresh
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"Changing offsets between refreshes should not create duplicates")
self.assertEqual(result["scheduled"], 0)
# ---------------------------------------------------------------------------
# Edge case tests: Redis unavailability, non-series recordings, robustness
# ---------------------------------------------------------------------------
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class RedisUnavailabilityTests(SeriesRuleDedupBaseTestCase):
"""Verify evaluation works when Redis is unavailable (lock cannot be acquired)."""
def test_proceeds_when_redis_down(self, mock_schedule, mock_artwork):
"""Evaluation succeeds (with dedup guards) when Redis raises on lock acquire."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")):
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
def test_dedup_still_works_without_lock(self, mock_schedule, mock_artwork):
"""Dedup guards prevent duplicates even when the lock is unavailable."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
# First call: Redis down, proceeds without lock
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")):
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
# EPG refresh
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
# Second call: Redis still down
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")):
result = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"Dedup guards prevent duplicates even without lock")
self.assertEqual(result["scheduled"], 0)
def test_lock_not_released_when_not_acquired(self, mock_schedule, mock_artwork):
"""release_task_lock is not called if acquire raised an exception."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")), \
patch("apps.channels.tasks.release_task_lock") as mock_release:
evaluate_series_rules_impl()
mock_release.assert_not_called()
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class NonSeriesRecordingTests(SeriesRuleDedupBaseTestCase):
"""Verify non-series recordings don't interfere with series rule dedup."""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_manual_recording_without_program_data_ignored(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Recordings without custom_properties.program are skipped by dedup key builder."""
from apps.channels.tasks import evaluate_series_rules_impl
# Manual recording with no program metadata
Recording.objects.create(
channel=self.channel,
start_time=self.now + timedelta(hours=2),
end_time=self.now + timedelta(hours=3),
custom_properties={},
)
prog = self._create_program(hours_from_now=2)
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 2)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_recurring_rule_recording_does_not_interfere(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Recordings from recurring rules (custom_properties.rule) don't block series rules."""
from apps.channels.tasks import evaluate_series_rules_impl
Recording.objects.create(
channel=self.channel,
start_time=self.now + timedelta(hours=2),
end_time=self.now + timedelta(hours=3),
custom_properties={"rule": {"id": 1, "name": "Daily News"}},
)
prog = self._create_program(hours_from_now=2)
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 2)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_recording_with_null_custom_properties_ignored(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Recordings with None custom_properties don't crash the dedup key builder."""
from apps.channels.tasks import evaluate_series_rules_impl
Recording.objects.create(
channel=self.channel,
start_time=self.now + timedelta(hours=2),
end_time=self.now + timedelta(hours=3),
custom_properties=None,
)
prog = self._create_program(hours_from_now=2)
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
@@ -0,0 +1,385 @@
"""Tests for ghost client detection and cleanup.
Covers:
- ClientManager.remove_ghost_clients() pipelined EXISTS logic
- channel_status detailed stats path removes ghost clients from Redis SET
- channel_status basic stats path removes ghost clients and corrects count
- _check_orphaned_metadata() validates client SET entries and cleans up
channels where all clients are ghosts
"""
from unittest.mock import MagicMock, patch, PropertyMock
from django.test import TestCase
from apps.proxy.ts_proxy.client_manager import ClientManager
from apps.proxy.ts_proxy.constants import ChannelMetadataField, ChannelState
from apps.proxy.ts_proxy.redis_keys import RedisKeys
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
CHANNEL_ID = "00000000-0000-0000-0000-000000000001"
def _make_proxy_server(redis_client=None):
"""Create a minimal mock ProxyServer with a redis_client."""
server = MagicMock()
server.redis_client = redis_client or MagicMock()
server.stream_managers = {}
server.client_managers = {}
server.worker_id = "test-worker-1"
return server
def _metadata_for_channel(state="active"):
"""Return a plausible channel metadata dict (bytes keys/values)."""
return {
ChannelMetadataField.STATE.encode(): state.encode(),
ChannelMetadataField.URL.encode(): b"http://example.com/stream",
ChannelMetadataField.STREAM_PROFILE.encode(): b"default",
ChannelMetadataField.OWNER.encode(): b"test-worker-1",
ChannelMetadataField.INIT_TIME.encode(): b"1773500000.0",
}
# ---------------------------------------------------------------------------
# Unit tests for ClientManager.remove_ghost_clients()
# ---------------------------------------------------------------------------
class RemoveGhostClientsTests(TestCase):
"""Directly exercises the static method that all callers rely on."""
def test_ghost_removed_and_returned(self):
"""Client ID in SET with no metadata hash should be SREM'd."""
redis = MagicMock()
redis.smembers.return_value = {b"ghost_001"}
pipe = MagicMock()
redis.pipeline.return_value = pipe
pipe.execute.return_value = [False] # EXISTS → False
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(result, [b"ghost_001"])
redis.srem.assert_called_once()
def test_live_client_preserved(self):
"""Client with valid metadata hash should NOT be removed."""
redis = MagicMock()
redis.smembers.return_value = {b"live_001"}
pipe = MagicMock()
redis.pipeline.return_value = pipe
pipe.execute.return_value = [True] # EXISTS → True
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(result, [])
redis.srem.assert_not_called()
def test_mixed_ghost_and_live(self):
"""Only ghost clients should be removed; live ones preserved."""
redis = MagicMock()
redis.smembers.return_value = {b"ghost_001", b"live_001"}
pipe = MagicMock()
redis.pipeline.return_value = pipe
# Order matches list(smembers), which is non-deterministic —
# map both IDs so the test is stable regardless of iteration order.
client_id_list = list(redis.smembers.return_value)
def exists_results():
return [
b"ghost_001" not in cid.decode() == False
for cid in client_id_list
]
# Simpler: mock based on key content
def pipe_exists(key):
pass # just enqueued; results come from execute()
pipe.exists.side_effect = pipe_exists
pipe.execute.return_value = [
"live" in cid.decode() for cid in client_id_list
]
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(len(result), 1)
self.assertTrue(any(b"ghost" in cid for cid in result))
redis.srem.assert_called_once()
def test_empty_set_returns_empty(self):
"""No clients means nothing to clean."""
redis = MagicMock()
redis.smembers.return_value = set()
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(result, [])
redis.pipeline.assert_not_called()
def test_pre_fetched_client_ids_skips_smembers(self):
"""When client_ids is passed, SMEMBERS should not be called."""
redis = MagicMock()
pipe = MagicMock()
redis.pipeline.return_value = pipe
pipe.execute.return_value = [False]
pre_fetched = {b"ghost_001"}
result = ClientManager.remove_ghost_clients(
redis, CHANNEL_ID, client_ids=pre_fetched
)
redis.smembers.assert_not_called()
self.assertEqual(len(result), 1)
# ---------------------------------------------------------------------------
# Detailed stats path: exercises get_detailed_channel_info()
# ---------------------------------------------------------------------------
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
class DetailedStatsGhostClientTests(TestCase):
"""get_detailed_channel_info() should remove ghost clients whose metadata
hash has expired from the Redis client SET."""
def _setup_redis(self, mock_proxy_cls, client_ids, hgetall_side_effect):
"""Wire up a mock ProxyServer with controlled Redis responses."""
redis = MagicMock()
server = _make_proxy_server(redis)
mock_proxy_cls.get_instance.return_value = server
redis.hgetall.side_effect = hgetall_side_effect
redis.smembers.return_value = client_ids
# buffer_index, ttl, exists all need safe defaults
redis.get.return_value = b"10"
redis.ttl.return_value = 300
redis.exists.return_value = True
return redis
def test_ghost_client_removed_from_set(self, mock_proxy_cls):
"""Ghost client should be SREM'd and excluded from result."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
def hgetall_side_effect(key):
if "clients:" in key:
return {} # ghost — metadata expired
return _metadata_for_channel()
redis = self._setup_redis(
mock_proxy_cls, {b"ghost_001"}, hgetall_side_effect
)
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
self.assertEqual(result['client_count'], 0)
self.assertEqual(len(result['clients']), 0)
redis.srem.assert_called_once()
def test_live_client_preserved(self, mock_proxy_cls):
"""Client with valid metadata should appear in results."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
def hgetall_side_effect(key):
if "clients:" in key:
return {
b'user_agent': b'VLC/3.0',
b'worker_id': b'test-worker-1',
b'connected_at': b'1773500000.0',
}
return _metadata_for_channel()
redis = self._setup_redis(
mock_proxy_cls, {b"live_001"}, hgetall_side_effect
)
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
self.assertEqual(result['client_count'], 1)
self.assertEqual(len(result['clients']), 1)
redis.srem.assert_not_called()
def test_mixed_ghost_and_live(self, mock_proxy_cls):
"""Only ghost clients should be removed; live ones preserved."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
def hgetall_side_effect(key):
if "clients:" in key:
if "ghost" in key:
return {}
return {
b'user_agent': b'VLC/3.0',
b'worker_id': b'test-worker-1',
}
return _metadata_for_channel()
redis = self._setup_redis(
mock_proxy_cls, {b"ghost_001", b"live_001"}, hgetall_side_effect
)
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
self.assertEqual(result['client_count'], 1)
self.assertEqual(len(result['clients']), 1)
redis.srem.assert_called_once()
# ---------------------------------------------------------------------------
# Basic stats path: exercises get_basic_channel_info()
# ---------------------------------------------------------------------------
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
class BasicStatsGhostClientTests(TestCase):
"""get_basic_channel_info() should call remove_ghost_clients(), skip
ghosts from display, and correct client_count."""
def _setup_redis(self, mock_proxy_cls, client_ids, ghost_ids):
"""Wire up mock ProxyServer. ghost_ids controls which EXISTS return False."""
redis = MagicMock()
server = _make_proxy_server(redis)
mock_proxy_cls.get_instance.return_value = server
redis.hgetall.return_value = _metadata_for_channel()
redis.get.return_value = b"10" # buffer_index
redis.scard.return_value = len(client_ids)
redis.smembers.return_value = client_ids
redis.hget.return_value = None # individual field lookups
# Pipeline for remove_ghost_clients
pipe = MagicMock()
redis.pipeline.return_value = pipe
client_id_list = list(client_ids)
pipe.execute.return_value = [
cid not in ghost_ids for cid in client_id_list
]
return redis
def test_ghost_removed_and_count_corrected(self, mock_proxy_cls):
"""Ghost client should be cleaned and client_count decremented."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
redis = self._setup_redis(
mock_proxy_cls,
client_ids={b"ghost_001"},
ghost_ids={b"ghost_001"},
)
result = ChannelStatus.get_basic_channel_info(CHANNEL_ID)
self.assertIsNotNone(result)
self.assertEqual(result['client_count'], 0)
redis.srem.assert_called_once()
def test_live_client_count_preserved(self, mock_proxy_cls):
"""Live clients should be counted correctly."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
redis = self._setup_redis(
mock_proxy_cls,
client_ids={b"live_001"},
ghost_ids=set(),
)
result = ChannelStatus.get_basic_channel_info(CHANNEL_ID)
self.assertIsNotNone(result)
self.assertEqual(result['client_count'], 1)
redis.srem.assert_not_called()
# ---------------------------------------------------------------------------
# Orphaned channel cleanup: exercises _check_orphaned_metadata()
# ---------------------------------------------------------------------------
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
class OrphanedChannelGhostValidationTests(TestCase):
"""_check_orphaned_metadata() should validate client SET entries when
owner is dead and client_count > 0. If all clients are ghosts, it
should clean up the channel."""
def _make_server_for_orphan_check(self, mock_proxy_cls, channel_id,
client_ids, ghost_ids, owner="dead-worker"):
"""Build a mock ProxyServer whose Redis state simulates an orphaned channel."""
redis = MagicMock()
server = _make_proxy_server(redis)
mock_proxy_cls.get_instance.return_value = server
metadata_key = RedisKeys.channel_metadata(channel_id)
metadata = _metadata_for_channel()
metadata[ChannelMetadataField.OWNER.encode()] = owner.encode()
# scan returns the one channel metadata key
redis.scan.return_value = (0, [metadata_key.encode()])
redis.hgetall.return_value = metadata
redis.scard.return_value = len(client_ids)
redis.smembers.return_value = client_ids
# Owner heartbeat is dead
redis.exists.side_effect = lambda key: (
False if "heartbeat" in key else True
)
# Pipeline for remove_ghost_clients
pipe = MagicMock()
redis.pipeline.return_value = pipe
client_id_list = list(client_ids)
pipe.execute.return_value = [
cid not in ghost_ids for cid in client_id_list
]
return server, redis
def test_all_ghosts_triggers_cleanup(self, mock_proxy_cls):
"""When all clients are ghosts, channel should be cleaned up."""
from apps.proxy.ts_proxy.server import ProxyServer
channel_id = "00000000-0000-0000-0000-000000000005"
server, redis = self._make_server_for_orphan_check(
mock_proxy_cls, channel_id,
client_ids={b"ghost_001", b"ghost_002"},
ghost_ids={b"ghost_001", b"ghost_002"},
)
# Call the real method on a real-ish ProxyServer
# The method lives on the server instance, so invoke it directly.
# We need to call _check_orphaned_metadata on the actual server mock,
# but it's a MagicMock. Instead, test via remove_ghost_clients directly
# and verify the cleanup decision logic.
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
real_count = max(0, len({b"ghost_001", b"ghost_002"}) - len(stale_ids))
self.assertEqual(len(stale_ids), 2)
self.assertEqual(real_count, 0)
redis.srem.assert_called_once()
def test_mixed_preserves_live_clients(self, mock_proxy_cls):
"""When some clients are live, real_count should be > 0."""
channel_id = "00000000-0000-0000-0000-000000000006"
server, redis = self._make_server_for_orphan_check(
mock_proxy_cls, channel_id,
client_ids={b"ghost_001", b"live_001"},
ghost_ids={b"ghost_001"},
)
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
real_count = max(0, 2 - len(stale_ids))
self.assertEqual(len(stale_ids), 1)
self.assertEqual(real_count, 1)
def test_no_ghosts_no_cleanup(self, mock_proxy_cls):
"""When all clients are live, no SREM should be called."""
channel_id = "00000000-0000-0000-0000-000000000007"
server, redis = self._make_server_for_orphan_check(
mock_proxy_cls, channel_id,
client_ids={b"live_001"},
ghost_ids=set(),
)
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
self.assertEqual(len(stale_ids), 0)
redis.srem.assert_not_called()
@@ -0,0 +1,231 @@
"""Tests for stuck INITIALIZING state fix.
Covers:
- stream_manager.run() finally block: ownership check + state guard fallback
- ChannelState.PRE_ACTIVE contains the correct states
- INITIALIZING is included in the cleanup task grace period check
"""
import time
import threading
from unittest.mock import MagicMock, patch
from django.test import TestCase
from apps.proxy.ts_proxy.constants import ChannelMetadataField, ChannelState
from apps.proxy.ts_proxy.redis_keys import RedisKeys
from apps.proxy.ts_proxy.stream_manager import StreamManager
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
CHANNEL_ID = "00000000-0000-0000-0000-000000000001"
def _make_stream_manager(tried_stream_ids=None, max_retries=3):
"""Build a StreamManager via __new__ (bypasses __init__) with the
minimum attributes required by the run() finally block."""
sm = StreamManager.__new__(StreamManager)
sm.channel_id = CHANNEL_ID
sm.worker_id = "worker-1"
sm.max_retries = max_retries
sm.tried_stream_ids = tried_stream_ids if tried_stream_ids is not None else set()
sm.running = False # while-loop exits immediately
sm.connected = False
sm.transcode_process_active = False
sm._buffer_check_timers = []
sm.url = "http://example.com/stream"
sm.url_switching = False
sm.url_switch_start_time = 0
sm.url_switch_timeout = 30
sm.stop_requested = False
sm.stopping = False
sm.socket = None
sm.transcode_process = None
sm.current_response = None
sm.current_session = None
sm.current_stream_id = None
buffer = MagicMock()
buffer.redis_client = MagicMock()
buffer.channel_id = CHANNEL_ID
sm.buffer = buffer
return sm
def _run_finally_block(sm, owner_value, current_state):
"""Invoke StreamManager.run() so its finally block executes against real code.
Patches threading.Thread and ConfigHelper so the try-block is inert
(self.running=False makes the while-loop exit immediately).
Returns True if the finally block wrote ERROR to Redis.
"""
redis = sm.buffer.redis_client
# Mock the owner key GET — the finally block calls redis.get(owner_key)
def get_side_effect(key):
if "owner" in key:
return owner_value
return None
redis.get.side_effect = get_side_effect
# Mock hget for state field lookup in the PRE_ACTIVE guard
if current_state is not None:
redis.hget.return_value = current_state.encode('utf-8')
else:
redis.hget.return_value = None
# Reset hset so we can detect whether ERROR was written
redis.hset.reset_mock()
redis.setex.reset_mock()
with patch.object(threading, 'Thread', return_value=MagicMock()):
with patch('apps.proxy.ts_proxy.stream_manager.ConfigHelper') as mock_cfg:
mock_cfg.max_stream_switches.return_value = 0
mock_cfg.max_retries.return_value = sm.max_retries
sm.run()
# Check if hset was called with ERROR state
if redis.hset.called:
mapping = redis.hset.call_args[1].get('mapping', {})
return mapping.get(ChannelMetadataField.STATE) == ChannelState.ERROR
return False
# ---------------------------------------------------------------------------
# stream_manager.run() finally block: ownership + state guard behavior
# ---------------------------------------------------------------------------
class StreamManagerFinallyBlockTests(TestCase):
"""The run() finally block writes ERROR if the worker is still the owner
(normal case) OR if ownership expired and the channel is still in a
pre-active state (no new owner has taken over)."""
# --- Owner still valid: always write ERROR ---
def test_owner_writes_error_regardless_of_state(self):
"""When we're still the owner, always write ERROR."""
sm = _make_stream_manager()
owner = sm.worker_id.encode('utf-8')
self.assertTrue(_run_finally_block(sm, owner, ChannelState.ACTIVE))
def test_owner_writes_error_on_initializing(self):
"""Owner + INITIALIZING = write ERROR."""
sm = _make_stream_manager()
owner = sm.worker_id.encode('utf-8')
self.assertTrue(_run_finally_block(sm, owner, ChannelState.INITIALIZING))
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
self.assertEqual(mapping[ChannelMetadataField.STATE], ChannelState.ERROR)
# --- Ownership expired, no new owner: use state guard ---
def test_no_owner_initializing_writes_error(self):
"""Ownership expired + INITIALIZING = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.INITIALIZING))
def test_no_owner_connecting_writes_error(self):
"""Ownership expired + CONNECTING = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.CONNECTING))
def test_no_owner_buffering_writes_error(self):
"""Ownership expired + BUFFERING = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.BUFFERING))
def test_no_owner_waiting_for_clients_writes_error(self):
"""Ownership expired + WAITING_FOR_CLIENTS = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.WAITING_FOR_CLIENTS))
def test_no_owner_active_does_not_write(self):
"""Ownership expired + ACTIVE = do NOT write ERROR."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, None, ChannelState.ACTIVE))
def test_no_owner_error_does_not_write(self):
"""Ownership expired + already ERROR = do NOT write again."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, None, ChannelState.ERROR))
def test_no_owner_no_state_does_not_write(self):
"""Ownership expired + no state metadata = do NOT write."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, None, None))
# --- New owner took over: never clobber ---
def test_new_owner_initializing_does_not_write(self):
"""Another worker owns the channel — do NOT clobber."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, b"other-worker", ChannelState.INITIALIZING))
def test_new_owner_active_does_not_write(self):
"""Another worker owns the channel and is ACTIVE — do NOT write."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, b"other-worker", ChannelState.ACTIVE))
# --- Stopping key and error messages ---
def test_stopping_key_set_on_error_update(self):
"""When ERROR is written, stopping key must also be set."""
sm = _make_stream_manager()
_run_finally_block(sm, None, ChannelState.INITIALIZING)
sm.buffer.redis_client.setex.assert_called_once()
args = sm.buffer.redis_client.setex.call_args[0]
self.assertIn("stopping", args[0])
self.assertEqual(args[1], 60)
def test_error_message_includes_stream_count(self):
"""When multiple streams were tried, error message reflects that."""
sm = _make_stream_manager(tried_stream_ids={1, 2, 3})
_run_finally_block(sm, None, ChannelState.INITIALIZING)
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
error_msg = mapping[ChannelMetadataField.ERROR_MESSAGE]
self.assertIn("3 stream options failed", error_msg)
def test_error_message_with_no_streams_tried(self):
"""When no alternate streams were tried, shows retry count."""
sm = _make_stream_manager(tried_stream_ids=set(), max_retries=5)
_run_finally_block(sm, None, ChannelState.INITIALIZING)
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
error_msg = mapping[ChannelMetadataField.ERROR_MESSAGE]
self.assertIn("5", error_msg)
# ---------------------------------------------------------------------------
# ChannelState.PRE_ACTIVE: verify contents and immutability
# ---------------------------------------------------------------------------
class PreActiveStateTests(TestCase):
"""Verify PRE_ACTIVE contains the correct states and is immutable."""
def test_initializing_in_pre_active(self):
self.assertIn(ChannelState.INITIALIZING, ChannelState.PRE_ACTIVE)
def test_connecting_in_pre_active(self):
self.assertIn(ChannelState.CONNECTING, ChannelState.PRE_ACTIVE)
def test_buffering_in_pre_active(self):
self.assertIn(ChannelState.BUFFERING, ChannelState.PRE_ACTIVE)
def test_waiting_for_clients_in_pre_active(self):
self.assertIn(ChannelState.WAITING_FOR_CLIENTS, ChannelState.PRE_ACTIVE)
def test_active_not_in_pre_active(self):
self.assertNotIn(ChannelState.ACTIVE, ChannelState.PRE_ACTIVE)
def test_error_not_in_pre_active(self):
self.assertNotIn(ChannelState.ERROR, ChannelState.PRE_ACTIVE)
def test_pre_active_is_frozenset(self):
self.assertIsInstance(ChannelState.PRE_ACTIVE, frozenset)
@@ -0,0 +1,331 @@
"""Tests for ts_proxy keepalive and stats-update behavior.
Covers:
- stream_generator._should_send_keepalive() owner vs non-owner worker paths
- stream_generator._should_send_keepalive() Redis last_data health check
- client_manager._do_stats_update() error handling and WebSocket dispatch
- client_manager.remove_client() non-blocking stats update
- Keepalive/DVR-timeout timing invariants
"""
import threading
import time
from unittest.mock import MagicMock, patch
from django.test import TestCase
# ---------------------------------------------------------------------------
# _should_send_keepalive: owner worker path
# ---------------------------------------------------------------------------
class OwnerWorkerKeepaliveTests(TestCase):
"""Owner worker has a stream_manager; keepalive logic uses it directly."""
def _make_generator(self, healthy, at_buffer_head, consecutive_empty):
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
gen = StreamGenerator.__new__(StreamGenerator)
gen.channel_id = "00000000-0000-0000-0000-000000000001"
gen.client_id = "test-client"
buffer = MagicMock()
buffer.index = 10 if at_buffer_head else 100
gen.local_index = 10
gen.buffer = buffer
stream_manager = MagicMock()
stream_manager.healthy = healthy
gen.stream_manager = stream_manager
gen.consecutive_empty = consecutive_empty
return gen
def test_owner_healthy_returns_false(self):
"""Owner worker, healthy stream -> no keepalive."""
gen = self._make_generator(healthy=True, at_buffer_head=True, consecutive_empty=10)
self.assertFalse(gen._should_send_keepalive(gen.local_index))
def test_owner_unhealthy_at_head_returns_true(self):
"""Owner worker, unhealthy stream, at buffer head -> send keepalive."""
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=10)
self.assertTrue(gen._should_send_keepalive(gen.local_index))
def test_owner_unhealthy_not_at_head_returns_false(self):
"""Owner worker, unhealthy stream, but NOT at buffer head -> no keepalive."""
gen = self._make_generator(healthy=False, at_buffer_head=False, consecutive_empty=10)
self.assertFalse(gen._should_send_keepalive(gen.local_index))
def test_owner_insufficient_consecutive_empty_returns_false(self):
"""Owner worker, unhealthy, at head but consecutive_empty < 5 -> no keepalive."""
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=3)
self.assertFalse(gen._should_send_keepalive(gen.local_index))
def test_owner_exactly_5_consecutive_empty_returns_true(self):
"""consecutive_empty == 5 is the minimum threshold."""
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=5)
self.assertTrue(gen._should_send_keepalive(gen.local_index))
# ---------------------------------------------------------------------------
# _should_send_keepalive: non-owner worker path
# ---------------------------------------------------------------------------
class NonOwnerWorkerKeepaliveTests(TestCase):
"""Non-owner worker has stream_manager=None; health determined from Redis."""
def _make_generator(self, consecutive_empty=10):
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
gen = StreamGenerator.__new__(StreamGenerator)
gen.channel_id = "00000000-0000-0000-0000-000000000002"
gen.client_id = "test-client-nonowner"
buffer = MagicMock()
buffer.index = 10
gen.local_index = 10
gen.buffer = buffer
gen.stream_manager = None # non-owner worker
gen.consecutive_empty = consecutive_empty
# Attributes added by health-check throttling (set in __init__)
gen._last_health_check_time = 0.0
gen._last_health_check_result = False
gen._health_check_interval = 2.0
gen.proxy_server = None
return gen
def _mock_proxy_server(self, last_data_value):
"""Return a mock ProxyServer with a redis_client pre-configured."""
server = MagicMock()
redis_client = MagicMock()
server.redis_client = redis_client
redis_client.get.return_value = last_data_value
return server
def test_non_owner_fresh_data_returns_false(self):
"""Non-owner, last_data < 10s ago -> stream healthy -> no keepalive."""
gen = self._make_generator()
fresh_ts = str(time.time() - 2.0).encode()
server = self._mock_proxy_server(fresh_ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result, "Fresh data should NOT trigger keepalive")
def test_non_owner_stale_data_returns_true(self):
"""Non-owner, last_data >= 10s ago -> stream unhealthy -> send keepalive."""
gen = self._make_generator()
stale_ts = str(time.time() - 12.0).encode()
server = self._mock_proxy_server(stale_ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertTrue(result, "Stale data (12s) should trigger keepalive")
def test_non_owner_exactly_at_timeout_returns_true(self):
"""Data age exactly equal to CONNECTION_TIMEOUT (10s) -> send keepalive."""
gen = self._make_generator()
ts = str(time.time() - 10.0).encode()
server = self._mock_proxy_server(ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertTrue(result, "Data at exactly timeout threshold should trigger keepalive")
def test_non_owner_no_redis_key_returns_true(self):
"""Non-owner, last_data key missing from Redis -> assume unhealthy."""
gen = self._make_generator()
server = self._mock_proxy_server(None)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertTrue(result, "Missing last_data key should trigger keepalive")
def test_non_owner_redis_client_none_returns_false(self):
"""Non-owner, redis_client is None (disconnected) -> conservative, no keepalive."""
gen = self._make_generator()
server = MagicMock()
server.redis_client = None
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result, "No redis_client -> conservative, no keepalive")
def test_non_owner_redis_exception_returns_false(self):
"""Non-owner, Redis raises an exception -> conservative, no keepalive."""
gen = self._make_generator()
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.side_effect = Exception("Redis error")
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result, "Redis error -> conservative, no keepalive")
def test_non_owner_not_at_buffer_head_returns_false(self):
"""Non-owner, NOT at buffer head -> no keepalive regardless of Redis."""
gen = self._make_generator()
gen.buffer.index = 100 # far ahead of local_index=10
server = self._mock_proxy_server(None)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result)
def test_non_owner_insufficient_consecutive_empty_returns_false(self):
"""Non-owner, at head, but consecutive_empty < 5 -> no keepalive."""
gen = self._make_generator(consecutive_empty=2)
stale_ts = str(time.time() - 30.0).encode()
server = self._mock_proxy_server(stale_ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result)
# ---------------------------------------------------------------------------
# _do_stats_update: error handling and WebSocket dispatch
# ---------------------------------------------------------------------------
class DoStatsUpdateTests(TestCase):
"""_do_stats_update runs the actual Redis scan + WebSocket call."""
def _make_client_manager(self):
from apps.proxy.ts_proxy.client_manager import ClientManager
cm = ClientManager.__new__(ClientManager)
cm.channel_id = "00000000-0000-0000-0000-000000000004"
cm._heartbeat_running = False
return cm
def test_do_stats_update_calls_send_websocket_update(self):
"""_do_stats_update must call send_websocket_update with channel_stats."""
cm = self._make_client_manager()
mock_redis = MagicMock()
mock_redis.scan.return_value = (0, [])
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update") as mock_ws, \
patch("redis.Redis.from_url", return_value=mock_redis):
cm._do_stats_update()
mock_ws.assert_called_once()
event_type = mock_ws.call_args[0][1]
self.assertEqual(event_type, "update")
payload = mock_ws.call_args[0][2]
self.assertEqual(payload["type"], "channel_stats")
def test_do_stats_update_does_not_raise_on_redis_error(self):
"""Redis failure must be swallowed (logged), not propagated."""
cm = self._make_client_manager()
with patch("redis.Redis.from_url", side_effect=Exception("Redis down")):
try:
cm._do_stats_update()
except Exception as e:
self.fail(f"_do_stats_update raised an exception: {e}")
def test_do_stats_update_scans_channel_client_keys(self):
"""Must scan for ts_proxy:channel:*:clients pattern."""
cm = self._make_client_manager()
mock_redis = MagicMock()
mock_redis.scan.return_value = (0, [])
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update"), \
patch("redis.Redis.from_url", return_value=mock_redis):
cm._do_stats_update()
scan_call = mock_redis.scan.call_args
self.assertIn("ts_proxy:channel:*:clients", str(scan_call))
# ---------------------------------------------------------------------------
# Integration: remove_client must not block on WebSocket
# ---------------------------------------------------------------------------
class ClientRemoveIntegrationTests(TestCase):
"""When remove_client() fires, _trigger_stats_update must not block."""
def test_remove_client_does_not_block_on_websocket(self):
"""remove_client() must return quickly even if WebSocket is slow."""
from apps.proxy.ts_proxy.client_manager import ClientManager
cm = ClientManager.__new__(ClientManager)
cm.channel_id = "00000000-0000-0000-0000-000000000005"
cm._heartbeat_running = False
cm.clients = {"test-client-1"}
cm.last_heartbeat_time = {"test-client-1": time.time()}
cm.last_active_time = time.time()
cm.client_set_key = f"ts_proxy:channel:{cm.channel_id}:clients"
cm.client_ttl = 60
cm.worker_id = "worker-1"
cm.proxy_server = MagicMock()
cm.proxy_server.am_i_owner.return_value = False
cm.lock = threading.Lock()
mock_redis = MagicMock()
mock_redis.hgetall.return_value = {b"ip_address": b"127.0.0.1"}
mock_redis.scard.return_value = 1
cm.redis_client = mock_redis
slow_ws_called = threading.Event()
def slow_websocket(*args, **kwargs):
time.sleep(2.0)
slow_ws_called.set()
start = time.time()
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update", side_effect=slow_websocket):
cm.remove_client("test-client-1")
elapsed = time.time() - start
self.assertLess(elapsed, 1.0,
f"remove_client() blocked for {elapsed:.2f}s waiting for WebSocket "
f"(should dispatch to background thread and return immediately)")
# ---------------------------------------------------------------------------
# DVR timeout threshold vs keepalive timing
# ---------------------------------------------------------------------------
class KeepaliveTimingTests(TestCase):
"""Verify that keepalive threshold gives sufficient margin before DVR timeout."""
def test_keepalive_threshold_less_than_dvr_timeout(self):
"""CONNECTION_TIMEOUT (keepalive trigger) must be < DVR read timeout (15s)."""
from apps.proxy.config import TSConfig as Config
connection_timeout = getattr(Config, "CONNECTION_TIMEOUT", 10)
dvr_read_timeout = 15 # hard-coded in run_recording: timeout=(10, 15)
self.assertLess(
connection_timeout,
dvr_read_timeout,
f"CONNECTION_TIMEOUT ({connection_timeout}s) must be < DVR timeout ({dvr_read_timeout}s) "
f"so keepalives fire before DVR times out",
)
def test_keepalive_interval_is_short(self):
"""KEEPALIVE_INTERVAL must be short enough to send multiple keepalives in the gap."""
from apps.proxy.config import TSConfig as Config
interval = getattr(Config, "KEEPALIVE_INTERVAL", 0.5)
connection_timeout = getattr(Config, "CONNECTION_TIMEOUT", 10)
remaining_window = 15 - connection_timeout
self.assertGreater(
remaining_window / interval,
3,
f"KEEPALIVE_INTERVAL ({interval}s) is too long: only "
f"{remaining_window/interval:.1f} keepalives would fit in the "
f"{remaining_window}s window before DVR timeout",
)
@@ -0,0 +1,195 @@
"""
Unit tests for the keepalive duration cap in StreamGenerator._stream_data_generator.
Verifies that a client held in keepalive mode is disconnected after
MAX_KEEPALIVE_DURATION seconds, and that the timer resets when real data resumes.
"""
import time
from unittest.mock import MagicMock, patch, call
from django.test import TestCase
def _make_generator(consecutive_empty=10, local_index=10, buffer_index=10):
"""Minimal StreamGenerator stub for testing _stream_data_generator logic."""
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
gen = StreamGenerator.__new__(StreamGenerator)
gen.channel_id = "00000000-0000-0000-0000-000000000099"
gen.client_id = "test-client-duration"
gen.consecutive_empty = consecutive_empty
gen.empty_reads = 0
gen.local_index = local_index
gen.bytes_sent = 0
gen.chunks_sent = 0
gen.last_yield_time = time.time()
gen.stream_start_time = time.time()
gen.last_stats_time = time.time()
gen.last_stats_bytes = 0
gen.current_rate = 0.0
gen.last_ttl_refresh = time.time()
gen.ttl_refresh_interval = 3
gen.is_owner_worker = False
gen.stream_manager = None
gen._last_health_check_time = 0.0
gen._last_health_check_result = False
gen._health_check_interval = 2.0
gen.proxy_server = None
buffer = MagicMock()
buffer.index = buffer_index
buffer.get_optimized_client_data.return_value = ([], local_index)
buffer.find_oldest_available_chunk.return_value = None
gen.buffer = buffer
return gen
class KeepaliveDurationCapTests(TestCase):
"""MAX_KEEPALIVE_DURATION cap disconnects clients stuck in keepalive mode."""
def _run_generator_to_break(self, gen, max_iterations=20):
"""Drive _stream_data_generator until it breaks or hits iteration limit."""
iterations = 0
for _ in gen._stream_data_generator():
iterations += 1
if iterations >= max_iterations:
break
return iterations
def test_cap_fires_after_max_duration_exceeded(self):
"""Generator exits when keepalive has run longer than MAX_KEEPALIVE_DURATION."""
gen = _make_generator()
with patch.object(gen, '_check_resources', return_value=True), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent') as mock_gevent, \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
MockConfig.MAX_KEEPALIVE_DURATION = 30
# First call: keepalive_start_time not yet set (returns current)
# Second call: inside the cap check — simulate time elapsed > 30s
mock_time.time.side_effect = [
1000.0, # keepalive_start_time assignment
1031.0, # cap check: 31s elapsed > 30s limit
]
packets = list(gen._stream_data_generator())
# No packets should be yielded — cap fires before yield
self.assertEqual(len(packets), 0)
def test_cap_does_not_fire_before_max_duration(self):
"""Generator yields keepalive packets while within MAX_KEEPALIVE_DURATION."""
gen = _make_generator()
call_count = 0
def time_side_effect():
nonlocal call_count
call_count += 1
# keepalive_start_time set at t=1000; cap checks always see <30s elapsed
if call_count == 1:
return 1000.0 # keepalive_start_time
return 1010.0 # always 10s elapsed — under the 30s cap
with patch.object(gen, '_check_resources', side_effect=[True, True, False]), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
MockConfig.MAX_KEEPALIVE_DURATION = 30
mock_time.time.side_effect = time_side_effect
packets = list(gen._stream_data_generator())
# Two iterations with _check_resources=True should yield two keepalive packets
self.assertGreater(len(packets), 0)
def test_timer_resets_when_real_data_resumes(self):
"""keepalive_start_time is cleared to None when real chunks are received."""
gen = _make_generator()
chunk = b'\x47' * 188
real_chunks = ([chunk], gen.local_index + 1)
no_chunks = ([], gen.local_index)
# Sequence: no data (keepalive), then real data, then stop
gen.buffer.get_optimized_client_data.side_effect = [
no_chunks, # iteration 1: keepalive
real_chunks, # iteration 2: real data — should reset timer
no_chunks, # iteration 3: keepalive again — timer restarts fresh
]
captured_start_times = []
original_gen = gen
with patch.object(gen, '_check_resources', side_effect=[True, True, True, False]), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch.object(gen, '_process_chunks', return_value=iter([chunk])), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
MockConfig.MAX_KEEPALIVE_DURATION = 300
mock_time.time.return_value = 1000.0
list(gen._stream_data_generator())
# Test passes if no exception and generator completes normally —
# if the timer were NOT reset, the second keepalive block would
# carry over the old start time rather than starting fresh.
def test_cap_uses_config_value(self):
"""Cap threshold reads MAX_KEEPALIVE_DURATION from Config, not a hardcoded value."""
gen = _make_generator()
with patch.object(gen, '_check_resources', return_value=True), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
# Set a custom cap of 60s
MockConfig.MAX_KEEPALIVE_DURATION = 60
mock_time.time.side_effect = [
1000.0, # keepalive_start_time
1050.0, # cap check: 50s elapsed — under 60s, should NOT fire
1000.0, # last_yield_time update
1070.0, # cap check on next iteration: 70s elapsed — fires
]
packets = list(gen._stream_data_generator())
# First iteration: 50s < 60s cap — one keepalive yielded
# Second iteration: 70s > 60s cap — generator exits
self.assertEqual(len(packets), 1)
+169
View File
@@ -0,0 +1,169 @@
"""Tests for the _validate_url() helper in tasks.py.
Covers:
- Rejection of None, empty, and non-string inputs
- Non-HTTP URLs pass through without network requests
- HTTP(S) URLs validated via HEAD request (2xx/3xx pass, 4xx/5xx fail)
- Network errors (timeout, connection) treated as failures
- Per-worker result cache: hits, expiry, eviction
"""
from unittest.mock import patch, MagicMock
from django.test import TestCase
from apps.channels.tasks import _validate_url, _url_validation_cache, _URL_CACHE_TTL
class ValidateUrlInputTests(TestCase):
"""Input validation — no network requests should be made."""
def setUp(self):
_url_validation_cache.clear()
def test_none_returns_false(self):
self.assertFalse(_validate_url(None))
def test_empty_string_returns_false(self):
self.assertFalse(_validate_url(""))
def test_non_string_returns_false(self):
self.assertFalse(_validate_url(123))
self.assertFalse(_validate_url(["http://example.com"]))
@patch("apps.channels.tasks.requests.head")
def test_non_http_url_returns_true_without_request(self, mock_head):
"""file:// and other non-HTTP schemes skip validation."""
self.assertTrue(_validate_url("file:///local/path.jpg"))
self.assertTrue(_validate_url("/data/images/poster.jpg"))
mock_head.assert_not_called()
class ValidateUrlNetworkTests(TestCase):
"""HTTP(S) URL validation via HEAD request."""
def setUp(self):
_url_validation_cache.clear()
@patch("apps.channels.tasks.requests.head")
def test_200_returns_true(self, mock_head):
mock_head.return_value = MagicMock(status_code=200)
self.assertTrue(_validate_url("https://example.com/poster.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_302_redirect_returns_true(self, mock_head):
mock_head.return_value = MagicMock(status_code=302)
self.assertTrue(_validate_url("https://example.com/redirect"))
@patch("apps.channels.tasks.requests.head")
def test_404_returns_false(self, mock_head):
mock_head.return_value = MagicMock(status_code=404)
self.assertFalse(_validate_url("https://dead-cdn.com/missing.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_500_returns_false(self, mock_head):
mock_head.return_value = MagicMock(status_code=500)
self.assertFalse(_validate_url("https://broken.com/error"))
@patch("apps.channels.tasks.requests.head")
def test_timeout_returns_false(self, mock_head):
import requests
mock_head.side_effect = requests.Timeout("timed out")
self.assertFalse(_validate_url("https://slow-cdn.com/poster.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_connection_error_returns_false(self, mock_head):
import requests
mock_head.side_effect = requests.ConnectionError("refused")
self.assertFalse(_validate_url("https://unreachable.com/poster.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_custom_timeout_passed_to_head(self, mock_head):
mock_head.return_value = MagicMock(status_code=200)
_validate_url("https://example.com/img.jpg", timeout=10)
mock_head.assert_called_once_with(
"https://example.com/img.jpg", timeout=10, allow_redirects=True
)
@patch("apps.channels.tasks.requests.get")
@patch("apps.channels.tasks.requests.head")
def test_405_falls_back_to_get(self, mock_head, mock_get):
"""When HEAD returns 405, fall back to a ranged GET request."""
mock_head.return_value = MagicMock(status_code=405)
mock_resp = MagicMock(status_code=200)
mock_get.return_value = mock_resp
self.assertTrue(_validate_url("https://no-head.com/poster.jpg"))
mock_get.assert_called_once()
mock_resp.close.assert_called_once()
@patch("apps.channels.tasks.requests.get")
@patch("apps.channels.tasks.requests.head")
def test_405_fallback_get_also_fails(self, mock_head, mock_get):
"""When HEAD returns 405 and GET also fails, return False."""
mock_head.return_value = MagicMock(status_code=405)
mock_get.return_value = MagicMock(status_code=403)
self.assertFalse(_validate_url("https://blocked.com/poster.jpg"))
class ValidateUrlCacheTests(TestCase):
"""Per-worker result caching."""
def setUp(self):
_url_validation_cache.clear()
@patch("apps.channels.tasks.requests.head")
def test_cache_hit_avoids_second_request(self, mock_head):
mock_head.return_value = MagicMock(status_code=200)
url = "https://cached.com/poster.jpg"
self.assertTrue(_validate_url(url))
self.assertTrue(_validate_url(url))
mock_head.assert_called_once()
@patch("apps.channels.tasks.requests.head")
def test_cache_hit_returns_false_for_failed_url(self, mock_head):
mock_head.return_value = MagicMock(status_code=404)
url = "https://dead.com/missing.jpg"
self.assertFalse(_validate_url(url))
self.assertFalse(_validate_url(url))
mock_head.assert_called_once()
@patch("apps.channels.tasks.time.monotonic")
@patch("apps.channels.tasks.requests.head")
def test_cache_expiry_triggers_new_request(self, mock_head, mock_time):
"""After TTL expires, a new HEAD request is made."""
mock_head.return_value = MagicMock(status_code=200)
url = "https://expiring.com/poster.jpg"
mock_time.return_value = 1000.0
self.assertTrue(_validate_url(url))
self.assertEqual(mock_head.call_count, 1)
# Within TTL — cache hit
mock_time.return_value = 1000.0 + _URL_CACHE_TTL - 1
self.assertTrue(_validate_url(url))
self.assertEqual(mock_head.call_count, 1)
# Past TTL — new request
mock_time.return_value = 1000.0 + _URL_CACHE_TTL + 1
self.assertTrue(_validate_url(url))
self.assertEqual(mock_head.call_count, 2)
@patch("apps.channels.tasks.time.monotonic")
@patch("apps.channels.tasks.requests.head")
def test_eviction_when_cache_exceeds_limit(self, mock_head, mock_time):
"""Expired entries are evicted when cache grows past 512 entries."""
mock_head.return_value = MagicMock(status_code=200)
# Fill cache with 513 entries at time 0
mock_time.return_value = 0.0
for i in range(513):
_url_validation_cache[f"https://fill-{i}.com/img.jpg"] = (True, 0.0)
# Advance past TTL and add one more — triggers eviction
mock_time.return_value = _URL_CACHE_TTL + 1
_validate_url("https://trigger-eviction.com/img.jpg")
# All 513 old entries expired and should be evicted
remaining = [k for k in _url_validation_cache if k.startswith("https://fill-")]
self.assertEqual(len(remaining), 0)
# The new entry should remain
self.assertIn("https://trigger-eviction.com/img.jpg", _url_validation_cache)