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
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:
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user