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,38 @@
|
||||
from django.contrib import admin
|
||||
from .models import Stream, Channel, ChannelGroup
|
||||
|
||||
@admin.register(Stream)
|
||||
class StreamAdmin(admin.ModelAdmin):
|
||||
list_display = (
|
||||
'id', # Primary Key
|
||||
'name',
|
||||
'channel_group',
|
||||
'url',
|
||||
'current_viewers',
|
||||
'updated_at',
|
||||
)
|
||||
|
||||
list_filter = ('channel_group',) # Filter by 'channel_group' (foreign key)
|
||||
|
||||
search_fields = ('id', 'name', 'url', 'channel_group__name') # Search by 'ChannelGroup' name
|
||||
|
||||
ordering = ('-updated_at',)
|
||||
|
||||
@admin.register(Channel)
|
||||
class ChannelAdmin(admin.ModelAdmin):
|
||||
list_display = (
|
||||
'id', # Primary Key
|
||||
'channel_number',
|
||||
'uuid',
|
||||
'name',
|
||||
'channel_group',
|
||||
'epg_data'
|
||||
)
|
||||
list_filter = ('channel_group',)
|
||||
search_fields = ('id', 'name', 'channel_group__name', 'epg_data') # Added 'id'
|
||||
ordering = ('channel_number',)
|
||||
|
||||
@admin.register(ChannelGroup)
|
||||
class ChannelGroupAdmin(admin.ModelAdmin):
|
||||
list_display = ('id', 'name') # Added 'id'
|
||||
search_fields = ('id', 'name') # Added 'id'
|
||||
@@ -0,0 +1,55 @@
|
||||
from django.urls import path, include
|
||||
from rest_framework.routers import DefaultRouter
|
||||
from .api_views import (
|
||||
StreamViewSet,
|
||||
ChannelViewSet,
|
||||
ChannelGroupViewSet,
|
||||
BulkDeleteStreamsAPIView,
|
||||
BulkDeleteChannelsAPIView,
|
||||
BulkDeleteLogosAPIView,
|
||||
CleanupUnusedLogosAPIView,
|
||||
LogoViewSet,
|
||||
ChannelProfileViewSet,
|
||||
UpdateChannelMembershipAPIView,
|
||||
BulkUpdateChannelMembershipAPIView,
|
||||
RecordingViewSet,
|
||||
RecurringRecordingRuleViewSet,
|
||||
GetChannelStreamsAPIView,
|
||||
SeriesRulesAPIView,
|
||||
DeleteSeriesRuleAPIView,
|
||||
EvaluateSeriesRulesAPIView,
|
||||
BulkRemoveSeriesRecordingsAPIView,
|
||||
BulkDeleteUpcomingRecordingsAPIView,
|
||||
ComskipConfigAPIView,
|
||||
)
|
||||
|
||||
app_name = 'channels' # for DRF routing
|
||||
|
||||
router = DefaultRouter()
|
||||
router.register(r'streams', StreamViewSet, basename='stream')
|
||||
router.register(r'groups', ChannelGroupViewSet, basename='channel-group')
|
||||
router.register(r'channels', ChannelViewSet, basename='channel')
|
||||
router.register(r'logos', LogoViewSet, basename='logo')
|
||||
router.register(r'profiles', ChannelProfileViewSet, basename='profile')
|
||||
router.register(r'recordings', RecordingViewSet, basename='recording')
|
||||
router.register(r'recurring-rules', RecurringRecordingRuleViewSet, basename='recurring-rule')
|
||||
|
||||
urlpatterns = [
|
||||
# Bulk delete is a single APIView, not a ViewSet
|
||||
path('streams/bulk-delete/', BulkDeleteStreamsAPIView.as_view(), name='bulk_delete_streams'),
|
||||
path('channels/bulk-delete/', BulkDeleteChannelsAPIView.as_view(), name='bulk_delete_channels'),
|
||||
path('logos/bulk-delete/', BulkDeleteLogosAPIView.as_view(), name='bulk_delete_logos'),
|
||||
path('logos/cleanup/', CleanupUnusedLogosAPIView.as_view(), name='cleanup_unused_logos'),
|
||||
path('channels/<int:channel_id>/streams/', GetChannelStreamsAPIView.as_view(), name='get_channel_streams'),
|
||||
path('profiles/<int:profile_id>/channels/<int:channel_id>/', UpdateChannelMembershipAPIView.as_view(), name='update_channel_membership'),
|
||||
path('profiles/<int:profile_id>/channels/bulk-update/', BulkUpdateChannelMembershipAPIView.as_view(), name='bulk_update_channel_membership'),
|
||||
# DVR series rules (order matters: specific routes before catch-all slug)
|
||||
path('series-rules/', SeriesRulesAPIView.as_view(), name='series_rules'),
|
||||
path('series-rules/evaluate/', EvaluateSeriesRulesAPIView.as_view(), name='evaluate_series_rules'),
|
||||
path('series-rules/bulk-remove/', BulkRemoveSeriesRecordingsAPIView.as_view(), name='bulk_remove_series_recordings'),
|
||||
path('series-rules/<path:tvg_id>/', DeleteSeriesRuleAPIView.as_view(), name='delete_series_rule'),
|
||||
path('recordings/bulk-delete-upcoming/', BulkDeleteUpcomingRecordingsAPIView.as_view(), name='bulk_delete_upcoming_recordings'),
|
||||
path('dvr/comskip-config/', ComskipConfigAPIView.as_view(), name='comskip_config'),
|
||||
]
|
||||
|
||||
urlpatterns += router.urls
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,11 @@
|
||||
from django.apps import AppConfig
|
||||
|
||||
class ChannelsConfig(AppConfig):
|
||||
default_auto_field = 'django.db.models.BigAutoField'
|
||||
name = 'apps.channels'
|
||||
verbose_name = "Channel & Stream Management"
|
||||
label = 'dispatcharr_channels'
|
||||
|
||||
def ready(self):
|
||||
# Import signals so they get registered.
|
||||
import apps.channels.signals
|
||||
@@ -0,0 +1,53 @@
|
||||
from django import forms
|
||||
from .models import Stream, Channel, ChannelGroup
|
||||
|
||||
#
|
||||
# ChannelGroup Form
|
||||
#
|
||||
class ChannelGroupForm(forms.ModelForm):
|
||||
class Meta:
|
||||
model = ChannelGroup
|
||||
fields = ['name']
|
||||
|
||||
|
||||
#
|
||||
# Channel Form
|
||||
#
|
||||
class ChannelForm(forms.ModelForm):
|
||||
# Explicitly define channel_number as FloatField to ensure decimal values work
|
||||
channel_number = forms.FloatField(
|
||||
required=False,
|
||||
widget=forms.NumberInput(attrs={'step': '0.1'}), # Allow decimal steps
|
||||
help_text="Channel number can include decimals (e.g., 1.1, 2.5)"
|
||||
)
|
||||
|
||||
channel_group = forms.ModelChoiceField(
|
||||
queryset=ChannelGroup.objects.all(),
|
||||
required=False,
|
||||
label="Channel Group",
|
||||
empty_label="--- No group ---"
|
||||
)
|
||||
|
||||
class Meta:
|
||||
model = Channel
|
||||
fields = [
|
||||
'channel_number',
|
||||
'name',
|
||||
'channel_group',
|
||||
]
|
||||
|
||||
|
||||
#
|
||||
# Example: Stream Form (optional if you want a ModelForm for Streams)
|
||||
#
|
||||
class StreamForm(forms.ModelForm):
|
||||
class Meta:
|
||||
model = Stream
|
||||
fields = [
|
||||
'name',
|
||||
'url',
|
||||
'logo_url',
|
||||
'epg_data',
|
||||
'local_file',
|
||||
'channel_group',
|
||||
]
|
||||
@@ -0,0 +1,77 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-05 22:07
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
initial = True
|
||||
|
||||
dependencies = [
|
||||
('core', '0001_initial'),
|
||||
('m3u', '0001_initial'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='ChannelGroup',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(max_length=100, unique=True)),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='Channel',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('channel_number', models.IntegerField()),
|
||||
('channel_name', models.CharField(max_length=255)),
|
||||
('logo_url', models.URLField(blank=True, max_length=2000, null=True)),
|
||||
('logo_file', models.ImageField(blank=True, null=True, upload_to='logos/')),
|
||||
('tvg_id', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('tvg_name', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('stream_profile', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='core.streamprofile')),
|
||||
('channel_group', models.ForeignKey(blank=True, help_text='Channel group this channel belongs to.', null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='dispatcharr_channels.channelgroup')),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='Stream',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(default='Default Stream', max_length=255)),
|
||||
('url', models.URLField()),
|
||||
('custom_url', models.URLField(blank=True, max_length=2000, null=True)),
|
||||
('logo_url', models.URLField(blank=True, max_length=2000, null=True)),
|
||||
('tvg_id', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('local_file', models.FileField(blank=True, null=True, upload_to='uploads/')),
|
||||
('current_viewers', models.PositiveIntegerField(default=0)),
|
||||
('updated_at', models.DateTimeField(auto_now=True)),
|
||||
('group_name', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('m3u_account', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.CASCADE, related_name='streams', to='m3u.m3uaccount')),
|
||||
('stream_profile', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='streams', to='core.streamprofile')),
|
||||
],
|
||||
options={
|
||||
'verbose_name': 'Stream',
|
||||
'verbose_name_plural': 'Streams',
|
||||
'ordering': ['-updated_at'],
|
||||
},
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='ChannelStream',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('order', models.PositiveIntegerField(default=0)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.channel')),
|
||||
('stream', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.stream')),
|
||||
],
|
||||
options={
|
||||
'ordering': ['order'],
|
||||
},
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='streams',
|
||||
field=models.ManyToManyField(blank=True, related_name='channels', through='dispatcharr_channels.ChannelStream', to='dispatcharr_channels.stream'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,27 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-16 12:21
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0001_initial'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RenameField(
|
||||
model_name='channel',
|
||||
old_name='channel_name',
|
||||
new_name='name',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='stream',
|
||||
name='url',
|
||||
),
|
||||
migrations.RenameField(
|
||||
model_name='stream',
|
||||
old_name='custom_url',
|
||||
new_name='url',
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,38 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-16 13:25
|
||||
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
# Generated by Django 5.1.6 on 2025-03-16 13:25
|
||||
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
def generate_uuids(apps, schema_editor):
|
||||
Channel = apps.get_model('dispatcharr_channels', 'Channel')
|
||||
for channel in Channel.objects.all():
|
||||
if not channel.uuid:
|
||||
channel.uuid = uuid.uuid4()
|
||||
channel.save()
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0002_rename_channel_name_channel_name_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='uuid',
|
||||
field=models.UUIDField(default=uuid.uuid4, editable=False),
|
||||
),
|
||||
migrations.RunPython(generate_uuids),
|
||||
migrations.AlterField(
|
||||
model_name='channel',
|
||||
name='uuid',
|
||||
field=models.UUIDField(default=uuid.uuid4, editable=False, unique=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-17 21:16
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0003_channel_uuid'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='is_custom',
|
||||
field=models.BooleanField(default=False, help_text='Whether this is a user-created stream or from an M3U account'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,44 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-19 16:33
|
||||
|
||||
import django.db.models.deletion
|
||||
import django.utils.timezone
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0004_stream_is_custom'),
|
||||
('m3u', '0003_create_custom_account'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='channel_group',
|
||||
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='streams', to='dispatcharr_channels.channelgroup'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='last_seen',
|
||||
field=models.DateTimeField(db_index=True, default=django.utils.timezone.now),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='channel',
|
||||
name='uuid',
|
||||
field=models.UUIDField(db_index=True, default=uuid.uuid4, editable=False, unique=True),
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='ChannelGroupM3UAccount',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('enabled', models.BooleanField(default=True)),
|
||||
('channel_group', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='m3u_account', to='dispatcharr_channels.channelgroup')),
|
||||
('m3u_account', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='channel_group', to='m3u.m3uaccount')),
|
||||
],
|
||||
options={
|
||||
'unique_together': {('channel_group', 'm3u_account')},
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,51 @@
|
||||
# In your app's migrations folder, create a new migration file
|
||||
# e.g., migrations/000X_migrate_channel_group_to_foreign_key.py
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
def migrate_channel_group(apps, schema_editor):
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
ChannelGroup = apps.get_model('dispatcharr_channels', 'ChannelGroup')
|
||||
ChannelGroupM3UAccount = apps.get_model('dispatcharr_channels', 'ChannelGroup')
|
||||
M3UAccount = apps.get_model('m3u', 'M3UAccount')
|
||||
|
||||
streams_to_update = []
|
||||
for stream in Stream.objects.all():
|
||||
# If the stream has a 'channel_group' string, try to find or create the ChannelGroup
|
||||
if stream.group_name: # group_name holds the channel group string
|
||||
channel_group_name = stream.group_name.strip()
|
||||
|
||||
# Try to find the ChannelGroup by name
|
||||
channel_group, created = ChannelGroup.objects.get_or_create(name=channel_group_name)
|
||||
|
||||
# Set the foreign key to the found or newly created ChannelGroup
|
||||
stream.channel_group = channel_group
|
||||
|
||||
streams_to_update.append(stream)
|
||||
|
||||
# If the stream has an M3U account, ensure the M3U account is linked
|
||||
if stream.m3u_account:
|
||||
ChannelGroupM3UAccount.objects.get_or_create(
|
||||
channel_group=channel_group,
|
||||
m3u_account=stream.m3u_account,
|
||||
enabled=True # Or set it to whatever the default logic is
|
||||
)
|
||||
|
||||
Stream.objects.bulk_update(streams_to_update, ['channel_group'])
|
||||
|
||||
def reverse_migration(apps, schema_editor):
|
||||
# This reverse migration would undo the changes, setting `channel_group` to `None` and clearing any relationships.
|
||||
Stream = apps.get_model('yourapp', 'Stream')
|
||||
for stream in Stream.objects.all():
|
||||
stream.channel_group = None
|
||||
stream.save()
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0005_stream_channel_group_stream_last_seen_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(migrate_channel_group, reverse_code=reverse_migration),
|
||||
]
|
||||
@@ -0,0 +1,17 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-19 16:43
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0006_migrate_stream_groups'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='stream',
|
||||
name='group_name',
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-19 18:21
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0007_remove_stream_group_name'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_hash',
|
||||
field=models.CharField(db_index=True, help_text='Unique hash for this stream from the M3U account', max_length=255, null=True, unique=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='logo_url',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,24 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-26 12:59
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0008_stream_stream_hash'),
|
||||
('epg', '0004_epgdata_epg_source_alter_epgdata_tvg_id'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='channel',
|
||||
name='tvg_name',
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='epg_data',
|
||||
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='epg.epgdata'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-01 17:36
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0009_remove_channel_tvg_name_channel_epg_data'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='custom_properties',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,35 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-01 22:14
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0010_stream_custom_properties'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='Logo',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(max_length=255)),
|
||||
('url', models.URLField(unique=True)),
|
||||
],
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='channel',
|
||||
name='logo_file',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='channel',
|
||||
name='logo_url',
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='logo',
|
||||
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='dispatcharr_channels.logo'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,33 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-02 23:27
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0011_logo_remove_channel_logo_file_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='ChannelProfile',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(max_length=100, unique=True)),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='ChannelProfileMembership',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('enabled', models.BooleanField(default=True)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.channel')),
|
||||
('channel_profile', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.channelprofile')),
|
||||
],
|
||||
options={
|
||||
'unique_together': {('channel_profile', 'channel')},
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-04 15:04
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0012_channelprofile_channelprofilemembership'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='logo',
|
||||
name='url',
|
||||
field=models.TextField(unique=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,24 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-05 22:25
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0013_alter_logo_url'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='Recording',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('start_time', models.DateTimeField()),
|
||||
('end_time', models.DateTimeField()),
|
||||
('task_id', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='recordings', to='dispatcharr_channels.channel')),
|
||||
],
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-07 16:47
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0014_recording'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='recording',
|
||||
name='custom_properties',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,38 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-18 16:21
|
||||
|
||||
from django.db import migrations, models
|
||||
from django.db.models import Count
|
||||
|
||||
def remove_duplicate_channel_streams(apps, schema_editor):
|
||||
ChannelStream = apps.get_model('dispatcharr_channels', 'ChannelStream')
|
||||
# Find duplicates by (channel, stream)
|
||||
duplicates = (
|
||||
ChannelStream.objects
|
||||
.values('channel', 'stream')
|
||||
.annotate(count=Count('id'))
|
||||
.filter(count__gt=1)
|
||||
)
|
||||
|
||||
for dupe in duplicates:
|
||||
# Get all duplicates for this pair
|
||||
dups = ChannelStream.objects.filter(
|
||||
channel=dupe['channel'],
|
||||
stream=dupe['stream']
|
||||
).order_by('id')
|
||||
|
||||
# Keep the first one, delete the rest
|
||||
dups.exclude(id=dups.first().id).delete()
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0015_recording_custom_properties'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(remove_duplicate_channel_streams),
|
||||
migrations.AddConstraint(
|
||||
model_name='channelstream',
|
||||
constraint=models.UniqueConstraint(fields=('channel', 'stream'), name='unique_channel_stream'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-21 20:47
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0016_channelstream_unique_channel_stream'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channelgroup',
|
||||
name='name',
|
||||
field=models.TextField(db_index=True, unique=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-27 14:12
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0017_alter_channelgroup_name'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='custom_properties',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-05-04 00:02
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0018_channelgroupm3uaccount_custom_properties_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='tvc_guide_stationid',
|
||||
field=models.CharField(blank=True, max_length=255, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-05-15 19:37
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0019_channel_tvc_guide_stationid'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channel',
|
||||
name='channel_number',
|
||||
field=models.FloatField(db_index=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-05-18 14:31
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0020_alter_channel_channel_number'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='user_level',
|
||||
field=models.IntegerField(default=0),
|
||||
),
|
||||
]
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
# Generated by Django 5.1.6 on 2025-07-13 23:08
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0021_channel_user_level'),
|
||||
('m3u', '0012_alter_m3uaccount_refresh_interval'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='auto_created',
|
||||
field=models.BooleanField(default=False, help_text='Whether this channel was automatically created via M3U auto channel sync'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='auto_created_by',
|
||||
field=models.ForeignKey(blank=True, help_text='The M3U account that auto-created this channel', null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='auto_created_channels', to='m3u.m3uaccount'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='auto_channel_sync',
|
||||
field=models.BooleanField(default=False, help_text='Automatically create/delete channels to match streams in this group'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='auto_sync_channel_start',
|
||||
field=models.FloatField(blank=True, help_text='Starting channel number for auto-created channels in this group', null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.1.6 on 2025-07-29 02:39
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0022_channel_auto_created_channel_auto_created_by_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_stats',
|
||||
field=models.JSONField(blank=True, help_text='JSON object containing stream statistics like video codec, resolution, etc.', null=True),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_stats_updated_at',
|
||||
field=models.DateTimeField(blank=True, db_index=True, help_text='When stream statistics were last updated', null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,19 @@
|
||||
# Generated by Django 5.2.4 on 2025-08-22 20:14
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0023_stream_stream_stats_stream_stream_stats_updated_at'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='channel_group',
|
||||
field=models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='m3u_accounts', to='dispatcharr_channels.channelgroup'),
|
||||
),
|
||||
]
|
||||
+28
@@ -0,0 +1,28 @@
|
||||
# Generated by Django 5.2.4 on 2025-09-02 14:30
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0024_alter_channelgroupm3uaccount_channel_group'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='custom_properties',
|
||||
field=models.JSONField(blank=True, default=dict, null=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='recording',
|
||||
name='custom_properties',
|
||||
field=models.JSONField(blank=True, default=dict, null=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='custom_properties',
|
||||
field=models.JSONField(blank=True, default=dict, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# Generated by Django 5.0.14 on 2025-09-18 14:56
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0025_alter_channelgroupm3uaccount_custom_properties_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='RecurringRecordingRule',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('days_of_week', models.JSONField(default=list)),
|
||||
('start_time', models.TimeField()),
|
||||
('end_time', models.TimeField()),
|
||||
('enabled', models.BooleanField(default=True)),
|
||||
('name', models.CharField(blank=True, max_length=255)),
|
||||
('created_at', models.DateTimeField(auto_now_add=True)),
|
||||
('updated_at', models.DateTimeField(auto_now=True)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='recurring_rules', to='dispatcharr_channels.channel')),
|
||||
],
|
||||
options={
|
||||
'ordering': ['channel', 'start_time'],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.2.4 on 2025-10-05 20:50
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0026_recurringrecordingrule'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='recurringrecordingrule',
|
||||
name='end_date',
|
||||
field=models.DateField(blank=True, null=True),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='recurringrecordingrule',
|
||||
name='start_date',
|
||||
field=models.DateField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,25 @@
|
||||
# Generated by Django 5.2.4 on 2025-10-06 22:55
|
||||
|
||||
import django.utils.timezone
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0027_recurringrecordingrule_end_date_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='created_at',
|
||||
field=models.DateTimeField(auto_now_add=True, default=django.utils.timezone.now, help_text='Timestamp when this channel was created'),
|
||||
preserve_default=False,
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='updated_at',
|
||||
field=models.DateTimeField(auto_now=True, help_text='Timestamp when this channel was last updated'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,54 @@
|
||||
# Generated migration to backfill stream_hash for existing custom streams
|
||||
|
||||
from django.db import migrations
|
||||
import hashlib
|
||||
|
||||
|
||||
def backfill_custom_stream_hashes(apps, schema_editor):
|
||||
"""
|
||||
Generate stream_hash for all custom streams that don't have one.
|
||||
Uses stream ID to create a stable hash that won't change when name/url is edited.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
|
||||
custom_streams_without_hash = Stream.objects.filter(
|
||||
is_custom=True,
|
||||
stream_hash__isnull=True
|
||||
)
|
||||
|
||||
updated_count = 0
|
||||
for stream in custom_streams_without_hash:
|
||||
# Generate a stable hash using the stream's ID
|
||||
# This ensures the hash never changes even if name/url is edited
|
||||
unique_string = f"custom_stream_{stream.id}"
|
||||
stream.stream_hash = hashlib.sha256(unique_string.encode()).hexdigest()
|
||||
stream.save(update_fields=['stream_hash'])
|
||||
updated_count += 1
|
||||
|
||||
if updated_count > 0:
|
||||
print(f"Backfilled stream_hash for {updated_count} custom streams")
|
||||
else:
|
||||
print("No custom streams needed stream_hash backfill")
|
||||
|
||||
|
||||
def reverse_backfill(apps, schema_editor):
|
||||
"""
|
||||
Reverse migration - clear stream_hash for custom streams.
|
||||
Note: This will break preview functionality for custom streams.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
|
||||
custom_streams = Stream.objects.filter(is_custom=True)
|
||||
count = custom_streams.update(stream_hash=None)
|
||||
print(f"Cleared stream_hash for {count} custom streams")
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0028_channel_created_at_channel_updated_at'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(backfill_custom_stream_hashes, reverse_backfill),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.2.4 on 2025-10-28 20:00
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0029_backfill_custom_stream_hashes'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='url',
|
||||
field=models.URLField(blank=True, max_length=4096, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,29 @@
|
||||
# Generated by Django 5.2.9 on 2026-01-09 18:19
|
||||
|
||||
import django.utils.timezone
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0030_alter_stream_url'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='is_stale',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this group relationship is stale (not seen in recent refresh, pending deletion)'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='last_seen',
|
||||
field=models.DateTimeField(db_index=True, default=django.utils.timezone.now, help_text='Last time this group was seen in the M3U source during a refresh'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='is_stale',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this stream is stale (not seen in recent refresh, pending deletion)'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.2.9 on 2026-01-17 16:56
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0031_channelgroupm3uaccount_is_stale_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='is_adult',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this channel contains adult content'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='is_adult',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this stream contains adult content'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,205 @@
|
||||
# Generated by Django - Add stream_id and channel_number fields with data migration
|
||||
|
||||
from django.db import migrations, models
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def populate_fields_and_rehash(apps, schema_editor):
|
||||
"""
|
||||
Populate stream_id and stream_chno from custom_properties for XC account streams,
|
||||
populate stream_chno from tvg-chno for standard M3U accounts,
|
||||
then rehash XC streams using stable hash keys.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
M3UAccount = apps.get_model('m3u', 'M3UAccount')
|
||||
CoreSettings = apps.get_model('core', 'CoreSettings')
|
||||
|
||||
# Get hash keys from settings
|
||||
try:
|
||||
stream_settings = CoreSettings.objects.get(key='stream_settings')
|
||||
hash_key_str = stream_settings.value.get('m3u_hash_key', '') if stream_settings.value else ''
|
||||
keys = [k.strip() for k in hash_key_str.split(',') if k.strip()] if hash_key_str else []
|
||||
except CoreSettings.DoesNotExist:
|
||||
keys = []
|
||||
|
||||
logger.info(f"Using hash keys: {keys}")
|
||||
|
||||
# Get XC account IDs
|
||||
xc_account_ids = set(
|
||||
M3UAccount.objects.filter(account_type='XC').values_list('id', flat=True)
|
||||
)
|
||||
|
||||
logger.info(f"Found {len(xc_account_ids)} XC accounts")
|
||||
|
||||
# Track hash collisions for XC streams
|
||||
hash_map = {} # new_hash -> stream_id
|
||||
duplicates_to_delete = []
|
||||
|
||||
# Process all streams in batches
|
||||
batch_size = 1000
|
||||
processed = 0
|
||||
updated = 0
|
||||
|
||||
total_count = Stream.objects.count()
|
||||
logger.info(f"Processing {total_count} total streams")
|
||||
|
||||
streams_to_update = []
|
||||
|
||||
for stream in Stream.objects.select_related('channel_group', 'm3u_account').iterator(chunk_size=batch_size):
|
||||
processed += 1
|
||||
needs_update = False
|
||||
|
||||
custom_props = stream.custom_properties or {}
|
||||
is_xc = stream.m3u_account_id in xc_account_ids if stream.m3u_account_id else False
|
||||
|
||||
# Extract stream_id (XC accounts only)
|
||||
if is_xc and isinstance(custom_props, dict):
|
||||
provider_stream_id = custom_props.get('stream_id')
|
||||
if provider_stream_id:
|
||||
try:
|
||||
stream.stream_id = int(provider_stream_id)
|
||||
needs_update = True
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# Extract stream_chno
|
||||
channel_num = None
|
||||
if isinstance(custom_props, dict):
|
||||
if is_xc:
|
||||
# XC accounts use 'num'
|
||||
channel_num = custom_props.get('num')
|
||||
else:
|
||||
# Standard M3U accounts use 'tvg-chno' or 'channel-number' (case insensitive check)
|
||||
for key in ['tvg-chno', 'TVG-CHNO', 'tvg-Chno', 'Tvg-Chno', 'channel-number', 'Channel-Number', 'CHANNEL-NUMBER']:
|
||||
if key in custom_props:
|
||||
channel_num = custom_props.get(key)
|
||||
break
|
||||
|
||||
if channel_num is not None:
|
||||
try:
|
||||
stream.stream_chno = float(channel_num)
|
||||
needs_update = True
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# Rehash XC streams only when 'url' is in hash keys (otherwise hash wouldn't change)
|
||||
if is_xc and stream.stream_id and keys and 'url' in keys:
|
||||
# For XC accounts, use stream_id instead of url when 'url' is in the hash keys
|
||||
# This ensures credential/URL changes don't break stream identity
|
||||
effective_url = stream.stream_id
|
||||
|
||||
# Get group name
|
||||
group_name = stream.channel_group.name if stream.channel_group else None
|
||||
|
||||
# Build hash parts
|
||||
stream_parts = {
|
||||
"name": stream.name,
|
||||
"url": effective_url,
|
||||
"tvg_id": stream.tvg_id,
|
||||
"m3u_id": stream.m3u_account_id,
|
||||
"group": group_name
|
||||
}
|
||||
hash_parts = {key: stream_parts[key] for key in keys if key in stream_parts}
|
||||
|
||||
# When using stream_id instead of URL, we MUST include m3u_id to prevent
|
||||
# collisions across different XC accounts (stream_id is only unique per account)
|
||||
if 'm3u_id' not in hash_parts:
|
||||
hash_parts['m3u_id'] = stream.m3u_account_id
|
||||
|
||||
# Generate hash
|
||||
serialized_obj = json.dumps(hash_parts, sort_keys=True)
|
||||
new_hash = hashlib.sha256(serialized_obj.encode()).hexdigest()
|
||||
|
||||
# Check for collisions
|
||||
if new_hash in hash_map:
|
||||
# Duplicate - mark for deletion (keep the first one)
|
||||
duplicates_to_delete.append(stream.id)
|
||||
continue
|
||||
|
||||
hash_map[new_hash] = stream.id
|
||||
stream.stream_hash = new_hash
|
||||
needs_update = True
|
||||
|
||||
if needs_update:
|
||||
streams_to_update.append(stream)
|
||||
updated += 1
|
||||
|
||||
# Bulk update in batches
|
||||
if len(streams_to_update) >= batch_size:
|
||||
Stream.objects.bulk_update(
|
||||
streams_to_update,
|
||||
['stream_id', 'stream_chno', 'stream_hash'],
|
||||
batch_size=500
|
||||
)
|
||||
logger.info(f"Updated batch: {processed}/{total_count} streams processed")
|
||||
streams_to_update = []
|
||||
|
||||
# Final batch
|
||||
if streams_to_update:
|
||||
Stream.objects.bulk_update(
|
||||
streams_to_update,
|
||||
['stream_id', 'stream_chno', 'stream_hash'],
|
||||
batch_size=500
|
||||
)
|
||||
|
||||
# Delete duplicates if any
|
||||
if duplicates_to_delete:
|
||||
logger.warning(f"Deleting {len(duplicates_to_delete)} duplicate streams due to hash collisions")
|
||||
Stream.objects.filter(id__in=duplicates_to_delete).delete()
|
||||
|
||||
logger.info(f"Migration complete: {updated} streams updated, {len(duplicates_to_delete)} duplicates removed")
|
||||
|
||||
|
||||
def reverse_migration(apps, schema_editor):
|
||||
"""
|
||||
Reverse migration - clear fields but don't attempt to reverse hash changes.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
Stream.objects.all().update(stream_id=None, stream_chno=None)
|
||||
logger.info("Cleared stream_id and stream_chno fields. Note: stream hashes were not reverted.")
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0032_channel_is_adult_stream_is_adult'),
|
||||
('m3u', '0018_add_profile_custom_properties'),
|
||||
('core', '0020_change_coresettings_value_to_jsonfield'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
# Schema changes - add fields WITHOUT indexes first
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_id',
|
||||
field=models.IntegerField(
|
||||
blank=True,
|
||||
help_text='Provider stream ID (e.g., XC stream_id) for stable identity across credential changes',
|
||||
null=True,
|
||||
),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_chno',
|
||||
field=models.FloatField(
|
||||
blank=True,
|
||||
help_text='Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1',
|
||||
null=True,
|
||||
),
|
||||
),
|
||||
# Data migration (may delete duplicates, which would conflict with pending index creation)
|
||||
migrations.RunPython(populate_fields_and_rehash, reverse_migration),
|
||||
# Add indexes AFTER data migration completes
|
||||
migrations.AddIndex(
|
||||
model_name='stream',
|
||||
index=models.Index(fields=['stream_id'], name='dispatcharr_stream_id_idx'),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='stream',
|
||||
index=models.Index(fields=['stream_chno'], name='dispatcharr_stream_chno_idx'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# Generated by Django 5.2.9 on 2026-02-01 03:21
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0033_stream_id_stream_chno'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveIndex(
|
||||
model_name='stream',
|
||||
name='dispatcharr_stream_id_idx',
|
||||
),
|
||||
migrations.RemoveIndex(
|
||||
model_name='stream',
|
||||
name='dispatcharr_stream_chno_idx',
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='stream_chno',
|
||||
field=models.FloatField(blank=True, db_index=True, help_text='Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1', null=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='stream_id',
|
||||
field=models.IntegerField(blank=True, db_index=True, help_text='Provider stream ID (e.g., XC stream_id) for stable identity across credential changes', null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,988 @@
|
||||
from django.db import models
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.conf import settings
|
||||
from core.models import StreamProfile, CoreSettings
|
||||
from core.utils import RedisClient
|
||||
from apps.proxy.ts_proxy.redis_keys import RedisKeys
|
||||
from apps.proxy.ts_proxy.constants import ChannelMetadataField
|
||||
import logging
|
||||
import uuid
|
||||
from django.utils import timezone
|
||||
import hashlib
|
||||
import json
|
||||
from apps.epg.models import EPGData
|
||||
from apps.accounts.models import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# If you have an M3UAccount model in apps.m3u, you can still import it:
|
||||
from apps.m3u.models import M3UAccount
|
||||
|
||||
|
||||
# Add fallback functions if Redis isn't available
|
||||
def get_total_viewers(channel_id):
|
||||
"""Get viewer count from Redis or return 0 if Redis isn't available"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
try:
|
||||
return int(redis_client.get(f"channel:{channel_id}:viewers") or 0)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
class ChannelGroup(models.Model):
|
||||
name = models.TextField(unique=True, db_index=True)
|
||||
|
||||
def related_channels(self):
|
||||
# local import if needed to avoid cyc. Usually fine in a single file though
|
||||
return Channel.objects.filter(channel_group=self)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def bulk_create_and_fetch(cls, objects):
|
||||
# Perform the bulk create operation
|
||||
cls.objects.bulk_create(objects)
|
||||
|
||||
# Use a unique field to fetch the created objects (assuming 'name' is unique)
|
||||
created_objects = cls.objects.filter(name__in=[obj.name for obj in objects])
|
||||
|
||||
return created_objects
|
||||
|
||||
|
||||
class Stream(models.Model):
|
||||
"""
|
||||
Represents a single stream (e.g. from an M3U source or custom URL).
|
||||
"""
|
||||
|
||||
name = models.CharField(max_length=255, default="Default Stream")
|
||||
url = models.URLField(max_length=4096, blank=True, null=True)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount,
|
||||
on_delete=models.CASCADE,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
logo_url = models.TextField(blank=True, null=True)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
local_file = models.FileField(upload_to="uploads/", blank=True, null=True)
|
||||
current_viewers = models.PositiveIntegerField(default=0)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
null=True,
|
||||
blank=True,
|
||||
on_delete=models.SET_NULL,
|
||||
related_name="streams",
|
||||
)
|
||||
is_custom = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this is a user-created stream or from an M3U account",
|
||||
)
|
||||
stream_hash = models.CharField(
|
||||
max_length=255,
|
||||
null=True,
|
||||
unique=True,
|
||||
help_text="Unique hash for this stream from the M3U account",
|
||||
db_index=True,
|
||||
)
|
||||
last_seen = models.DateTimeField(db_index=True, default=timezone.now)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream is stale (not seen in recent refresh, pending deletion)"
|
||||
)
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream contains adult content"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
stream_id = models.IntegerField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider stream ID (e.g., XC stream_id) for stable identity across credential changes"
|
||||
)
|
||||
stream_chno = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1"
|
||||
)
|
||||
|
||||
# Stream statistics fields
|
||||
stream_stats = models.JSONField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="JSON object containing stream statistics like video codec, resolution, etc."
|
||||
)
|
||||
stream_stats_updated_at = models.DateTimeField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="When stream statistics were last updated",
|
||||
db_index=True
|
||||
)
|
||||
|
||||
class Meta:
|
||||
# If you use m3u_account, you might do unique_together = ('name','url','m3u_account')
|
||||
verbose_name = "Stream"
|
||||
verbose_name_plural = "Streams"
|
||||
ordering = ["-updated_at"]
|
||||
|
||||
def __str__(self):
|
||||
return self.name or self.url or f"Stream ID {self.id}"
|
||||
|
||||
@classmethod
|
||||
def generate_hash_key(cls, name, url, tvg_id, keys=None, m3u_id=None, group=None,
|
||||
account_type=None, stream_id=None):
|
||||
if keys is None:
|
||||
keys = CoreSettings.get_m3u_hash_key().split(",")
|
||||
|
||||
# For XC accounts, use stream_id instead of url when 'url' is in the hash keys
|
||||
# This ensures credential/URL changes don't break stream identity
|
||||
effective_url = url
|
||||
use_stream_id = account_type == 'XC' and stream_id and 'url' in keys
|
||||
if use_stream_id:
|
||||
effective_url = stream_id
|
||||
|
||||
stream_parts = {"name": name, "url": effective_url, "tvg_id": tvg_id, "m3u_id": m3u_id, "group": group}
|
||||
|
||||
hash_parts = {key: stream_parts[key] for key in keys if key in stream_parts}
|
||||
|
||||
# When using stream_id instead of URL, we MUST include m3u_id to prevent
|
||||
# collisions across different XC accounts (stream_id is only unique per account)
|
||||
if use_stream_id and 'm3u_id' not in hash_parts:
|
||||
hash_parts['m3u_id'] = m3u_id
|
||||
|
||||
# Serialize and hash the dictionary
|
||||
serialized_obj = json.dumps(
|
||||
hash_parts, sort_keys=True
|
||||
) # sort_keys ensures consistent ordering
|
||||
hash_object = hashlib.sha256(serialized_obj.encode())
|
||||
return hash_object.hexdigest()
|
||||
|
||||
@classmethod
|
||||
def update_or_create_by_hash(cls, hash_value, **fields_to_update):
|
||||
try:
|
||||
# Try to find the Stream object with the given hash
|
||||
stream = cls.objects.get(stream_hash=hash_value)
|
||||
# If it exists, update the fields
|
||||
for field, value in fields_to_update.items():
|
||||
setattr(stream, field, value)
|
||||
stream.save() # Save the updated object
|
||||
return stream, False # False means it was updated, not created
|
||||
except cls.DoesNotExist:
|
||||
# If it doesn't exist, create a new object with the given hash
|
||||
fields_to_update["stream_hash"] = (
|
||||
hash_value # Make sure the hash field is set
|
||||
)
|
||||
stream = cls.objects.create(**fields_to_update)
|
||||
return stream, True # True means it was created
|
||||
|
||||
def get_stream_profile(self):
|
||||
"""
|
||||
Get the stream profile for this stream.
|
||||
Uses the stream's own profile if set, otherwise returns the default.
|
||||
"""
|
||||
if self.stream_profile:
|
||||
return self.stream_profile
|
||||
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available profile for this stream and reserves a connection slot.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
profile_id = redis_client.get(f"stream_profile:{self.id}")
|
||||
if profile_id:
|
||||
profile_id = int(profile_id)
|
||||
return self.id, profile_id, None
|
||||
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = self.m3u_account
|
||||
m3u_profiles = m3u_account.profiles.all()
|
||||
default_profile = next((obj for obj in m3u_profiles if obj.is_default), None)
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
logger.info(profile)
|
||||
# Skip inactive profiles
|
||||
if profile.is_active == False:
|
||||
continue
|
||||
|
||||
# Atomic slot reservation: INCR first, check, rollback if over capacity
|
||||
if profile.max_streams == 0:
|
||||
reserved = True
|
||||
else:
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
if new_count <= profile.max_streams:
|
||||
reserved = True
|
||||
else:
|
||||
redis_client.decr(profile_connections_key)
|
||||
reserved = False
|
||||
|
||||
if reserved:
|
||||
redis_client.set(f"channel_stream:{self.id}", self.id)
|
||||
redis_client.set(f"stream_profile:{self.id}", profile.id)
|
||||
return self.id, profile.id, None
|
||||
|
||||
return None, None, "All active M3U profiles have reached maximum connection limits"
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = self.id
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not profile_id:
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: no profile found in "
|
||||
f"stream_profile:{stream_id}"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
|
||||
profile_id = int(profile_id)
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: found profile_id={profile_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
class ChannelManager(models.Manager):
|
||||
def active(self):
|
||||
return self.all()
|
||||
|
||||
|
||||
class Channel(models.Model):
|
||||
channel_number = models.FloatField(db_index=True)
|
||||
name = models.CharField(max_length=255)
|
||||
logo = models.ForeignKey(
|
||||
"Logo",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
# M2M to Stream now in the same file
|
||||
streams = models.ManyToManyField(
|
||||
Stream, blank=True, through="ChannelStream", related_name="channels"
|
||||
)
|
||||
|
||||
channel_group = models.ForeignKey(
|
||||
"ChannelGroup",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
help_text="Channel group this channel belongs to.",
|
||||
)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
tvc_guide_stationid = models.CharField(max_length=255, blank=True, null=True)
|
||||
|
||||
epg_data = models.ForeignKey(
|
||||
EPGData,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
uuid = models.UUIDField(
|
||||
default=uuid.uuid4, editable=False, unique=True, db_index=True
|
||||
)
|
||||
|
||||
user_level = models.IntegerField(default=0)
|
||||
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this channel contains adult content"
|
||||
)
|
||||
|
||||
auto_created = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this channel was automatically created via M3U auto channel sync"
|
||||
)
|
||||
auto_created_by = models.ForeignKey(
|
||||
"m3u.M3UAccount",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="auto_created_channels",
|
||||
help_text="The M3U account that auto-created this channel"
|
||||
)
|
||||
|
||||
created_at = models.DateTimeField(
|
||||
auto_now_add=True,
|
||||
help_text="Timestamp when this channel was created"
|
||||
)
|
||||
updated_at = models.DateTimeField(
|
||||
auto_now=True,
|
||||
help_text="Timestamp when this channel was last updated"
|
||||
)
|
||||
|
||||
def clean(self):
|
||||
# Enforce unique channel_number within a given group
|
||||
existing = Channel.objects.filter(
|
||||
channel_number=self.channel_number, channel_group=self.channel_group
|
||||
).exclude(id=self.id)
|
||||
if existing.exists():
|
||||
raise ValidationError(
|
||||
f"Channel number {self.channel_number} already exists in group {self.channel_group}."
|
||||
)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_number} - {self.name}"
|
||||
|
||||
@classmethod
|
||||
def get_next_available_channel_number(cls, starting_from=1):
|
||||
used_numbers = set(cls.objects.all().values_list("channel_number", flat=True))
|
||||
n = starting_from
|
||||
while n in used_numbers:
|
||||
n += 1
|
||||
return n
|
||||
|
||||
# @TODO: honor stream's stream profile
|
||||
def get_stream_profile(self):
|
||||
stream_profile = self.stream_profile
|
||||
if not stream_profile:
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def _account_active_connections(self, redis_client, m3u_account, cache: dict[int, int]) -> int:
|
||||
"""
|
||||
Return active connection count for an M3U account across all its active profiles.
|
||||
Uses Redis `profile_connections:*`, which is shared by live and VOD proxy paths.
|
||||
"""
|
||||
if not m3u_account:
|
||||
return 0
|
||||
account_id = int(m3u_account.id)
|
||||
if account_id in cache:
|
||||
return cache[account_id]
|
||||
|
||||
total = 0
|
||||
try:
|
||||
for profile in m3u_account.profiles.filter(is_active=True):
|
||||
total += int(redis_client.get(f"profile_connections:{profile.id}") or 0)
|
||||
except Exception:
|
||||
total = 0
|
||||
cache[account_id] = total
|
||||
return total
|
||||
|
||||
def _pick_channel_to_preempt(
|
||||
self,
|
||||
profile_id,
|
||||
requester_level,
|
||||
redis_client,
|
||||
exclude_channel_ids=None,
|
||||
cooldown_seconds=30,
|
||||
):
|
||||
"""
|
||||
Pick the lowest-impact channel to terminate on the given profile.
|
||||
Returns: Optional[int] channel_id to preempt
|
||||
"""
|
||||
exclude_channel_ids = set(exclude_channel_ids or [])
|
||||
candidates = []
|
||||
|
||||
# 1) Try to get active channel IDs for this profile from an index set if available
|
||||
ch_set_key = f"ts_proxy:profile:{profile_id}:channels"
|
||||
try:
|
||||
ch_ids = { (int(x) if not isinstance(x, int) else x) for x in (redis_client.smembers(ch_set_key) or set()) }
|
||||
except Exception:
|
||||
ch_ids = set()
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
# 2) Fallback: scan metadata keys and filter by m3u_profile == profile_id
|
||||
if not ch_ids:
|
||||
cursor = 0
|
||||
pattern = "ts_proxy:channel:*:metadata"
|
||||
while True:
|
||||
cursor, keys = redis_client.scan(cursor=cursor, match=pattern, count=500)
|
||||
if keys:
|
||||
# Prefer HGET m3u_profile if metadata is a hash
|
||||
pipe = redis_client.pipeline()
|
||||
for k in keys:
|
||||
pipe.hget(k, "m3u_profile")
|
||||
prof_vals = pipe.execute()
|
||||
for k, prof_val in zip(keys, prof_vals):
|
||||
try:
|
||||
pid = int(prof_val) if prof_val is not None else None
|
||||
except Exception:
|
||||
pid = None
|
||||
|
||||
if pid == profile_id:
|
||||
parts = k.split(":") # ts_proxy:channel:{id}:metadata
|
||||
if len(parts) >= 4:
|
||||
try:
|
||||
ch_ids.add(int(parts[2]))
|
||||
except Exception:
|
||||
pass
|
||||
if cursor == 0:
|
||||
break
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
if not ch_ids:
|
||||
return None
|
||||
|
||||
# 3) Score candidates
|
||||
for ch_id in ch_ids:
|
||||
if ch_id in exclude_channel_ids:
|
||||
continue
|
||||
|
||||
# Skip if recently preempted
|
||||
last_preempt_key = f"ts_proxy:channel:{ch_id}:last_preempt"
|
||||
try:
|
||||
last_preempt = float(redis_client.get(last_preempt_key) or 0.0)
|
||||
except Exception:
|
||||
last_preempt = 0.0
|
||||
if last_preempt and (time.time() - last_preempt) < cooldown_seconds:
|
||||
continue
|
||||
|
||||
# Clients and their levels
|
||||
clients_key = f"ts_proxy:channel:{ch_id}:clients"
|
||||
member_ids = list(redis_client.smembers(clients_key) or [])
|
||||
viewer_count = len(member_ids)
|
||||
max_viewer_level = 0
|
||||
if viewer_count:
|
||||
pipe = redis_client.pipeline()
|
||||
for cid in member_ids:
|
||||
pipe.hget(f"ts_proxy:channel:{ch_id}:clients:{cid}", "user_level")
|
||||
levels_raw = pipe.execute()
|
||||
levels = []
|
||||
for lv in levels_raw:
|
||||
try:
|
||||
levels.append(int(lv or 0))
|
||||
except Exception:
|
||||
levels.append(0)
|
||||
max_viewer_level = max(levels or [0])
|
||||
|
||||
# Only preempt if requester strictly outranks this channel's viewers
|
||||
if requester_level <= max_viewer_level:
|
||||
continue
|
||||
|
||||
# Metadata (protected/recording/started_at_ts)
|
||||
meta_key = f"ts_proxy:channel:{ch_id}:metadata"
|
||||
try:
|
||||
protected, recording, started_at_ts = redis_client.hmget(
|
||||
meta_key, "protected", "recording", "started_at_ts"
|
||||
)
|
||||
except Exception:
|
||||
protected = recording = started_at_ts = None
|
||||
|
||||
protected = str(protected or "0") in ("1", "true", "True")
|
||||
recording = str(recording or "0") in ("1", "true", "True")
|
||||
if protected or recording:
|
||||
continue
|
||||
|
||||
try:
|
||||
started_at_ts = float(started_at_ts) if started_at_ts is not None else None
|
||||
except Exception:
|
||||
started_at_ts = None
|
||||
if started_at_ts is None:
|
||||
started_at_ts = time.time() # treat unknown as newest
|
||||
|
||||
# Score: lower is safer to terminate
|
||||
has_viewers = 1 if viewer_count > 0 else 0
|
||||
score = (has_viewers, max_viewer_level, viewer_count, started_at_ts)
|
||||
candidates.append((score, ch_id))
|
||||
|
||||
logger.debug("Candidate channels after scoring:")
|
||||
logger.debug(candidates)
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
victim_id = candidates[0][1]
|
||||
|
||||
# Mark preempt timestamp to avoid thrashing
|
||||
try:
|
||||
redis_client.set(f"ts_proxy:channel:{victim_id}:last_preempt", str(time.time()), ex=3600)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return victim_id
|
||||
|
||||
def _check_and_reserve_profile_slot(self, profile, redis_client):
|
||||
"""
|
||||
Atomically check and reserve a connection slot for the given profile.
|
||||
|
||||
Uses an INCR-first-then-check pattern to eliminate the TOCTOU race
|
||||
condition where separate GET + check + INCR operations could allow
|
||||
concurrent requests to both pass the capacity check.
|
||||
|
||||
For profiles with max_streams=0 (unlimited), no reservation is needed.
|
||||
|
||||
Args:
|
||||
profile: M3UAccountProfile instance
|
||||
redis_client: Redis client instance
|
||||
|
||||
Returns:
|
||||
tuple: (reserved: bool, current_count: int)
|
||||
"""
|
||||
if profile.max_streams == 0:
|
||||
return (True, 0)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
|
||||
# Atomically increment first — this is a single Redis command
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
|
||||
if new_count <= profile.max_streams:
|
||||
return (True, new_count)
|
||||
|
||||
# Over capacity — roll back the increment
|
||||
redis_client.decr(profile_connections_key)
|
||||
return (False, new_count - 1)
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available stream for the requested channel and returns the selected stream and profile.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
error_reason = None
|
||||
|
||||
# Check if this channel has any streams
|
||||
if not self.streams.exists():
|
||||
error_reason = "No streams assigned to channel"
|
||||
return None, None, error_reason
|
||||
|
||||
# Check if a stream is already active for this channel
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if stream_id_bytes:
|
||||
try:
|
||||
stream_id = int(stream_id_bytes)
|
||||
profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id_bytes:
|
||||
try:
|
||||
profile_id = int(profile_id_bytes)
|
||||
return stream_id, profile_id, None
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid profile ID retrieved from Redis: {profile_id_bytes}"
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid stream ID retrieved from Redis: {stream_id_bytes}"
|
||||
)
|
||||
|
||||
# No existing active stream, attempt to assign a new one
|
||||
has_streams_but_maxed_out = False
|
||||
has_active_profiles = False
|
||||
account_load_cache: dict[int, int] = {}
|
||||
|
||||
ordered_streams = list(self.streams.all().order_by("channelstream__order"))
|
||||
original_order = {stream.id: idx for idx, stream in enumerate(ordered_streams)}
|
||||
ordered_streams.sort(
|
||||
key=lambda s: (
|
||||
self._account_active_connections(redis_client, s.m3u_account, account_load_cache),
|
||||
original_order.get(s.id, 0),
|
||||
)
|
||||
)
|
||||
|
||||
# Iterate through channel streams and their profiles
|
||||
for stream in ordered_streams:
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = stream.m3u_account
|
||||
if not m3u_account:
|
||||
logger.debug(f"Stream {stream.id} has no M3U account")
|
||||
continue
|
||||
if m3u_account.is_active == False:
|
||||
logger.debug(f"M3U account {m3u_account.id} is inactive, skipping.")
|
||||
continue
|
||||
|
||||
m3u_profiles = m3u_account.profiles.filter(is_active=True)
|
||||
default_profile = next(
|
||||
(obj for obj in m3u_profiles if obj.is_default), None
|
||||
)
|
||||
|
||||
if not default_profile:
|
||||
logger.debug(f"M3U account {m3u_account.id} has no active default profile")
|
||||
continue
|
||||
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
has_active_profiles = True
|
||||
|
||||
# Atomically check and reserve a slot (INCR-first pattern)
|
||||
reserved, current_count = self._check_and_reserve_profile_slot(
|
||||
profile, redis_client
|
||||
)
|
||||
|
||||
if reserved:
|
||||
# Slot reserved — assign stream to this channel
|
||||
redis_client.set(f"channel_stream:{self.id}", stream.id)
|
||||
redis_client.set(f"stream_profile:{stream.id}", profile.id)
|
||||
|
||||
return (
|
||||
stream.id,
|
||||
profile.id,
|
||||
None,
|
||||
) # Return newly assigned stream and matched profile
|
||||
else:
|
||||
# At capacity: try to preempt a lower-impact channel on this profile
|
||||
victim_channel_id = self._pick_channel_to_preempt(
|
||||
profile_id=profile.id,
|
||||
requester_level=requester.user_level if requester else 100,
|
||||
redis_client=redis_client,
|
||||
exclude_channel_ids=None,
|
||||
)
|
||||
if victim_channel_id:
|
||||
logger.info(f"Preempting channel {victim_channel_id} for new stream on profile {profile.id}")
|
||||
# return self.id, profile.id, victim_channel_id
|
||||
|
||||
# This profile is at max connections
|
||||
has_streams_but_maxed_out = True
|
||||
logger.debug(
|
||||
f"Profile {profile.id} at max connections: "
|
||||
f"{current_count}/{profile.max_streams}"
|
||||
)
|
||||
|
||||
# No available streams - determine specific reason
|
||||
if has_streams_but_maxed_out:
|
||||
error_reason = "All active M3U profiles have reached maximum connection limits"
|
||||
elif has_active_profiles:
|
||||
error_reason = "No compatible active profile found for any assigned stream"
|
||||
else:
|
||||
error_reason = "No active profiles found for any assigned stream"
|
||||
|
||||
return None, None, error_reason
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no stream/profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id:
|
||||
# Primary key missing — try metadata hash fallback.
|
||||
# The proxy may have already cleaned up channel_stream/stream_profile
|
||||
# keys, but the metadata hash can still have the stream_id and profile.
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_stream_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.STREAM_ID
|
||||
)
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
|
||||
if meta_stream_id and meta_profile_id:
|
||||
stream_id = int(meta_stream_id)
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered stream_id={stream_id}, "
|
||||
f"profile_id={profile_id} from metadata fallback"
|
||||
)
|
||||
# Clean up any remaining keys
|
||||
redis_client.delete(f"channel_stream:{self.id}")
|
||||
redis_client.delete(f"stream_profile:{stream_id}")
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# won't find them and DECR again
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
current_count = int(
|
||||
redis_client.get(profile_connections_key) or 0
|
||||
)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
return True
|
||||
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: no stream info found in primary keys "
|
||||
f"or metadata fallback"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"channel_stream:{self.id}") # Remove active stream
|
||||
|
||||
stream_id = int(stream_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found stream_id={stream_id} for "
|
||||
f"channel_stream:{self.id}"
|
||||
)
|
||||
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id:
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
profile_id = int(profile_id)
|
||||
else:
|
||||
# stream_profile key missing — try metadata hash fallback
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
if meta_profile_id:
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered profile_id={profile_id} "
|
||||
f"from metadata fallback (stream_profile:{stream_id} was missing)"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"Channel {self.uuid}: no profile found for "
|
||||
f"stream_profile:{stream_id} or in metadata fallback"
|
||||
)
|
||||
return False
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found profile_id={profile_id} for "
|
||||
f"stream {stream_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# (e.g. from _clean_redis_keys or ChannelService.stop_channel)
|
||||
# won't find them via fallback and DECR again
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
def update_stream_profile(self, new_profile_id):
|
||||
"""
|
||||
Updates the profile for the current stream and adjusts connection counts.
|
||||
|
||||
Args:
|
||||
new_profile_id: The ID of the new stream profile to use
|
||||
|
||||
Returns:
|
||||
bool: True if successful, False otherwise
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
# Get current stream ID
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id_bytes:
|
||||
logger.debug("No active stream found for channel")
|
||||
return False
|
||||
|
||||
stream_id = int(stream_id_bytes)
|
||||
|
||||
# Get current profile ID
|
||||
current_profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not current_profile_id_bytes:
|
||||
logger.debug("No profile found for current stream")
|
||||
return False
|
||||
|
||||
current_profile_id = int(current_profile_id_bytes)
|
||||
|
||||
# Don't do anything if the profile is already set to the requested one
|
||||
if current_profile_id == new_profile_id:
|
||||
return True
|
||||
|
||||
# Use pipeline for atomic profile switch to prevent counter drift
|
||||
# if an exception occurs between DECR and INCR
|
||||
old_profile_connections_key = f"profile_connections:{current_profile_id}"
|
||||
new_profile_connections_key = f"profile_connections:{new_profile_id}"
|
||||
old_count = int(redis_client.get(old_profile_connections_key) or 0)
|
||||
|
||||
pipe = redis_client.pipeline()
|
||||
if old_count > 0:
|
||||
pipe.decr(old_profile_connections_key)
|
||||
pipe.set(f"stream_profile:{stream_id}", new_profile_id)
|
||||
pipe.incr(new_profile_connections_key)
|
||||
pipe.execute()
|
||||
logger.info(
|
||||
f"Updated stream {stream_id} profile from {current_profile_id} to {new_profile_id}"
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
class ChannelProfile(models.Model):
|
||||
name = models.CharField(max_length=100, unique=True)
|
||||
|
||||
|
||||
class ChannelProfileMembership(models.Model):
|
||||
channel_profile = models.ForeignKey(ChannelProfile, on_delete=models.CASCADE)
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
enabled = models.BooleanField(
|
||||
default=True
|
||||
) # Track if the channel is enabled for this group
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_profile", "channel")
|
||||
|
||||
|
||||
class ChannelStream(models.Model):
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
stream = models.ForeignKey(Stream, on_delete=models.CASCADE)
|
||||
order = models.PositiveIntegerField(default=0) # Ordering field
|
||||
|
||||
class Meta:
|
||||
ordering = ["order"] # Ensure streams are retrieved in order
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=["channel", "stream"], name="unique_channel_stream"
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class ChannelGroupM3UAccount(models.Model):
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup, on_delete=models.CASCADE, related_name="m3u_accounts"
|
||||
)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount, on_delete=models.CASCADE, related_name="channel_group"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
enabled = models.BooleanField(default=True)
|
||||
auto_channel_sync = models.BooleanField(
|
||||
default=False,
|
||||
help_text='Automatically create/delete channels to match streams in this group'
|
||||
)
|
||||
auto_sync_channel_start = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text='Starting channel number for auto-created channels in this group'
|
||||
)
|
||||
last_seen = models.DateTimeField(
|
||||
default=timezone.now,
|
||||
db_index=True,
|
||||
help_text='Last time this group was seen in the M3U source during a refresh'
|
||||
)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text='Whether this group relationship is stale (not seen in recent refresh, pending deletion)'
|
||||
)
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_group", "m3u_account")
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_group.name} - {self.m3u_account.name} (Enabled: {self.enabled})"
|
||||
|
||||
|
||||
class Logo(models.Model):
|
||||
name = models.CharField(max_length=255)
|
||||
url = models.TextField(unique=True)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
|
||||
class Recording(models.Model):
|
||||
channel = models.ForeignKey(
|
||||
"Channel", on_delete=models.CASCADE, related_name="recordings"
|
||||
)
|
||||
start_time = models.DateTimeField()
|
||||
end_time = models.DateTimeField()
|
||||
task_id = models.CharField(max_length=255, null=True, blank=True)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel.name} - {self.start_time} to {self.end_time}"
|
||||
|
||||
|
||||
class RecurringRecordingRule(models.Model):
|
||||
"""Rule describing a recurring manual DVR schedule."""
|
||||
|
||||
channel = models.ForeignKey(
|
||||
"Channel",
|
||||
on_delete=models.CASCADE,
|
||||
related_name="recurring_rules",
|
||||
)
|
||||
days_of_week = models.JSONField(default=list)
|
||||
start_time = models.TimeField()
|
||||
end_time = models.TimeField()
|
||||
enabled = models.BooleanField(default=True)
|
||||
name = models.CharField(max_length=255, blank=True)
|
||||
start_date = models.DateField(null=True, blank=True)
|
||||
end_date = models.DateField(null=True, blank=True)
|
||||
created_at = models.DateTimeField(auto_now_add=True)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
|
||||
class Meta:
|
||||
ordering = ["channel", "start_time"]
|
||||
|
||||
def __str__(self):
|
||||
channel_name = getattr(self.channel, "name", str(self.channel_id))
|
||||
return f"Recurring rule for {channel_name}"
|
||||
|
||||
def cleaned_days(self):
|
||||
try:
|
||||
return sorted({int(d) for d in (self.days_of_week or []) if 0 <= int(d) <= 6})
|
||||
except Exception:
|
||||
return []
|
||||
@@ -0,0 +1,958 @@
|
||||
from django.db import models
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.conf import settings
|
||||
from core.models import StreamProfile, CoreSettings
|
||||
from core.utils import RedisClient
|
||||
from apps.proxy.ts_proxy.redis_keys import RedisKeys
|
||||
from apps.proxy.ts_proxy.constants import ChannelMetadataField
|
||||
import logging
|
||||
import uuid
|
||||
from django.utils import timezone
|
||||
import hashlib
|
||||
import json
|
||||
from apps.epg.models import EPGData
|
||||
from apps.accounts.models import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# If you have an M3UAccount model in apps.m3u, you can still import it:
|
||||
from apps.m3u.models import M3UAccount
|
||||
|
||||
|
||||
# Add fallback functions if Redis isn't available
|
||||
def get_total_viewers(channel_id):
|
||||
"""Get viewer count from Redis or return 0 if Redis isn't available"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
try:
|
||||
return int(redis_client.get(f"channel:{channel_id}:viewers") or 0)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
class ChannelGroup(models.Model):
|
||||
name = models.TextField(unique=True, db_index=True)
|
||||
|
||||
def related_channels(self):
|
||||
# local import if needed to avoid cyc. Usually fine in a single file though
|
||||
return Channel.objects.filter(channel_group=self)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def bulk_create_and_fetch(cls, objects):
|
||||
# Perform the bulk create operation
|
||||
cls.objects.bulk_create(objects)
|
||||
|
||||
# Use a unique field to fetch the created objects (assuming 'name' is unique)
|
||||
created_objects = cls.objects.filter(name__in=[obj.name for obj in objects])
|
||||
|
||||
return created_objects
|
||||
|
||||
|
||||
class Stream(models.Model):
|
||||
"""
|
||||
Represents a single stream (e.g. from an M3U source or custom URL).
|
||||
"""
|
||||
|
||||
name = models.CharField(max_length=255, default="Default Stream")
|
||||
url = models.URLField(max_length=4096, blank=True, null=True)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount,
|
||||
on_delete=models.CASCADE,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
logo_url = models.TextField(blank=True, null=True)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
local_file = models.FileField(upload_to="uploads/", blank=True, null=True)
|
||||
current_viewers = models.PositiveIntegerField(default=0)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
null=True,
|
||||
blank=True,
|
||||
on_delete=models.SET_NULL,
|
||||
related_name="streams",
|
||||
)
|
||||
is_custom = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this is a user-created stream or from an M3U account",
|
||||
)
|
||||
stream_hash = models.CharField(
|
||||
max_length=255,
|
||||
null=True,
|
||||
unique=True,
|
||||
help_text="Unique hash for this stream from the M3U account",
|
||||
db_index=True,
|
||||
)
|
||||
last_seen = models.DateTimeField(db_index=True, default=timezone.now)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream is stale (not seen in recent refresh, pending deletion)"
|
||||
)
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream contains adult content"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
stream_id = models.IntegerField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider stream ID (e.g., XC stream_id) for stable identity across credential changes"
|
||||
)
|
||||
stream_chno = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1"
|
||||
)
|
||||
|
||||
# Stream statistics fields
|
||||
stream_stats = models.JSONField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="JSON object containing stream statistics like video codec, resolution, etc."
|
||||
)
|
||||
stream_stats_updated_at = models.DateTimeField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="When stream statistics were last updated",
|
||||
db_index=True
|
||||
)
|
||||
|
||||
class Meta:
|
||||
# If you use m3u_account, you might do unique_together = ('name','url','m3u_account')
|
||||
verbose_name = "Stream"
|
||||
verbose_name_plural = "Streams"
|
||||
ordering = ["-updated_at"]
|
||||
|
||||
def __str__(self):
|
||||
return self.name or self.url or f"Stream ID {self.id}"
|
||||
|
||||
@classmethod
|
||||
def generate_hash_key(cls, name, url, tvg_id, keys=None, m3u_id=None, group=None,
|
||||
account_type=None, stream_id=None):
|
||||
if keys is None:
|
||||
keys = CoreSettings.get_m3u_hash_key().split(",")
|
||||
|
||||
# For XC accounts, use stream_id instead of url when 'url' is in the hash keys
|
||||
# This ensures credential/URL changes don't break stream identity
|
||||
effective_url = url
|
||||
use_stream_id = account_type == 'XC' and stream_id and 'url' in keys
|
||||
if use_stream_id:
|
||||
effective_url = stream_id
|
||||
|
||||
stream_parts = {"name": name, "url": effective_url, "tvg_id": tvg_id, "m3u_id": m3u_id, "group": group}
|
||||
|
||||
hash_parts = {key: stream_parts[key] for key in keys if key in stream_parts}
|
||||
|
||||
# When using stream_id instead of URL, we MUST include m3u_id to prevent
|
||||
# collisions across different XC accounts (stream_id is only unique per account)
|
||||
if use_stream_id and 'm3u_id' not in hash_parts:
|
||||
hash_parts['m3u_id'] = m3u_id
|
||||
|
||||
# Serialize and hash the dictionary
|
||||
serialized_obj = json.dumps(
|
||||
hash_parts, sort_keys=True
|
||||
) # sort_keys ensures consistent ordering
|
||||
hash_object = hashlib.sha256(serialized_obj.encode())
|
||||
return hash_object.hexdigest()
|
||||
|
||||
@classmethod
|
||||
def update_or_create_by_hash(cls, hash_value, **fields_to_update):
|
||||
try:
|
||||
# Try to find the Stream object with the given hash
|
||||
stream = cls.objects.get(stream_hash=hash_value)
|
||||
# If it exists, update the fields
|
||||
for field, value in fields_to_update.items():
|
||||
setattr(stream, field, value)
|
||||
stream.save() # Save the updated object
|
||||
return stream, False # False means it was updated, not created
|
||||
except cls.DoesNotExist:
|
||||
# If it doesn't exist, create a new object with the given hash
|
||||
fields_to_update["stream_hash"] = (
|
||||
hash_value # Make sure the hash field is set
|
||||
)
|
||||
stream = cls.objects.create(**fields_to_update)
|
||||
return stream, True # True means it was created
|
||||
|
||||
def get_stream_profile(self):
|
||||
"""
|
||||
Get the stream profile for this stream.
|
||||
Uses the stream's own profile if set, otherwise returns the default.
|
||||
"""
|
||||
if self.stream_profile:
|
||||
return self.stream_profile
|
||||
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available profile for this stream and reserves a connection slot.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
profile_id = redis_client.get(f"stream_profile:{self.id}")
|
||||
if profile_id:
|
||||
profile_id = int(profile_id)
|
||||
return self.id, profile_id, None
|
||||
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = self.m3u_account
|
||||
m3u_profiles = m3u_account.profiles.all()
|
||||
default_profile = next((obj for obj in m3u_profiles if obj.is_default), None)
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
logger.info(profile)
|
||||
# Skip inactive profiles
|
||||
if profile.is_active == False:
|
||||
continue
|
||||
|
||||
# Atomic slot reservation: INCR first, check, rollback if over capacity
|
||||
if profile.max_streams == 0:
|
||||
reserved = True
|
||||
else:
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
if new_count <= profile.max_streams:
|
||||
reserved = True
|
||||
else:
|
||||
redis_client.decr(profile_connections_key)
|
||||
reserved = False
|
||||
|
||||
if reserved:
|
||||
redis_client.set(f"channel_stream:{self.id}", self.id)
|
||||
redis_client.set(f"stream_profile:{self.id}", profile.id)
|
||||
return self.id, profile.id, None
|
||||
|
||||
return None, None, "All active M3U profiles have reached maximum connection limits"
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = self.id
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not profile_id:
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: no profile found in "
|
||||
f"stream_profile:{stream_id}"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
|
||||
profile_id = int(profile_id)
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: found profile_id={profile_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
class ChannelManager(models.Manager):
|
||||
def active(self):
|
||||
return self.all()
|
||||
|
||||
|
||||
class Channel(models.Model):
|
||||
channel_number = models.FloatField(db_index=True)
|
||||
name = models.CharField(max_length=255)
|
||||
logo = models.ForeignKey(
|
||||
"Logo",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
# M2M to Stream now in the same file
|
||||
streams = models.ManyToManyField(
|
||||
Stream, blank=True, through="ChannelStream", related_name="channels"
|
||||
)
|
||||
|
||||
channel_group = models.ForeignKey(
|
||||
"ChannelGroup",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
help_text="Channel group this channel belongs to.",
|
||||
)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
tvc_guide_stationid = models.CharField(max_length=255, blank=True, null=True)
|
||||
|
||||
epg_data = models.ForeignKey(
|
||||
EPGData,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
uuid = models.UUIDField(
|
||||
default=uuid.uuid4, editable=False, unique=True, db_index=True
|
||||
)
|
||||
|
||||
user_level = models.IntegerField(default=0)
|
||||
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this channel contains adult content"
|
||||
)
|
||||
|
||||
auto_created = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this channel was automatically created via M3U auto channel sync"
|
||||
)
|
||||
auto_created_by = models.ForeignKey(
|
||||
"m3u.M3UAccount",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="auto_created_channels",
|
||||
help_text="The M3U account that auto-created this channel"
|
||||
)
|
||||
|
||||
created_at = models.DateTimeField(
|
||||
auto_now_add=True,
|
||||
help_text="Timestamp when this channel was created"
|
||||
)
|
||||
updated_at = models.DateTimeField(
|
||||
auto_now=True,
|
||||
help_text="Timestamp when this channel was last updated"
|
||||
)
|
||||
|
||||
def clean(self):
|
||||
# Enforce unique channel_number within a given group
|
||||
existing = Channel.objects.filter(
|
||||
channel_number=self.channel_number, channel_group=self.channel_group
|
||||
).exclude(id=self.id)
|
||||
if existing.exists():
|
||||
raise ValidationError(
|
||||
f"Channel number {self.channel_number} already exists in group {self.channel_group}."
|
||||
)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_number} - {self.name}"
|
||||
|
||||
@classmethod
|
||||
def get_next_available_channel_number(cls, starting_from=1):
|
||||
used_numbers = set(cls.objects.all().values_list("channel_number", flat=True))
|
||||
n = starting_from
|
||||
while n in used_numbers:
|
||||
n += 1
|
||||
return n
|
||||
|
||||
# @TODO: honor stream's stream profile
|
||||
def get_stream_profile(self):
|
||||
stream_profile = self.stream_profile
|
||||
if not stream_profile:
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def _pick_channel_to_preempt(
|
||||
self,
|
||||
profile_id,
|
||||
requester_level,
|
||||
redis_client,
|
||||
exclude_channel_ids=None,
|
||||
cooldown_seconds=30,
|
||||
):
|
||||
"""
|
||||
Pick the lowest-impact channel to terminate on the given profile.
|
||||
Returns: Optional[int] channel_id to preempt
|
||||
"""
|
||||
exclude_channel_ids = set(exclude_channel_ids or [])
|
||||
candidates = []
|
||||
|
||||
# 1) Try to get active channel IDs for this profile from an index set if available
|
||||
ch_set_key = f"ts_proxy:profile:{profile_id}:channels"
|
||||
try:
|
||||
ch_ids = { (int(x) if not isinstance(x, int) else x) for x in (redis_client.smembers(ch_set_key) or set()) }
|
||||
except Exception:
|
||||
ch_ids = set()
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
# 2) Fallback: scan metadata keys and filter by m3u_profile == profile_id
|
||||
if not ch_ids:
|
||||
cursor = 0
|
||||
pattern = "ts_proxy:channel:*:metadata"
|
||||
while True:
|
||||
cursor, keys = redis_client.scan(cursor=cursor, match=pattern, count=500)
|
||||
if keys:
|
||||
# Prefer HGET m3u_profile if metadata is a hash
|
||||
pipe = redis_client.pipeline()
|
||||
for k in keys:
|
||||
pipe.hget(k, "m3u_profile")
|
||||
prof_vals = pipe.execute()
|
||||
for k, prof_val in zip(keys, prof_vals):
|
||||
try:
|
||||
pid = int(prof_val) if prof_val is not None else None
|
||||
except Exception:
|
||||
pid = None
|
||||
|
||||
if pid == profile_id:
|
||||
parts = k.split(":") # ts_proxy:channel:{id}:metadata
|
||||
if len(parts) >= 4:
|
||||
try:
|
||||
ch_ids.add(int(parts[2]))
|
||||
except Exception:
|
||||
pass
|
||||
if cursor == 0:
|
||||
break
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
if not ch_ids:
|
||||
return None
|
||||
|
||||
# 3) Score candidates
|
||||
for ch_id in ch_ids:
|
||||
if ch_id in exclude_channel_ids:
|
||||
continue
|
||||
|
||||
# Skip if recently preempted
|
||||
last_preempt_key = f"ts_proxy:channel:{ch_id}:last_preempt"
|
||||
try:
|
||||
last_preempt = float(redis_client.get(last_preempt_key) or 0.0)
|
||||
except Exception:
|
||||
last_preempt = 0.0
|
||||
if last_preempt and (time.time() - last_preempt) < cooldown_seconds:
|
||||
continue
|
||||
|
||||
# Clients and their levels
|
||||
clients_key = f"ts_proxy:channel:{ch_id}:clients"
|
||||
member_ids = list(redis_client.smembers(clients_key) or [])
|
||||
viewer_count = len(member_ids)
|
||||
max_viewer_level = 0
|
||||
if viewer_count:
|
||||
pipe = redis_client.pipeline()
|
||||
for cid in member_ids:
|
||||
pipe.hget(f"ts_proxy:channel:{ch_id}:clients:{cid}", "user_level")
|
||||
levels_raw = pipe.execute()
|
||||
levels = []
|
||||
for lv in levels_raw:
|
||||
try:
|
||||
levels.append(int(lv or 0))
|
||||
except Exception:
|
||||
levels.append(0)
|
||||
max_viewer_level = max(levels or [0])
|
||||
|
||||
# Only preempt if requester strictly outranks this channel's viewers
|
||||
if requester_level <= max_viewer_level:
|
||||
continue
|
||||
|
||||
# Metadata (protected/recording/started_at_ts)
|
||||
meta_key = f"ts_proxy:channel:{ch_id}:metadata"
|
||||
try:
|
||||
protected, recording, started_at_ts = redis_client.hmget(
|
||||
meta_key, "protected", "recording", "started_at_ts"
|
||||
)
|
||||
except Exception:
|
||||
protected = recording = started_at_ts = None
|
||||
|
||||
protected = str(protected or "0") in ("1", "true", "True")
|
||||
recording = str(recording or "0") in ("1", "true", "True")
|
||||
if protected or recording:
|
||||
continue
|
||||
|
||||
try:
|
||||
started_at_ts = float(started_at_ts) if started_at_ts is not None else None
|
||||
except Exception:
|
||||
started_at_ts = None
|
||||
if started_at_ts is None:
|
||||
started_at_ts = time.time() # treat unknown as newest
|
||||
|
||||
# Score: lower is safer to terminate
|
||||
has_viewers = 1 if viewer_count > 0 else 0
|
||||
score = (has_viewers, max_viewer_level, viewer_count, started_at_ts)
|
||||
candidates.append((score, ch_id))
|
||||
|
||||
logger.debug("Candidate channels after scoring:")
|
||||
logger.debug(candidates)
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
victim_id = candidates[0][1]
|
||||
|
||||
# Mark preempt timestamp to avoid thrashing
|
||||
try:
|
||||
redis_client.set(f"ts_proxy:channel:{victim_id}:last_preempt", str(time.time()), ex=3600)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return victim_id
|
||||
|
||||
def _check_and_reserve_profile_slot(self, profile, redis_client):
|
||||
"""
|
||||
Atomically check and reserve a connection slot for the given profile.
|
||||
|
||||
Uses an INCR-first-then-check pattern to eliminate the TOCTOU race
|
||||
condition where separate GET + check + INCR operations could allow
|
||||
concurrent requests to both pass the capacity check.
|
||||
|
||||
For profiles with max_streams=0 (unlimited), no reservation is needed.
|
||||
|
||||
Args:
|
||||
profile: M3UAccountProfile instance
|
||||
redis_client: Redis client instance
|
||||
|
||||
Returns:
|
||||
tuple: (reserved: bool, current_count: int)
|
||||
"""
|
||||
if profile.max_streams == 0:
|
||||
return (True, 0)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
|
||||
# Atomically increment first — this is a single Redis command
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
|
||||
if new_count <= profile.max_streams:
|
||||
return (True, new_count)
|
||||
|
||||
# Over capacity — roll back the increment
|
||||
redis_client.decr(profile_connections_key)
|
||||
return (False, new_count - 1)
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available stream for the requested channel and returns the selected stream and profile.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
error_reason = None
|
||||
|
||||
# Check if this channel has any streams
|
||||
if not self.streams.exists():
|
||||
error_reason = "No streams assigned to channel"
|
||||
return None, None, error_reason
|
||||
|
||||
# Check if a stream is already active for this channel
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if stream_id_bytes:
|
||||
try:
|
||||
stream_id = int(stream_id_bytes)
|
||||
profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id_bytes:
|
||||
try:
|
||||
profile_id = int(profile_id_bytes)
|
||||
return stream_id, profile_id, None
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid profile ID retrieved from Redis: {profile_id_bytes}"
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid stream ID retrieved from Redis: {stream_id_bytes}"
|
||||
)
|
||||
|
||||
# No existing active stream, attempt to assign a new one
|
||||
has_streams_but_maxed_out = False
|
||||
has_active_profiles = False
|
||||
|
||||
# Iterate through channel streams and their profiles
|
||||
for stream in self.streams.all().order_by("channelstream__order"):
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = stream.m3u_account
|
||||
if not m3u_account:
|
||||
logger.debug(f"Stream {stream.id} has no M3U account")
|
||||
continue
|
||||
if m3u_account.is_active == False:
|
||||
logger.debug(f"M3U account {m3u_account.id} is inactive, skipping.")
|
||||
continue
|
||||
|
||||
m3u_profiles = m3u_account.profiles.filter(is_active=True)
|
||||
default_profile = next(
|
||||
(obj for obj in m3u_profiles if obj.is_default), None
|
||||
)
|
||||
|
||||
if not default_profile:
|
||||
logger.debug(f"M3U account {m3u_account.id} has no active default profile")
|
||||
continue
|
||||
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
has_active_profiles = True
|
||||
|
||||
# Atomically check and reserve a slot (INCR-first pattern)
|
||||
reserved, current_count = self._check_and_reserve_profile_slot(
|
||||
profile, redis_client
|
||||
)
|
||||
|
||||
if reserved:
|
||||
# Slot reserved — assign stream to this channel
|
||||
redis_client.set(f"channel_stream:{self.id}", stream.id)
|
||||
redis_client.set(f"stream_profile:{stream.id}", profile.id)
|
||||
|
||||
return (
|
||||
stream.id,
|
||||
profile.id,
|
||||
None,
|
||||
) # Return newly assigned stream and matched profile
|
||||
else:
|
||||
# At capacity: try to preempt a lower-impact channel on this profile
|
||||
victim_channel_id = self._pick_channel_to_preempt(
|
||||
profile_id=profile.id,
|
||||
requester_level=requester.user_level if requester else 100,
|
||||
redis_client=redis_client,
|
||||
exclude_channel_ids=None,
|
||||
)
|
||||
if victim_channel_id:
|
||||
logger.info(f"Preempting channel {victim_channel_id} for new stream on profile {profile.id}")
|
||||
# return self.id, profile.id, victim_channel_id
|
||||
|
||||
# This profile is at max connections
|
||||
has_streams_but_maxed_out = True
|
||||
logger.debug(
|
||||
f"Profile {profile.id} at max connections: "
|
||||
f"{current_count}/{profile.max_streams}"
|
||||
)
|
||||
|
||||
# No available streams - determine specific reason
|
||||
if has_streams_but_maxed_out:
|
||||
error_reason = "All active M3U profiles have reached maximum connection limits"
|
||||
elif has_active_profiles:
|
||||
error_reason = "No compatible active profile found for any assigned stream"
|
||||
else:
|
||||
error_reason = "No active profiles found for any assigned stream"
|
||||
|
||||
return None, None, error_reason
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no stream/profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id:
|
||||
# Primary key missing — try metadata hash fallback.
|
||||
# The proxy may have already cleaned up channel_stream/stream_profile
|
||||
# keys, but the metadata hash can still have the stream_id and profile.
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_stream_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.STREAM_ID
|
||||
)
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
|
||||
if meta_stream_id and meta_profile_id:
|
||||
stream_id = int(meta_stream_id)
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered stream_id={stream_id}, "
|
||||
f"profile_id={profile_id} from metadata fallback"
|
||||
)
|
||||
# Clean up any remaining keys
|
||||
redis_client.delete(f"channel_stream:{self.id}")
|
||||
redis_client.delete(f"stream_profile:{stream_id}")
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# won't find them and DECR again
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
current_count = int(
|
||||
redis_client.get(profile_connections_key) or 0
|
||||
)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
return True
|
||||
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: no stream info found in primary keys "
|
||||
f"or metadata fallback"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"channel_stream:{self.id}") # Remove active stream
|
||||
|
||||
stream_id = int(stream_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found stream_id={stream_id} for "
|
||||
f"channel_stream:{self.id}"
|
||||
)
|
||||
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id:
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
profile_id = int(profile_id)
|
||||
else:
|
||||
# stream_profile key missing — try metadata hash fallback
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
if meta_profile_id:
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered profile_id={profile_id} "
|
||||
f"from metadata fallback (stream_profile:{stream_id} was missing)"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"Channel {self.uuid}: no profile found for "
|
||||
f"stream_profile:{stream_id} or in metadata fallback"
|
||||
)
|
||||
return False
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found profile_id={profile_id} for "
|
||||
f"stream {stream_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# (e.g. from _clean_redis_keys or ChannelService.stop_channel)
|
||||
# won't find them via fallback and DECR again
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
def update_stream_profile(self, new_profile_id):
|
||||
"""
|
||||
Updates the profile for the current stream and adjusts connection counts.
|
||||
|
||||
Args:
|
||||
new_profile_id: The ID of the new stream profile to use
|
||||
|
||||
Returns:
|
||||
bool: True if successful, False otherwise
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
# Get current stream ID
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id_bytes:
|
||||
logger.debug("No active stream found for channel")
|
||||
return False
|
||||
|
||||
stream_id = int(stream_id_bytes)
|
||||
|
||||
# Get current profile ID
|
||||
current_profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not current_profile_id_bytes:
|
||||
logger.debug("No profile found for current stream")
|
||||
return False
|
||||
|
||||
current_profile_id = int(current_profile_id_bytes)
|
||||
|
||||
# Don't do anything if the profile is already set to the requested one
|
||||
if current_profile_id == new_profile_id:
|
||||
return True
|
||||
|
||||
# Use pipeline for atomic profile switch to prevent counter drift
|
||||
# if an exception occurs between DECR and INCR
|
||||
old_profile_connections_key = f"profile_connections:{current_profile_id}"
|
||||
new_profile_connections_key = f"profile_connections:{new_profile_id}"
|
||||
old_count = int(redis_client.get(old_profile_connections_key) or 0)
|
||||
|
||||
pipe = redis_client.pipeline()
|
||||
if old_count > 0:
|
||||
pipe.decr(old_profile_connections_key)
|
||||
pipe.set(f"stream_profile:{stream_id}", new_profile_id)
|
||||
pipe.incr(new_profile_connections_key)
|
||||
pipe.execute()
|
||||
logger.info(
|
||||
f"Updated stream {stream_id} profile from {current_profile_id} to {new_profile_id}"
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
class ChannelProfile(models.Model):
|
||||
name = models.CharField(max_length=100, unique=True)
|
||||
|
||||
|
||||
class ChannelProfileMembership(models.Model):
|
||||
channel_profile = models.ForeignKey(ChannelProfile, on_delete=models.CASCADE)
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
enabled = models.BooleanField(
|
||||
default=True
|
||||
) # Track if the channel is enabled for this group
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_profile", "channel")
|
||||
|
||||
|
||||
class ChannelStream(models.Model):
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
stream = models.ForeignKey(Stream, on_delete=models.CASCADE)
|
||||
order = models.PositiveIntegerField(default=0) # Ordering field
|
||||
|
||||
class Meta:
|
||||
ordering = ["order"] # Ensure streams are retrieved in order
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=["channel", "stream"], name="unique_channel_stream"
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class ChannelGroupM3UAccount(models.Model):
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup, on_delete=models.CASCADE, related_name="m3u_accounts"
|
||||
)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount, on_delete=models.CASCADE, related_name="channel_group"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
enabled = models.BooleanField(default=True)
|
||||
auto_channel_sync = models.BooleanField(
|
||||
default=False,
|
||||
help_text='Automatically create/delete channels to match streams in this group'
|
||||
)
|
||||
auto_sync_channel_start = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text='Starting channel number for auto-created channels in this group'
|
||||
)
|
||||
last_seen = models.DateTimeField(
|
||||
default=timezone.now,
|
||||
db_index=True,
|
||||
help_text='Last time this group was seen in the M3U source during a refresh'
|
||||
)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text='Whether this group relationship is stale (not seen in recent refresh, pending deletion)'
|
||||
)
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_group", "m3u_account")
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_group.name} - {self.m3u_account.name} (Enabled: {self.enabled})"
|
||||
|
||||
|
||||
class Logo(models.Model):
|
||||
name = models.CharField(max_length=255)
|
||||
url = models.TextField(unique=True)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
|
||||
class Recording(models.Model):
|
||||
channel = models.ForeignKey(
|
||||
"Channel", on_delete=models.CASCADE, related_name="recordings"
|
||||
)
|
||||
start_time = models.DateTimeField()
|
||||
end_time = models.DateTimeField()
|
||||
task_id = models.CharField(max_length=255, null=True, blank=True)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel.name} - {self.start_time} to {self.end_time}"
|
||||
|
||||
|
||||
class RecurringRecordingRule(models.Model):
|
||||
"""Rule describing a recurring manual DVR schedule."""
|
||||
|
||||
channel = models.ForeignKey(
|
||||
"Channel",
|
||||
on_delete=models.CASCADE,
|
||||
related_name="recurring_rules",
|
||||
)
|
||||
days_of_week = models.JSONField(default=list)
|
||||
start_time = models.TimeField()
|
||||
end_time = models.TimeField()
|
||||
enabled = models.BooleanField(default=True)
|
||||
name = models.CharField(max_length=255, blank=True)
|
||||
start_date = models.DateField(null=True, blank=True)
|
||||
end_date = models.DateField(null=True, blank=True)
|
||||
created_at = models.DateTimeField(auto_now_add=True)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
|
||||
class Meta:
|
||||
ordering = ["channel", "start_time"]
|
||||
|
||||
def __str__(self):
|
||||
channel_name = getattr(self.channel, "name", str(self.channel_id))
|
||||
return f"Recurring rule for {channel_name}"
|
||||
|
||||
def cleaned_days(self):
|
||||
try:
|
||||
return sorted({int(d) for d in (self.days_of_week or []) if 0 <= int(d) <= 6})
|
||||
except Exception:
|
||||
return []
|
||||
@@ -0,0 +1,537 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from rest_framework import serializers
|
||||
from .models import (
|
||||
Stream,
|
||||
Channel,
|
||||
ChannelGroup,
|
||||
ChannelStream,
|
||||
ChannelGroupM3UAccount,
|
||||
Logo,
|
||||
ChannelProfile,
|
||||
ChannelProfileMembership,
|
||||
Recording,
|
||||
RecurringRecordingRule,
|
||||
)
|
||||
from apps.epg.serializers import EPGDataSerializer
|
||||
from core.models import StreamProfile
|
||||
from apps.epg.models import EPGData
|
||||
from django.urls import reverse
|
||||
from rest_framework import serializers
|
||||
from django.utils import timezone
|
||||
from core.utils import validate_flexible_url
|
||||
|
||||
|
||||
class LogoSerializer(serializers.ModelSerializer):
|
||||
cache_url = serializers.SerializerMethodField()
|
||||
channel_count = serializers.SerializerMethodField()
|
||||
is_used = serializers.SerializerMethodField()
|
||||
channel_names = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Logo
|
||||
fields = ["id", "name", "url", "cache_url", "channel_count", "is_used", "channel_names"]
|
||||
|
||||
def validate_url(self, value):
|
||||
"""Validate that the URL is unique for creation or update"""
|
||||
if self.instance and self.instance.url == value:
|
||||
return value
|
||||
|
||||
if Logo.objects.filter(url=value).exists():
|
||||
raise serializers.ValidationError("A logo with this URL already exists.")
|
||||
|
||||
return value
|
||||
|
||||
def create(self, validated_data):
|
||||
"""Handle logo creation with proper URL validation"""
|
||||
return Logo.objects.create(**validated_data)
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
"""Handle logo updates"""
|
||||
for attr, value in validated_data.items():
|
||||
setattr(instance, attr, value)
|
||||
instance.save()
|
||||
return instance
|
||||
|
||||
def get_cache_url(self, obj):
|
||||
# return f"/api/channels/logos/{obj.id}/cache/"
|
||||
request = self.context.get("request")
|
||||
if request:
|
||||
return request.build_absolute_uri(
|
||||
reverse("api:channels:logo-cache", args=[obj.id])
|
||||
)
|
||||
return reverse("api:channels:logo-cache", args=[obj.id])
|
||||
|
||||
def get_channel_count(self, obj):
|
||||
"""Get the number of channels using this logo"""
|
||||
return obj.channels.count()
|
||||
|
||||
def get_is_used(self, obj):
|
||||
"""Check if this logo is used by any channels"""
|
||||
return obj.channels.exists()
|
||||
|
||||
def get_channel_names(self, obj):
|
||||
"""Get the names of channels using this logo (limited to first 5)"""
|
||||
names = []
|
||||
|
||||
# Get channel names
|
||||
channels = obj.channels.all()[:5]
|
||||
for channel in channels:
|
||||
names.append(f"Channel: {channel.name}")
|
||||
|
||||
# Calculate total count for "more" message
|
||||
total_count = self.get_channel_count(obj)
|
||||
if total_count > 5:
|
||||
names.append(f"...and {total_count - 5} more")
|
||||
|
||||
return names
|
||||
|
||||
|
||||
#
|
||||
# Stream
|
||||
#
|
||||
class StreamSerializer(serializers.ModelSerializer):
|
||||
url = serializers.CharField(
|
||||
required=False,
|
||||
allow_blank=True,
|
||||
allow_null=True,
|
||||
validators=[validate_flexible_url]
|
||||
)
|
||||
stream_profile_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=StreamProfile.objects.all(),
|
||||
source="stream_profile",
|
||||
allow_null=True,
|
||||
required=False,
|
||||
)
|
||||
read_only_fields = ["is_custom", "m3u_account", "stream_hash", "stream_id", "stream_chno"]
|
||||
|
||||
class Meta:
|
||||
model = Stream
|
||||
fields = [
|
||||
"id",
|
||||
"name",
|
||||
"url",
|
||||
"m3u_account", # Uncomment if using M3U fields
|
||||
"logo_url",
|
||||
"tvg_id",
|
||||
"local_file",
|
||||
"current_viewers",
|
||||
"updated_at",
|
||||
"last_seen",
|
||||
"is_stale",
|
||||
"is_adult",
|
||||
"stream_profile_id",
|
||||
"is_custom",
|
||||
"channel_group",
|
||||
"stream_hash",
|
||||
"stream_stats",
|
||||
"stream_stats_updated_at",
|
||||
"stream_id",
|
||||
"stream_chno",
|
||||
]
|
||||
|
||||
def get_fields(self):
|
||||
fields = super().get_fields()
|
||||
|
||||
# Unable to edit specific properties if this stream was created from an M3U account
|
||||
if (
|
||||
self.instance
|
||||
and getattr(self.instance, "m3u_account", None)
|
||||
and not self.instance.is_custom
|
||||
):
|
||||
fields["id"].read_only = True
|
||||
fields["name"].read_only = True
|
||||
fields["url"].read_only = True
|
||||
fields["m3u_account"].read_only = True
|
||||
fields["tvg_id"].read_only = True
|
||||
fields["channel_group"].read_only = True
|
||||
|
||||
return fields
|
||||
|
||||
|
||||
class ChannelGroupM3UAccountSerializer(serializers.ModelSerializer):
|
||||
m3u_accounts = serializers.IntegerField(source="m3u_accounts.id", read_only=True)
|
||||
enabled = serializers.BooleanField()
|
||||
auto_channel_sync = serializers.BooleanField(default=False)
|
||||
auto_sync_channel_start = serializers.FloatField(allow_null=True, required=False)
|
||||
custom_properties = serializers.JSONField(required=False)
|
||||
|
||||
class Meta:
|
||||
model = ChannelGroupM3UAccount
|
||||
fields = ["m3u_accounts", "channel_group", "enabled", "auto_channel_sync", "auto_sync_channel_start", "custom_properties", "is_stale", "last_seen"]
|
||||
|
||||
def to_representation(self, instance):
|
||||
data = super().to_representation(instance)
|
||||
|
||||
custom_props = instance.custom_properties or {}
|
||||
|
||||
return data
|
||||
|
||||
def to_internal_value(self, data):
|
||||
# Accept both dict and JSON string for custom_properties (for backward compatibility)
|
||||
val = data.get("custom_properties")
|
||||
if isinstance(val, str):
|
||||
try:
|
||||
data["custom_properties"] = json.loads(val)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return super().to_internal_value(data)
|
||||
|
||||
#
|
||||
# Channel Group
|
||||
#
|
||||
class ChannelGroupSerializer(serializers.ModelSerializer):
|
||||
channel_count = serializers.SerializerMethodField()
|
||||
m3u_account_count = serializers.SerializerMethodField()
|
||||
m3u_accounts = ChannelGroupM3UAccountSerializer(
|
||||
many=True,
|
||||
read_only=True
|
||||
)
|
||||
|
||||
class Meta:
|
||||
model = ChannelGroup
|
||||
fields = ["id", "name", "channel_count", "m3u_account_count", "m3u_accounts"]
|
||||
|
||||
def get_channel_count(self, obj):
|
||||
"""Get count of channels in this group"""
|
||||
return obj.channels.count()
|
||||
|
||||
def get_m3u_account_count(self, obj):
|
||||
"""Get count of M3U accounts associated with this group"""
|
||||
return obj.m3u_accounts.count()
|
||||
|
||||
|
||||
class ChannelProfileSerializer(serializers.ModelSerializer):
|
||||
channels = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = ChannelProfile
|
||||
fields = ["id", "name", "channels"]
|
||||
|
||||
def get_channels(self, obj):
|
||||
memberships = ChannelProfileMembership.objects.filter(
|
||||
channel_profile=obj, enabled=True
|
||||
)
|
||||
return [membership.channel.id for membership in memberships]
|
||||
|
||||
|
||||
class ChannelProfileMembershipSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = ChannelProfileMembership
|
||||
fields = ["channel", "enabled"]
|
||||
|
||||
|
||||
class ChanneProfilelMembershipUpdateSerializer(serializers.Serializer):
|
||||
channel_id = serializers.IntegerField() # Ensure channel_id is an integer
|
||||
enabled = serializers.BooleanField()
|
||||
|
||||
|
||||
class BulkChannelProfileMembershipSerializer(serializers.Serializer):
|
||||
channels = serializers.ListField(
|
||||
child=ChanneProfilelMembershipUpdateSerializer(), # Use the nested serializer
|
||||
allow_empty=False,
|
||||
)
|
||||
|
||||
def validate_channels(self, value):
|
||||
if not value:
|
||||
raise serializers.ValidationError("At least one channel must be provided.")
|
||||
return value
|
||||
|
||||
|
||||
#
|
||||
# Channel
|
||||
#
|
||||
class ChannelSerializer(serializers.ModelSerializer):
|
||||
# Show nested group data, or ID
|
||||
# Ensure channel_number is explicitly typed as FloatField and properly validated
|
||||
channel_number = serializers.FloatField(
|
||||
allow_null=True,
|
||||
required=False,
|
||||
error_messages={"invalid": "Channel number must be a valid decimal number."},
|
||||
)
|
||||
channel_group_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=ChannelGroup.objects.all(), source="channel_group", required=False
|
||||
)
|
||||
epg_data_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=EPGData.objects.all(),
|
||||
source="epg_data",
|
||||
required=False,
|
||||
allow_null=True,
|
||||
)
|
||||
|
||||
stream_profile_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=StreamProfile.objects.all(),
|
||||
source="stream_profile",
|
||||
allow_null=True,
|
||||
required=False,
|
||||
)
|
||||
|
||||
streams = serializers.PrimaryKeyRelatedField(
|
||||
queryset=Stream.objects.all(), many=True, required=False
|
||||
)
|
||||
|
||||
logo_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=Logo.objects.all(),
|
||||
source="logo",
|
||||
allow_null=True,
|
||||
required=False,
|
||||
)
|
||||
|
||||
auto_created_by_name = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Channel
|
||||
fields = [
|
||||
"id",
|
||||
"channel_number",
|
||||
"name",
|
||||
"channel_group_id",
|
||||
"tvg_id",
|
||||
"tvc_guide_stationid",
|
||||
"epg_data_id",
|
||||
"streams",
|
||||
"stream_profile_id",
|
||||
"uuid",
|
||||
"logo_id",
|
||||
"user_level",
|
||||
"is_adult",
|
||||
"auto_created",
|
||||
"auto_created_by",
|
||||
"auto_created_by_name",
|
||||
]
|
||||
|
||||
def to_representation(self, instance):
|
||||
include_streams = self.context.get("include_streams", False)
|
||||
|
||||
if include_streams:
|
||||
self.fields["streams"] = serializers.SerializerMethodField()
|
||||
return super().to_representation(instance)
|
||||
else:
|
||||
# Fix: For PATCH/PUT responses, ensure streams are ordered
|
||||
representation = super().to_representation(instance)
|
||||
if "streams" in representation:
|
||||
representation["streams"] = list(
|
||||
instance.streams.all()
|
||||
.order_by("channelstream__order")
|
||||
.values_list("id", flat=True)
|
||||
)
|
||||
return representation
|
||||
|
||||
def get_logo(self, obj):
|
||||
return LogoSerializer(obj.logo).data
|
||||
|
||||
def get_streams(self, obj):
|
||||
"""Retrieve ordered stream IDs for GET requests."""
|
||||
return StreamSerializer(
|
||||
obj.streams.all().order_by("channelstream__order"), many=True
|
||||
).data
|
||||
|
||||
def create(self, validated_data):
|
||||
streams = validated_data.pop("streams", [])
|
||||
channel_number = validated_data.pop(
|
||||
"channel_number", Channel.get_next_available_channel_number()
|
||||
)
|
||||
validated_data["channel_number"] = channel_number
|
||||
|
||||
# Auto-assign Default Group if no channel_group is specified
|
||||
if "channel_group" not in validated_data or validated_data.get("channel_group") is None:
|
||||
from apps.channels.models import ChannelGroup
|
||||
default_group, _ = ChannelGroup.objects.get_or_create(name="Default Group")
|
||||
validated_data["channel_group"] = default_group
|
||||
|
||||
channel = Channel.objects.create(**validated_data)
|
||||
|
||||
# Add streams in the specified order
|
||||
for index, stream in enumerate(streams):
|
||||
ChannelStream.objects.create(
|
||||
channel=channel, stream_id=stream.id, order=index
|
||||
)
|
||||
|
||||
return channel
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
streams = validated_data.pop("streams", None)
|
||||
|
||||
# Update standard fields
|
||||
for attr, value in validated_data.items():
|
||||
setattr(instance, attr, value)
|
||||
|
||||
instance.save()
|
||||
|
||||
if streams is not None:
|
||||
# Normalize stream IDs
|
||||
normalized_ids = [
|
||||
stream.id if hasattr(stream, "id") else stream for stream in streams
|
||||
]
|
||||
print(normalized_ids)
|
||||
|
||||
# Get current mapping of stream_id -> ChannelStream
|
||||
current_links = {
|
||||
cs.stream_id: cs for cs in instance.channelstream_set.all()
|
||||
}
|
||||
|
||||
# Track existing stream IDs
|
||||
existing_ids = set(current_links.keys())
|
||||
new_ids = set(normalized_ids)
|
||||
|
||||
# Delete any links not in the new list
|
||||
to_remove = existing_ids - new_ids
|
||||
if to_remove:
|
||||
instance.channelstream_set.filter(stream_id__in=to_remove).delete()
|
||||
|
||||
# Update or create with new order
|
||||
for order, stream_id in enumerate(normalized_ids):
|
||||
if stream_id in current_links:
|
||||
cs = current_links[stream_id]
|
||||
if cs.order != order:
|
||||
cs.order = order
|
||||
cs.save(update_fields=["order"])
|
||||
else:
|
||||
ChannelStream.objects.create(
|
||||
channel=instance, stream_id=stream_id, order=order
|
||||
)
|
||||
|
||||
return instance
|
||||
|
||||
def validate_channel_number(self, value):
|
||||
"""Ensure channel_number is properly processed as a float"""
|
||||
if value is None:
|
||||
return value
|
||||
|
||||
try:
|
||||
# Ensure it's processed as a float
|
||||
return float(value)
|
||||
except (ValueError, TypeError):
|
||||
raise serializers.ValidationError(
|
||||
"Channel number must be a valid decimal number."
|
||||
)
|
||||
|
||||
def validate_stream_profile(self, value):
|
||||
"""Handle special case where empty/0 values mean 'use default' (null)"""
|
||||
if value == "0" or value == 0 or value == "" or value is None:
|
||||
return None
|
||||
return value # PrimaryKeyRelatedField will handle the conversion to object
|
||||
|
||||
def get_auto_created_by_name(self, obj):
|
||||
"""Get the name of the M3U account that auto-created this channel."""
|
||||
if obj.auto_created_by:
|
||||
return obj.auto_created_by.name
|
||||
return None
|
||||
|
||||
|
||||
class RecordingSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = Recording
|
||||
fields = "__all__"
|
||||
read_only_fields = ["task_id"]
|
||||
|
||||
def validate(self, data):
|
||||
from core.models import CoreSettings
|
||||
start_time = data.get("start_time")
|
||||
end_time = data.get("end_time")
|
||||
|
||||
if start_time and timezone.is_naive(start_time):
|
||||
start_time = timezone.make_aware(start_time, timezone.get_current_timezone())
|
||||
data["start_time"] = start_time
|
||||
if end_time and timezone.is_naive(end_time):
|
||||
end_time = timezone.make_aware(end_time, timezone.get_current_timezone())
|
||||
data["end_time"] = end_time
|
||||
|
||||
# If this is an EPG-based recording (program provided), apply global pre/post offsets
|
||||
try:
|
||||
cp = data.get("custom_properties") or {}
|
||||
is_epg_based = isinstance(cp, dict) and isinstance(cp.get("program"), (dict,))
|
||||
except Exception:
|
||||
is_epg_based = False
|
||||
|
||||
if is_epg_based and start_time and end_time:
|
||||
try:
|
||||
pre_min = int(CoreSettings.get_dvr_pre_offset_minutes())
|
||||
except Exception:
|
||||
pre_min = 0
|
||||
try:
|
||||
post_min = int(CoreSettings.get_dvr_post_offset_minutes())
|
||||
except Exception:
|
||||
post_min = 0
|
||||
from datetime import timedelta
|
||||
try:
|
||||
if pre_min and pre_min > 0:
|
||||
start_time = start_time - timedelta(minutes=pre_min)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if post_min and post_min > 0:
|
||||
end_time = end_time + timedelta(minutes=post_min)
|
||||
except Exception:
|
||||
pass
|
||||
# write back adjusted times so scheduling uses them
|
||||
data["start_time"] = start_time
|
||||
data["end_time"] = end_time
|
||||
|
||||
now = timezone.now() # timezone-aware current time
|
||||
|
||||
if end_time < now:
|
||||
raise serializers.ValidationError("End time must be in the future.")
|
||||
|
||||
if start_time < now:
|
||||
# Optional: Adjust start_time if it's in the past but end_time is in the future
|
||||
data["start_time"] = now # or: timezone.now() + timedelta(seconds=1)
|
||||
if end_time <= data["start_time"]:
|
||||
raise serializers.ValidationError("End time must be after start time.")
|
||||
|
||||
return data
|
||||
|
||||
|
||||
class RecurringRecordingRuleSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = RecurringRecordingRule
|
||||
fields = "__all__"
|
||||
read_only_fields = ["created_at", "updated_at"]
|
||||
|
||||
def validate_days_of_week(self, value):
|
||||
if not value:
|
||||
raise serializers.ValidationError("Select at least one day of the week")
|
||||
cleaned = []
|
||||
for entry in value:
|
||||
try:
|
||||
iv = int(entry)
|
||||
except (TypeError, ValueError):
|
||||
raise serializers.ValidationError("Days of week must be integers 0-6")
|
||||
if iv < 0 or iv > 6:
|
||||
raise serializers.ValidationError("Days of week must be between 0 (Monday) and 6 (Sunday)")
|
||||
cleaned.append(iv)
|
||||
return sorted(set(cleaned))
|
||||
|
||||
def validate(self, attrs):
|
||||
start = attrs.get("start_time") or getattr(self.instance, "start_time", None)
|
||||
end = attrs.get("end_time") or getattr(self.instance, "end_time", None)
|
||||
start_date = attrs.get("start_date") if "start_date" in attrs else getattr(self.instance, "start_date", None)
|
||||
end_date = attrs.get("end_date") if "end_date" in attrs else getattr(self.instance, "end_date", None)
|
||||
if start_date is None:
|
||||
existing_start = getattr(self.instance, "start_date", None)
|
||||
if existing_start is None:
|
||||
raise serializers.ValidationError("Start date is required")
|
||||
if start_date and end_date and end_date < start_date:
|
||||
raise serializers.ValidationError("End date must be on or after start date")
|
||||
if end_date is None:
|
||||
existing_end = getattr(self.instance, "end_date", None)
|
||||
if existing_end is None:
|
||||
raise serializers.ValidationError("End date is required")
|
||||
if start and end and start_date and end_date:
|
||||
start_dt = datetime.combine(start_date, start)
|
||||
end_dt = datetime.combine(end_date, end)
|
||||
if end_dt <= start_dt:
|
||||
raise serializers.ValidationError("End datetime must be after start datetime")
|
||||
elif start and end and end == start:
|
||||
raise serializers.ValidationError("End time must be different from start time")
|
||||
# Normalize empty strings to None for dates
|
||||
if attrs.get("end_date") == "":
|
||||
attrs["end_date"] = None
|
||||
if attrs.get("start_date") == "":
|
||||
attrs["start_date"] = None
|
||||
return super().validate(attrs)
|
||||
|
||||
def create(self, validated_data):
|
||||
return super().create(validated_data)
|
||||
@@ -0,0 +1,236 @@
|
||||
# apps/channels/signals.py
|
||||
|
||||
from django.db.models.signals import m2m_changed, pre_save, post_save, post_delete
|
||||
from django.dispatch import receiver
|
||||
from django.utils.timezone import now, is_aware, make_aware
|
||||
from celery.result import AsyncResult
|
||||
from django_celery_beat.models import ClockedSchedule, PeriodicTask
|
||||
from .models import Channel, Stream, ChannelProfile, ChannelProfileMembership, Recording
|
||||
from apps.m3u.models import M3UAccount
|
||||
from apps.epg.tasks import parse_programs_for_tvg_id
|
||||
import json
|
||||
import logging
|
||||
from .tasks import run_recording, prefetch_recording_artwork
|
||||
from datetime import timedelta
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@receiver(m2m_changed, sender=Channel.streams.through)
|
||||
def update_channel_tvg_id_and_logo(sender, instance, action, reverse, model, pk_set, **kwargs):
|
||||
"""
|
||||
Whenever streams are added to a channel:
|
||||
1) If the channel doesn't have a tvg_id, fill it from the first newly-added stream that has one.
|
||||
"""
|
||||
# We only care about post_add, i.e. once the new streams are fully associated
|
||||
if action == "post_add":
|
||||
# --- 1) Populate channel.tvg_id if empty ---
|
||||
if not instance.tvg_id:
|
||||
# Look for newly added streams that have a nonempty tvg_id
|
||||
streams_with_tvg = model.objects.filter(pk__in=pk_set).exclude(tvg_id__exact='')
|
||||
if streams_with_tvg.exists():
|
||||
instance.tvg_id = streams_with_tvg.first().tvg_id
|
||||
instance.save(update_fields=['tvg_id'])
|
||||
|
||||
@receiver(pre_save, sender=Stream)
|
||||
def set_default_m3u_account(sender, instance, **kwargs):
|
||||
"""
|
||||
This function will be triggered before saving a Stream instance.
|
||||
It sets the default m3u_account if not provided.
|
||||
"""
|
||||
if not instance.m3u_account:
|
||||
instance.is_custom = True
|
||||
default_account = M3UAccount.get_custom_account()
|
||||
|
||||
if default_account:
|
||||
instance.m3u_account = default_account
|
||||
else:
|
||||
raise ValueError("No default M3UAccount found.")
|
||||
|
||||
@receiver(post_save, sender=Stream)
|
||||
def generate_custom_stream_hash(sender, instance, created, **kwargs):
|
||||
"""
|
||||
Generate a stable stream_hash for custom streams after creation.
|
||||
Uses the stream's ID to ensure the hash never changes even if name/url is edited.
|
||||
"""
|
||||
if instance.is_custom and not instance.stream_hash and created:
|
||||
import hashlib
|
||||
# Use stream ID for a stable, unique hash that never changes
|
||||
unique_string = f"custom_stream_{instance.id}"
|
||||
instance.stream_hash = hashlib.sha256(unique_string.encode()).hexdigest()
|
||||
# Use update to avoid triggering signals again
|
||||
Stream.objects.filter(id=instance.id).update(stream_hash=instance.stream_hash)
|
||||
|
||||
@receiver(post_save, sender=Channel)
|
||||
def refresh_epg_programs(sender, instance, created, **kwargs):
|
||||
"""
|
||||
When a channel is saved, check if the EPG data has changed.
|
||||
If so, trigger a refresh of the program data for the EPG.
|
||||
"""
|
||||
# Check if this is an update (not a new channel) and the epg_data has changed
|
||||
if not created and kwargs.get('update_fields') and 'epg_data' in kwargs['update_fields']:
|
||||
logger.info(f"Channel {instance.id} ({instance.name}) EPG data updated, refreshing program data")
|
||||
if instance.epg_data:
|
||||
logger.info(f"Triggering EPG program refresh for {instance.epg_data.tvg_id}")
|
||||
parse_programs_for_tvg_id.delay(instance.epg_data.id)
|
||||
# For new channels with EPG data, also refresh
|
||||
elif created and instance.epg_data:
|
||||
logger.info(f"New channel {instance.id} ({instance.name}) created with EPG data, refreshing program data")
|
||||
parse_programs_for_tvg_id.delay(instance.epg_data.id)
|
||||
|
||||
@receiver(post_save, sender=ChannelProfile)
|
||||
def create_profile_memberships(sender, instance, created, **kwargs):
|
||||
if created:
|
||||
channels = Channel.objects.all()
|
||||
ChannelProfileMembership.objects.bulk_create([
|
||||
ChannelProfileMembership(channel_profile=instance, channel=channel)
|
||||
for channel in channels
|
||||
])
|
||||
|
||||
def _dvr_task_name(recording_id):
|
||||
"""Predictable PeriodicTask name for a DVR recording."""
|
||||
return f"dvr-recording-{recording_id}"
|
||||
|
||||
|
||||
def schedule_recording_task(instance, eta=None):
|
||||
"""Schedule a recording task via ClockedSchedule + one-off PeriodicTask.
|
||||
|
||||
The task is stored in the database and dispatched by Celery Beat at the
|
||||
scheduled time with no countdown. This avoids the Redis visibility_timeout
|
||||
redelivery bug that caused duplicate recordings when using apply_async
|
||||
with long countdowns.
|
||||
"""
|
||||
if eta is None:
|
||||
eta = instance.start_time
|
||||
if eta is not None and not is_aware(eta):
|
||||
eta = make_aware(eta)
|
||||
# Clamp to now so Beat dispatches immediately for past/current start times
|
||||
if eta <= now():
|
||||
eta = now()
|
||||
|
||||
task_args = [
|
||||
instance.id,
|
||||
instance.channel_id,
|
||||
str(instance.start_time),
|
||||
str(instance.end_time),
|
||||
]
|
||||
|
||||
clocked, _ = ClockedSchedule.objects.get_or_create(clocked_time=eta)
|
||||
task_name = _dvr_task_name(instance.id)
|
||||
PeriodicTask.objects.update_or_create(
|
||||
name=task_name,
|
||||
defaults={
|
||||
"task": "apps.channels.tasks.run_recording",
|
||||
"clocked": clocked,
|
||||
"args": json.dumps(task_args),
|
||||
"one_off": True,
|
||||
"enabled": True,
|
||||
"interval": None,
|
||||
"crontab": None,
|
||||
"solar": None,
|
||||
},
|
||||
)
|
||||
return task_name
|
||||
|
||||
|
||||
def revoke_task(task_id):
|
||||
"""Cancel a pending recording task.
|
||||
|
||||
task_id is normally a PeriodicTask name (e.g. "dvr-recording-42").
|
||||
For backwards compatibility with legacy Celery async-result UUIDs,
|
||||
falls back to AsyncResult.revoke().
|
||||
"""
|
||||
if not task_id:
|
||||
return
|
||||
# Primary path: delete the PeriodicTask and clean up its ClockedSchedule
|
||||
try:
|
||||
pt = PeriodicTask.objects.get(name=task_id)
|
||||
old_clocked = pt.clocked
|
||||
pt.delete()
|
||||
if old_clocked and not PeriodicTask.objects.filter(clocked=old_clocked).exists():
|
||||
old_clocked.delete()
|
||||
return
|
||||
except PeriodicTask.DoesNotExist:
|
||||
pass
|
||||
# Fallback for legacy Celery task UUIDs
|
||||
try:
|
||||
AsyncResult(task_id).revoke()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@receiver(pre_save, sender=Recording)
|
||||
def revoke_old_task_on_update(sender, instance, **kwargs):
|
||||
if not instance.pk:
|
||||
return # New instance
|
||||
try:
|
||||
old = Recording.objects.get(pk=instance.pk)
|
||||
if old.task_id and (
|
||||
old.start_time != instance.start_time or
|
||||
old.end_time != instance.end_time or
|
||||
old.channel_id != instance.channel_id
|
||||
):
|
||||
# Do NOT revoke while the recording is actively streaming.
|
||||
# run_recording re-reads end_time from the DB every ~2 s and extends
|
||||
# its internal deadline dynamically — revoking here would kill the task.
|
||||
old_status = (old.custom_properties or {}).get("status", "")
|
||||
if old_status == "recording":
|
||||
return
|
||||
revoke_task(old.task_id)
|
||||
instance.task_id = None
|
||||
except Recording.DoesNotExist:
|
||||
pass
|
||||
|
||||
@receiver(post_save, sender=Recording)
|
||||
def schedule_task_on_save(sender, instance, created, **kwargs):
|
||||
try:
|
||||
# Skip processing for internal field-only saves (metadata updates,
|
||||
# task_id assignment, end_time extensions) to prevent re-entrant
|
||||
# artwork dispatch and redundant recording_updated WS events.
|
||||
update_fields = kwargs.get('update_fields')
|
||||
if not created and update_fields is not None and set(update_fields) <= {'custom_properties', 'task_id', 'end_time'}:
|
||||
return
|
||||
|
||||
if not instance.task_id:
|
||||
start_time = instance.start_time
|
||||
end_time = instance.end_time
|
||||
|
||||
# Make datetimes aware (in UTC)
|
||||
if not is_aware(start_time):
|
||||
start_time = make_aware(start_time)
|
||||
if end_time and not is_aware(end_time):
|
||||
end_time = make_aware(end_time)
|
||||
|
||||
current_time = now()
|
||||
|
||||
if start_time > current_time - timedelta(seconds=1):
|
||||
# Future recording — schedule at start_time
|
||||
logger.info(f"Recording {instance.id}: scheduling task at {start_time}")
|
||||
task_id = schedule_recording_task(instance, eta=start_time)
|
||||
instance.task_id = task_id
|
||||
instance.save(update_fields=['task_id'])
|
||||
elif end_time and end_time > current_time:
|
||||
# Currently-playing — start immediately (e.g. series rule for in-progress program)
|
||||
logger.info(f"Recording {instance.id}: start_time in past but end_time still future, scheduling immediately")
|
||||
task_id = schedule_recording_task(instance, eta=current_time)
|
||||
instance.task_id = task_id
|
||||
instance.save(update_fields=['task_id'])
|
||||
else:
|
||||
logger.info(f"Recording {instance.id}: start_time and end_time both in past, not scheduling")
|
||||
# Kick off poster/artwork prefetch to enrich Upcoming cards.
|
||||
# Skip when the recording is already active or finished — run_recording
|
||||
# handles its own poster resolution, and scheduling artwork prefetch
|
||||
# while the task is running causes a race that can overwrite status.
|
||||
cp = instance.custom_properties or {}
|
||||
rec_status = cp.get("status", "")
|
||||
if rec_status not in ("recording", "completed", "stopped", "interrupted"):
|
||||
try:
|
||||
prefetch_recording_artwork.apply_async(args=[instance.id], countdown=1)
|
||||
except Exception as e:
|
||||
print("Error scheduling artwork prefetch:", e)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
print("Error in post_save signal:", e)
|
||||
traceback.print_exc()
|
||||
|
||||
@receiver(post_delete, sender=Recording)
|
||||
def revoke_task_on_delete(sender, instance, **kwargs):
|
||||
revoke_task(instance.task_id)
|
||||
Executable
+3973
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
@@ -0,0 +1,12 @@
|
||||
from django.urls import path
|
||||
from .views import StreamDashboardView, channels_dashboard_view
|
||||
|
||||
app_name = 'channels_dashboard'
|
||||
|
||||
urlpatterns = [
|
||||
# Example “dashboard” routes for streams
|
||||
path('streams/', StreamDashboardView.as_view(), name='stream_dashboard'),
|
||||
|
||||
# Example “dashboard” route for channels
|
||||
path('channels/', channels_dashboard_view, name='channels_dashboard'),
|
||||
]
|
||||
@@ -0,0 +1,25 @@
|
||||
import threading
|
||||
|
||||
lock = threading.Lock()
|
||||
# Dictionary to track usage: {account_id: current_usage}
|
||||
active_streams_map = {}
|
||||
|
||||
def increment_stream_count(account):
|
||||
with lock:
|
||||
current_usage = active_streams_map.get(account.id, 0)
|
||||
current_usage += 1
|
||||
active_streams_map[account.id] = current_usage
|
||||
account.active_streams = current_usage
|
||||
account.save(update_fields=['active_streams'])
|
||||
|
||||
def decrement_stream_count(account):
|
||||
with lock:
|
||||
current_usage = active_streams_map.get(account.id, 0)
|
||||
if current_usage > 0:
|
||||
current_usage -= 1
|
||||
if current_usage == 0:
|
||||
del active_streams_map[account.id]
|
||||
else:
|
||||
active_streams_map[account.id] = current_usage
|
||||
account.active_streams = current_usage
|
||||
account.save(update_fields=['active_streams'])
|
||||
@@ -0,0 +1,41 @@
|
||||
from django.views import View
|
||||
from django.http import JsonResponse
|
||||
from django.utils.decorators import method_decorator
|
||||
from django.contrib.auth.decorators import login_required
|
||||
from django.views.decorators.csrf import csrf_exempt
|
||||
from django.shortcuts import render
|
||||
|
||||
from .models import Stream
|
||||
|
||||
@method_decorator(csrf_exempt, name='dispatch')
|
||||
@method_decorator(login_required, name='dispatch')
|
||||
class StreamDashboardView(View):
|
||||
"""
|
||||
Example “dashboard” style view for Streams
|
||||
"""
|
||||
def get(self, request, *args, **kwargs):
|
||||
streams = Stream.objects.values(
|
||||
'id', 'name', 'url',
|
||||
'channel_group', 'current_viewers'
|
||||
)
|
||||
return JsonResponse({'data': list(streams)}, safe=False)
|
||||
|
||||
def post(self, request, *args, **kwargs):
|
||||
"""
|
||||
Creates a new Stream from JSON data
|
||||
"""
|
||||
import json
|
||||
try:
|
||||
data = json.loads(request.body)
|
||||
new_stream = Stream.objects.create(**data)
|
||||
return JsonResponse({
|
||||
'id': new_stream.id,
|
||||
'message': 'Stream created successfully!'
|
||||
}, status=201)
|
||||
except Exception as e:
|
||||
return JsonResponse({'error': str(e)}, status=400)
|
||||
|
||||
|
||||
@login_required
|
||||
def channels_dashboard_view(request):
|
||||
return render(request, 'channels/channels.html')
|
||||
Reference in New Issue
Block a user