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

This commit is contained in:
root
2026-05-09 21:24:50 +02:00
commit f56b088643
721 changed files with 177870 additions and 0 deletions
View File
+16
View File
@@ -0,0 +1,16 @@
from django.contrib import admin
from django.contrib.auth.admin import UserAdmin, GroupAdmin
from django.contrib.auth.models import Group
from .models import User
@admin.register(User)
class CustomUserAdmin(UserAdmin):
fieldsets = (
(None, {'fields': ('username', 'password', 'avatar_config', 'groups')}),
('Permissions', {'fields': ('is_staff', 'is_superuser', 'user_permissions')}),
('Important dates', {'fields': ('last_login', 'date_joined')}),
)
# Unregister default Group admin and re-register it.
admin.site.unregister(Group)
admin.site.register(Group, GroupAdmin)
+42
View File
@@ -0,0 +1,42 @@
from django.urls import path, include
from rest_framework.routers import DefaultRouter
from .api_views import (
AuthViewSet,
UserViewSet,
GroupViewSet,
APIKeyViewSet,
TokenObtainPairView,
TokenRefreshView,
list_permissions,
initialize_superuser,
)
from rest_framework_simplejwt import views as jwt_views
app_name = "accounts"
# 🔹 Register ViewSets with a Router
router = DefaultRouter()
router.register(r"users", UserViewSet, basename="user")
router.register(r"groups", GroupViewSet, basename="group")
router.register(r"api-keys", APIKeyViewSet, basename="api-key")
# 🔹 Custom Authentication Endpoints
auth_view = AuthViewSet.as_view({"post": "login"})
logout_view = AuthViewSet.as_view({"post": "logout"})
# 🔹 Define API URL patterns
urlpatterns = [
# Authentication
path("auth/login/", auth_view, name="user-login"),
path("auth/logout/", logout_view, name="user-logout"),
# Superuser API
path("initialize-superuser/", initialize_superuser, name="initialize_superuser"),
# Permissions API
path("permissions/", list_permissions, name="list-permissions"),
path("token/", TokenObtainPairView.as_view(), name="token_obtain_pair"),
path("token/refresh/", TokenRefreshView.as_view(), name="token_refresh"),
]
# 🔹 Include ViewSet routes
urlpatterns += router.urls
+366
View File
@@ -0,0 +1,366 @@
from django.contrib.auth import authenticate, login, logout
import logging
from django.contrib.auth.models import Group, Permission
from django.http import JsonResponse, HttpResponse
from django.views.decorators.csrf import csrf_exempt
from rest_framework.decorators import api_view, permission_classes, action
from rest_framework.response import Response
from rest_framework import viewsets, status, serializers
from rest_framework.throttling import AnonRateThrottle
from drf_spectacular.utils import extend_schema, OpenApiParameter, inline_serializer
from drf_spectacular.types import OpenApiTypes
import json
import secrets
from .permissions import IsAdmin, Authenticated
from dispatcharr.utils import network_access_allowed
from .models import User
from .serializers import UserSerializer, GroupSerializer, PermissionSerializer
from rest_framework_simplejwt.views import TokenObtainPairView, TokenRefreshView
logger = logging.getLogger(__name__)
class LoginRateThrottle(AnonRateThrottle):
scope = "login"
class TokenObtainPairView(TokenObtainPairView):
throttle_classes = [LoginRateThrottle]
def post(self, request, *args, **kwargs):
if not network_access_allowed(request, "UI"):
# Log blocked login attempt due to network restrictions
from core.utils import log_system_event
username = request.data.get("username", 'unknown')
client_ip = request.META.get('REMOTE_ADDR', 'unknown')
user_agent = request.META.get('HTTP_USER_AGENT', 'unknown')
logger.info(f"Login blocked by network policy: user={username} ip={client_ip} ua={user_agent}")
log_system_event(
event_type='login_failed',
user=username,
client_ip=client_ip,
user_agent=user_agent,
reason='Network access denied',
)
return Response({"error": "Forbidden"}, status=status.HTTP_403_FORBIDDEN)
# Get the response from the parent class first
username = request.data.get("username")
# Log login attempt
from core.utils import log_system_event
client_ip = request.META.get('REMOTE_ADDR', 'unknown')
user_agent = request.META.get('HTTP_USER_AGENT', 'unknown')
try:
logger.debug(f"Attempting JWT login for user={username}")
response = super().post(request, *args, **kwargs)
# If login was successful, update last_login and log success
if response.status_code == 200:
if username:
from django.utils import timezone
try:
user = User.objects.get(username=username)
user.last_login = timezone.now()
user.save(update_fields=['last_login'])
# Log successful login
log_system_event(
event_type='login_success',
user=username,
client_ip=client_ip,
user_agent=user_agent,
)
logger.info(f"Login success: user={username} ip={client_ip}")
except User.DoesNotExist:
pass # User doesn't exist, but login somehow succeeded
else:
# Log failed login attempt
log_system_event(
event_type='login_failed',
user=username or 'unknown',
client_ip=client_ip,
user_agent=user_agent,
reason='Invalid credentials',
)
logger.info(f"Login failed: user={username} ip={client_ip}")
return response
except Exception as e:
# If parent class raises an exception (e.g., validation error), log failed attempt
log_system_event(
event_type='login_failed',
user=username or 'unknown',
client_ip=client_ip,
user_agent=user_agent,
reason=f'Authentication error: {str(e)[:100]}',
)
logger.error(f"Login error for user={username}: {e}")
raise # Re-raise the exception to maintain normal error flow
class TokenRefreshView(TokenRefreshView):
def post(self, request, *args, **kwargs):
# Custom logic here
if not network_access_allowed(request, "UI"):
# Log blocked token refresh attempt due to network restrictions
from core.utils import log_system_event
client_ip = request.META.get('REMOTE_ADDR', 'unknown')
user_agent = request.META.get('HTTP_USER_AGENT', 'unknown')
logger.info(f"Token refresh blocked by network policy: ip={client_ip} ua={user_agent}")
log_system_event(
event_type='login_failed',
user='token_refresh',
client_ip=client_ip,
user_agent=user_agent,
reason='Network access denied (token refresh)',
)
return Response({"error": "Unauthorized"}, status=status.HTTP_403_FORBIDDEN)
return super().post(request, *args, **kwargs)
@csrf_exempt # In production, consider CSRF protection strategies or ensure this endpoint is only accessible when no superuser exists.
def initialize_superuser(request):
# If an admin-level user already exists, the system is configured
if User.objects.filter(user_level__gte=10).exists():
return JsonResponse({"superuser_exists": True})
if request.method == "POST":
try:
data = json.loads(request.body)
username = data.get("username")
password = data.get("password")
email = data.get("email", "")
if not username or not password:
return JsonResponse(
{"error": "Username and password are required."}, status=400
)
# Create the superuser
User.objects.create_superuser(
username=username, password=password, email=email, user_level=10
)
return JsonResponse({"superuser_exists": True})
except Exception as e:
return JsonResponse({"error": str(e)}, status=500)
# For GET requests, indicate no superuser exists
return JsonResponse({"superuser_exists": False})
# 🔹 1) Authentication APIs
class AuthViewSet(viewsets.ViewSet):
"""Handles user login and logout"""
def get_permissions(self):
"""
Login doesn't require auth, but logout does
"""
if self.action == 'logout':
return [Authenticated()]
return []
@extend_schema(
description="Alias for POST /api/accounts/token/ — returns JWT access and refresh tokens.",
request=inline_serializer(
name="LoginRequest",
fields={
"username": serializers.CharField(),
"password": serializers.CharField(),
},
),
)
def login(self, request):
"""Delegates to TokenObtainPairView (JWT login). Throttling, logging, and
network access checks are handled there."""
view = TokenObtainPairView.as_view()
return view(request._request)
@extend_schema(
description="Log out the current user",
)
def logout(self, request):
"""Logs out the authenticated user"""
# Log logout event before actually logging out
from core.utils import log_system_event
username = request.user.username if request.user and request.user.is_authenticated else 'unknown'
client_ip = request.META.get('REMOTE_ADDR', 'unknown')
user_agent = request.META.get('HTTP_USER_AGENT', 'unknown')
log_system_event(
event_type='logout',
user=username,
client_ip=client_ip,
user_agent=user_agent,
)
logger.info(f"Logout: user={username} ip={client_ip}")
logout(request)
return Response({"message": "Logout successful"})
# 🔹 2) User Management APIs
class UserViewSet(viewsets.ModelViewSet):
"""Handles CRUD operations for Users"""
queryset = User.objects.all().prefetch_related('channel_profiles')
serializer_class = UserSerializer
def get_permissions(self):
if self.action == "me":
return [Authenticated()]
return [IsAdmin()]
@extend_schema(
description="Retrieve a list of users",
responses={200: UserSerializer(many=True)},
)
def list(self, request, *args, **kwargs):
return super().list(request, *args, **kwargs)
@extend_schema(description="Retrieve a specific user by ID")
def retrieve(self, request, *args, **kwargs):
return super().retrieve(request, *args, **kwargs)
@extend_schema(description="Create a new user")
def create(self, request, *args, **kwargs):
return super().create(request, *args, **kwargs)
@extend_schema(description="Update a user")
def update(self, request, *args, **kwargs):
return super().update(request, *args, **kwargs)
@extend_schema(description="Delete a user")
def destroy(self, request, *args, **kwargs):
return super().destroy(request, *args, **kwargs)
@extend_schema(
description="Get or update active user information. PATCH updates custom_properties with merge semantics.",
methods=["GET", "PATCH"],
)
@action(detail=False, methods=["get", "patch"], url_path="me")
def me(self, request):
user = request.user
if request.method == "PATCH":
ALLOWED_FIELDS = {"custom_properties", "first_name", "last_name", "email", "password"}
disallowed = set(request.data.keys()) - ALLOWED_FIELDS
for key in disallowed:
request.data.pop(key, None)
# Strip admin-managed keys from custom_properties so users cannot
# set their own XC credentials via this endpoint.
ADMIN_ONLY_PROPS = {"xc_password"}
cp = request.data.get("custom_properties")
if isinstance(cp, dict):
for key in ADMIN_ONLY_PROPS:
cp.pop(key, None)
serializer = UserSerializer(user, data=request.data, partial=True)
serializer.is_valid(raise_exception=True)
serializer.save()
return Response(serializer.data)
serializer = UserSerializer(user)
return Response(serializer.data)
# 🔹 3) Group Management APIs
class GroupViewSet(viewsets.ModelViewSet):
"""Handles CRUD operations for Groups"""
queryset = Group.objects.all()
serializer_class = GroupSerializer
permission_classes = [Authenticated]
@extend_schema(
description="Retrieve a list of groups",
responses={200: GroupSerializer(many=True)},
)
def list(self, request, *args, **kwargs):
return super().list(request, *args, **kwargs)
@extend_schema(description="Retrieve a specific group by ID")
def retrieve(self, request, *args, **kwargs):
return super().retrieve(request, *args, **kwargs)
@extend_schema(description="Create a new group")
def create(self, request, *args, **kwargs):
return super().create(request, *args, **kwargs)
@extend_schema(description="Update a group")
def update(self, request, *args, **kwargs):
return super().update(request, *args, **kwargs)
@extend_schema(description="Delete a group")
def destroy(self, request, *args, **kwargs):
return super().destroy(request, *args, **kwargs)
# API Key management
class APIKeyViewSet(viewsets.ViewSet):
permission_classes = [Authenticated]
def list(self, request):
user = request.user
return Response({"key": user.api_key})
@action(detail=False, methods=["post"], url_path="generate")
def generate(self, request):
target_user = request.user
user_id = request.data.get("user_id")
if user_id:
from .permissions import IsAdmin
if not IsAdmin().has_permission(request, self):
return Response({"detail": "Not allowed to create keys for other users."}, status=status.HTTP_403_FORBIDDEN)
try:
target_user = User.objects.get(id=int(user_id))
except Exception:
return Response({"detail": "User not found."}, status=status.HTTP_404_NOT_FOUND)
raw = secrets.token_urlsafe(40)
target_user.api_key = raw
target_user.save(update_fields=["api_key"])
user_data = UserSerializer(target_user).data
return Response({"key": raw, "user": user_data}, status=status.HTTP_201_CREATED)
@action(detail=False, methods=["post"], url_path="revoke")
def revoke(self, request):
target_user = request.user
user_id = request.data.get("user_id")
if user_id:
from .permissions import IsAdmin
if not IsAdmin().has_permission(request, self):
return Response({"detail": "Not allowed to revoke keys for other users."}, status=status.HTTP_403_FORBIDDEN)
try:
target_user = User.objects.get(id=int(user_id))
except Exception:
return Response({"detail": "User not found."}, status=status.HTTP_404_NOT_FOUND)
target_user.api_key = None
target_user.save(update_fields=["api_key"])
return Response({"success": True})
# 🔹 4) Permissions List API
@extend_schema(
description="Retrieve a list of all permissions",
responses={200: PermissionSerializer(many=True)},
)
@api_view(["GET"])
@permission_classes([Authenticated])
def list_permissions(request):
"""Returns a list of all available permissions"""
permissions = Permission.objects.all()
serializer = PermissionSerializer(permissions, many=True)
return Response(serializer.data)
+7
View File
@@ -0,0 +1,7 @@
from django.apps import AppConfig
class AccountsConfig(AppConfig):
default_auto_field = "django.db.models.BigAutoField"
name = "apps.accounts"
verbose_name = "Accounts & Authentication"
+86
View File
@@ -0,0 +1,86 @@
from rest_framework import authentication
from rest_framework import exceptions
from django.conf import settings
from drf_spectacular.extensions import OpenApiAuthenticationExtension
from .models import User
class JWTAuthenticationScheme(OpenApiAuthenticationExtension):
target_class = "rest_framework_simplejwt.authentication.JWTAuthentication"
name = "jwtAuth"
def get_security_definition(self, auto_schema):
return {
"type": "http",
"scheme": "bearer",
"bearerFormat": "JWT",
"description": (
"JWT Bearer authentication.\n\n"
"Obtain a token pair via `POST /api/accounts/token/` using your username and password, "
"then paste the **access token** here — Swagger adds the `Bearer ` prefix automatically.\n\n"
"Access tokens expire after 30 minutes. Refresh using `POST /api/accounts/token/refresh/`."
),
}
class ApiKeyAuthenticationScheme(OpenApiAuthenticationExtension):
target_class = "apps.accounts.authentication.ApiKeyAuthentication"
name = "ApiKeyAuth"
def get_security_definition(self, auto_schema):
return {
"type": "apiKey",
"in": "header",
"name": "X-API-Key",
"description": (
"API key authentication.\n\n"
"Pass your personal API key in the `X-API-Key` request header. "
"Keys can be generated via `POST /api/accounts/api-keys/generate/` "
"and revoked via `POST /api/accounts/api-keys/revoke/`."
),
}
class ApiKeyAuthentication(authentication.BaseAuthentication):
"""
Accepts header `Authorization: ApiKey <key>` or `X-API-Key: <key>`.
"""
keyword = "ApiKey"
def authenticate(self, request):
# Check X-API-Key header first
raw_key = request.META.get("HTTP_X_API_KEY")
if not raw_key:
auth = authentication.get_authorization_header(request).split()
if not auth:
return None
if len(auth) != 2:
return None
scheme = auth[0].decode().lower()
if scheme != self.keyword.lower():
return None
raw_key = auth[1].decode()
if not raw_key:
return None
if not raw_key:
return None
try:
user = User.objects.get(api_key=raw_key)
except User.DoesNotExist:
raise exceptions.AuthenticationFailed("Invalid API key")
if not user.is_active:
raise exceptions.AuthenticationFailed("User inactive")
return (user, None)
def authenticate_header(self, request):
return self.keyword
+59
View File
@@ -0,0 +1,59 @@
from django import forms
from django.contrib.auth.forms import UserCreationForm
from django.contrib.auth.models import Permission
from django.contrib.auth.models import Group as AuthGroup
from apps.channels.models import ChannelGroup
from .models import User
from .models import User
class UserRegistrationForm(UserCreationForm):
groups = forms.ModelMultipleChoiceField(
queryset=AuthGroup.objects.all(),
required=False,
widget=forms.CheckboxSelectMultiple
)
class Meta:
model = User
fields = ['username', 'groups', 'password1', 'password2', ]
def save(self, commit=True):
user = super().save(commit=False)
if commit:
user.save()
self.save_m2m() # Save the many-to-many field data
return user
class GroupForm(forms.ModelForm):
permissions = forms.ModelMultipleChoiceField(
queryset=Permission.objects.all(),
widget=forms.CheckboxSelectMultiple,
required=False
)
class Meta:
model = AuthGroup
fields = ['name', 'permissions']
class UserEditForm(forms.ModelForm):
auth_groups = forms.ModelMultipleChoiceField(
queryset=AuthGroup.objects.all(),
widget=forms.CheckboxSelectMultiple,
required=False,
label="Auth Groups"
)
channel_groups = forms.ModelMultipleChoiceField(
queryset=ChannelGroup.objects.all(),
widget=forms.CheckboxSelectMultiple,
required=False,
label="Channel Groups"
)
class Meta:
model = User
fields = ['username', 'email', 'auth_groups', 'channel_groups']
+47
View File
@@ -0,0 +1,47 @@
# Generated by Django 5.1.6 on 2025-03-05 22:07
import django.contrib.auth.models
import django.contrib.auth.validators
import django.utils.timezone
from django.db import migrations, models
class Migration(migrations.Migration):
initial = True
dependencies = [
('auth', '0012_alter_user_first_name_max_length'),
('dispatcharr_channels', '0001_initial'),
]
operations = [
migrations.CreateModel(
name='User',
fields=[
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('password', models.CharField(max_length=128, verbose_name='password')),
('last_login', models.DateTimeField(blank=True, null=True, verbose_name='last login')),
('is_superuser', models.BooleanField(default=False, help_text='Designates that this user has all permissions without explicitly assigning them.', verbose_name='superuser status')),
('username', models.CharField(error_messages={'unique': 'A user with that username already exists.'}, help_text='Required. 150 characters or fewer. Letters, digits and @/./+/-/_ only.', max_length=150, unique=True, validators=[django.contrib.auth.validators.UnicodeUsernameValidator()], verbose_name='username')),
('first_name', models.CharField(blank=True, max_length=150, verbose_name='first name')),
('last_name', models.CharField(blank=True, max_length=150, verbose_name='last name')),
('email', models.EmailField(blank=True, max_length=254, verbose_name='email address')),
('is_staff', models.BooleanField(default=False, help_text='Designates whether the user can log into this admin site.', verbose_name='staff status')),
('is_active', models.BooleanField(default=True, help_text='Designates whether this user should be treated as active. Unselect this instead of deleting accounts.', verbose_name='active')),
('date_joined', models.DateTimeField(default=django.utils.timezone.now, verbose_name='date joined')),
('avatar_config', models.JSONField(blank=True, default=dict, null=True)),
('channel_groups', models.ManyToManyField(blank=True, related_name='users', to='dispatcharr_channels.channelgroup')),
('groups', models.ManyToManyField(blank=True, help_text='The groups this user belongs to. A user will get all permissions granted to each of their groups.', related_name='user_set', related_query_name='user', to='auth.group', verbose_name='groups')),
('user_permissions', models.ManyToManyField(blank=True, help_text='Specific permissions for this user.', related_name='user_set', related_query_name='user', to='auth.permission', verbose_name='user permissions')),
],
options={
'verbose_name': 'user',
'verbose_name_plural': 'users',
'abstract': False,
},
managers=[
('objects', django.contrib.auth.models.UserManager()),
],
),
]
@@ -0,0 +1,43 @@
# Generated by Django 5.1.6 on 2025-05-18 15:47
from django.db import migrations, models
def set_user_level_to_10(apps, schema_editor):
User = apps.get_model("accounts", "User")
User.objects.update(user_level=10)
class Migration(migrations.Migration):
dependencies = [
("accounts", "0001_initial"),
("dispatcharr_channels", "0021_channel_user_level"),
]
operations = [
migrations.RemoveField(
model_name="user",
name="channel_groups",
),
migrations.AddField(
model_name="user",
name="channel_profiles",
field=models.ManyToManyField(
blank=True,
related_name="users",
to="dispatcharr_channels.channelprofile",
),
),
migrations.AddField(
model_name="user",
name="user_level",
field=models.IntegerField(default=0),
),
migrations.AddField(
model_name="user",
name="custom_properties",
field=models.TextField(blank=True, null=True),
),
migrations.RunPython(set_user_level_to_10),
]
@@ -0,0 +1,18 @@
# Generated by Django 5.2.4 on 2025-09-02 14:30
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('accounts', '0002_remove_user_channel_groups_user_channel_profiles_and_more'),
]
operations = [
migrations.AlterField(
model_name='user',
name='custom_properties',
field=models.JSONField(blank=True, default=dict, null=True),
),
]
@@ -0,0 +1,18 @@
# Generated by Django 5.2.11 on 2026-02-21 18:14
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('accounts', '0003_alter_user_custom_properties'),
]
operations = [
migrations.AddField(
model_name='user',
name='api_key',
field=models.CharField(blank=True, db_index=True, max_length=200, null=True),
),
]
@@ -0,0 +1,20 @@
# Generated by Django 5.2.11 on 2026-02-26 19:24
import apps.accounts.models
from django.db import migrations
class Migration(migrations.Migration):
dependencies = [
('accounts', '0004_user_api_key'),
]
operations = [
migrations.AlterModelManagers(
name='user',
managers=[
('objects', apps.accounts.models.CustomUserManager()),
],
),
]
@@ -0,0 +1,18 @@
# Generated by Django 5.2.11 on 2026-03-19 13:46
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('accounts', '0005_alter_user_managers'),
]
operations = [
migrations.AddField(
model_name='user',
name='stream_limit',
field=models.IntegerField(default=0),
),
]
+48
View File
@@ -0,0 +1,48 @@
# apps/accounts/models.py
from django.db import models
from django.contrib.auth.models import AbstractUser, Permission, UserManager
class CustomUserManager(UserManager):
def create_superuser(self, username, email=None, password=None, **extra_fields):
extra_fields.setdefault('user_level', 10)
return super().create_superuser(username, email, password, **extra_fields)
class User(AbstractUser):
objects = CustomUserManager()
"""
Custom user model for Dispatcharr.
Inherits from Django's AbstractUser to add additional fields if needed.
"""
class UserLevel(models.IntegerChoices):
STREAMER = 0, "Streamer"
STANDARD = 1, "Standard User"
ADMIN = 10, "Admin"
avatar_config = models.JSONField(default=dict, blank=True, null=True)
channel_profiles = models.ManyToManyField(
"dispatcharr_channels.ChannelProfile",
blank=True,
related_name="users",
)
user_level = models.IntegerField(default=UserLevel.STREAMER)
custom_properties = models.JSONField(default=dict, blank=True, null=True)
api_key = models.CharField(max_length=200, blank=True, null=True, db_index=True)
stream_limit = models.IntegerField(default=0)
def __str__(self):
return self.username
def get_groups(self):
"""
Returns the groups (roles) the user belongs to.
"""
return self.groups.all()
def get_permissions(self):
"""
Returns the permissions assigned to the user and their groups.
"""
return self.user_permissions.all() | Permission.objects.filter(group__user=self)
+56
View File
@@ -0,0 +1,56 @@
from rest_framework.permissions import IsAuthenticated
from .models import User
from dispatcharr.utils import network_access_allowed
class Authenticated(IsAuthenticated):
def has_permission(self, request, view):
is_authenticated = super().has_permission(request, view)
network_allowed = network_access_allowed(request, "UI")
return is_authenticated and network_allowed
class IsStandardUser(Authenticated):
def has_permission(self, request, view):
if not super().has_permission(request, view):
return False
return request.user and request.user.user_level >= User.UserLevel.STANDARD
class IsAdmin(Authenticated):
def has_permission(self, request, view):
if not super().has_permission(request, view):
return False
return request.user.user_level >= 10
class IsOwnerOfObject(Authenticated):
def has_object_permission(self, request, view, obj):
if not super().has_permission(request, view):
return False
is_admin = IsAdmin().has_permission(request, view)
is_owner = request.user in obj.users.all()
return is_admin or is_owner
permission_classes_by_action = {
"list": [IsStandardUser],
"create": [IsAdmin],
"retrieve": [IsStandardUser],
"update": [IsAdmin],
"partial_update": [IsAdmin],
"destroy": [IsAdmin],
}
permission_classes_by_method = {
"GET": [IsStandardUser],
"POST": [IsAdmin],
"PATCH": [IsAdmin],
"PUT": [IsAdmin],
"DELETE": [IsAdmin],
}
+145
View File
@@ -0,0 +1,145 @@
import json
from rest_framework import serializers
from django.contrib.auth.models import Group, Permission
from .models import User
from apps.channels.models import ChannelProfile
# Valid navigation item IDs for validation
VALID_NAV_ITEM_IDS = {
'channels', 'vods', 'sources', 'guide', 'dvr',
'stats', 'plugins', 'integrations', 'system', 'settings'
}
MAX_CUSTOM_PROPS_SIZE = 10240 # 10KB limit
def validate_nav_array(value, field_name):
"""Validate that a value is an array of valid nav item ID strings."""
if not isinstance(value, list):
raise serializers.ValidationError(f"{field_name} must be an array")
if len(value) > 50:
raise serializers.ValidationError(f"{field_name} exceeds maximum length of 50 items")
for item in value:
if not isinstance(item, str):
raise serializers.ValidationError(f"{field_name} items must be strings")
if item not in VALID_NAV_ITEM_IDS:
raise serializers.ValidationError(f"'{item}' is not a valid navigation item ID")
# 🔹 Fix for Permission serialization
class PermissionSerializer(serializers.ModelSerializer):
class Meta:
model = Permission
fields = ["id", "name", "codename"]
# 🔹 Fix for Group serialization
class GroupSerializer(serializers.ModelSerializer):
permissions = serializers.PrimaryKeyRelatedField(
many=True, queryset=Permission.objects.all()
) # ✅ Fixes ManyToManyField `_meta` error
class Meta:
model = Group
fields = ["id", "name", "permissions"]
# 🔹 Fix for User serialization
class UserSerializer(serializers.ModelSerializer):
password = serializers.CharField(write_only=True, required=False)
channel_profiles = serializers.PrimaryKeyRelatedField(
queryset=ChannelProfile.objects.all(), many=True, required=False
)
api_key = serializers.CharField(read_only=True, allow_null=True)
class Meta:
model = User
fields = [
"id",
"username",
"api_key",
"email",
"user_level",
"password",
"channel_profiles",
"custom_properties",
"avatar_config",
"stream_limit",
"is_staff",
"is_superuser",
"last_login",
"date_joined",
"first_name",
"last_name",
]
def validate_custom_properties(self, value):
"""Validate custom_properties structure and size."""
if value is None:
return {}
if not isinstance(value, dict):
raise serializers.ValidationError("custom_properties must be a dictionary")
# Size limit check
try:
if len(json.dumps(value)) > MAX_CUSTOM_PROPS_SIZE:
raise serializers.ValidationError(
f"custom_properties exceeds maximum size of {MAX_CUSTOM_PROPS_SIZE} bytes"
)
except (TypeError, ValueError):
raise serializers.ValidationError("custom_properties contains non-serializable data")
# Validate navOrder if present
if 'navOrder' in value:
validate_nav_array(value['navOrder'], 'navOrder')
# Validate hiddenNav if present
if 'hiddenNav' in value:
validate_nav_array(value['hiddenNav'], 'hiddenNav')
return value
def create(self, validated_data):
channel_profiles = validated_data.pop("channel_profiles", [])
user = User(**validated_data)
user.set_password(validated_data["password"])
user.save()
user.channel_profiles.set(channel_profiles)
return user
def update(self, instance, validated_data):
password = validated_data.pop("password", None)
channel_profiles = validated_data.pop("channel_profiles", None)
# Merge custom_properties instead of replacing (prevents data loss)
# Strip null values — sending null for a key omits it rather than overwriting with null
custom_properties = validated_data.pop("custom_properties", None)
if custom_properties is not None:
existing = instance.custom_properties or {}
cleaned = {k: v for k, v in custom_properties.items() if v is not None}
merged = {**existing, **cleaned}
# Scrub stale nav IDs so the DB self-heals on next save
for nav_field in ('navOrder', 'hiddenNav'):
if nav_field in merged and isinstance(merged[nav_field], list):
merged[nav_field] = [
item for item in merged[nav_field]
if item in VALID_NAV_ITEM_IDS
]
instance.custom_properties = merged
for attr, value in validated_data.items():
setattr(instance, attr, value)
if password:
instance.set_password(password)
instance.save()
if channel_profiles is not None:
instance.channel_profiles.set(channel_profiles)
return instance
+15
View File
@@ -0,0 +1,15 @@
# apps/accounts/signals.py
# Example: automatically create something on user creation
from django.db.models.signals import post_save
from django.dispatch import receiver
from .models import User
@receiver(post_save, sender=User)
def handle_new_user(sender, instance, created, **kwargs):
if created:
# e.g. initialize default avatar config
if not instance.avatar_config:
instance.avatar_config = {"style": "circle"}
instance.save()
+72
View File
@@ -0,0 +1,72 @@
from django.test import TestCase
from django.contrib.auth import get_user_model
from rest_framework.test import APIClient
User = get_user_model()
class InitializeSuperuserTests(TestCase):
"""Tests for the initialize_superuser endpoint"""
def setUp(self):
self.client = APIClient()
self.url = "/api/accounts/initialize-superuser/"
def test_returns_true_when_superuser_exists(self):
"""Superuser with is_superuser=True should be detected"""
User.objects.create_superuser(
username="admin", password="testpass123", user_level=10
)
response = self.client.get(self.url)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.json()["superuser_exists"])
def test_returns_true_when_admin_level_user_exists(self):
"""User with user_level=10 but is_superuser=False should be detected"""
user = User.objects.create_user(username="admin", password="testpass123")
user.user_level = 10
user.is_superuser = False
user.save()
response = self.client.get(self.url)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.json()["superuser_exists"])
def test_returns_false_when_no_admin_exists(self):
"""No admin or superuser should return false"""
# Create a non-admin user
User.objects.create_user(username="regular", password="testpass123")
response = self.client.get(self.url)
self.assertEqual(response.status_code, 200)
self.assertFalse(response.json()["superuser_exists"])
def test_returns_false_when_no_users_exist(self):
"""Empty database should return false"""
response = self.client.get(self.url)
self.assertEqual(response.status_code, 200)
self.assertFalse(response.json()["superuser_exists"])
def test_create_superuser_when_none_exists(self):
"""POST should create superuser when none exists"""
response = self.client.post(
self.url,
{"username": "newadmin", "password": "testpass123", "email": "admin@test.com"},
format="json",
)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.json()["superuser_exists"])
self.assertTrue(User.objects.filter(username="newadmin", user_level=10).exists())
def test_cannot_create_superuser_when_admin_exists(self):
"""POST should fail when an admin-level user already exists"""
user = User.objects.create_user(username="existing", password="testpass123")
user.user_level = 10
user.save()
response = self.client.post(
self.url,
{"username": "newadmin", "password": "testpass123"},
format="json",
)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.json()["superuser_exists"])
# Should NOT have created a new user
self.assertFalse(User.objects.filter(username="newadmin").exists())
+12
View File
@@ -0,0 +1,12 @@
from django.urls import path
from django.contrib.auth import views as auth_views
app_name = 'accounts'
urlpatterns = [
# Login view using Django's built-in authentication
path('login/', auth_views.LoginView.as_view(template_name='login.html'), name='login'),
# Logout view using Django's built-in authentication
path('logout/', auth_views.LogoutView.as_view(next_page='accounts:login'), name='logout'),
# Onetime use superuser account creation
]
+29
View File
@@ -0,0 +1,29 @@
from django.urls import path, include, re_path
from drf_spectacular.views import SpectacularAPIView, SpectacularSwaggerView, SpectacularRedocView
app_name = 'api'
urlpatterns = [
path('accounts/', include(('apps.accounts.api_urls', 'accounts'), namespace='accounts')),
path('channels/', include(('apps.channels.api_urls', 'channels'), namespace='channels')),
path('epg/', include(('apps.epg.api_urls', 'epg'), namespace='epg')),
path('hdhr/', include(('apps.hdhr.api_urls', 'hdhr'), namespace='hdhr')),
path('m3u/', include(('apps.m3u.api_urls', 'm3u'), namespace='m3u')),
path('core/', include(('core.api_urls', 'core'), namespace='core')),
path('plugins/', include(('apps.plugins.api_urls', 'plugins'), namespace='plugins')),
path('vod/', include(('apps.vod.api_urls', 'vod'), namespace='vod')),
path('backups/', include(('apps.backups.api_urls', 'backups'), namespace='backups')),
path('connect/', include(('apps.connect.api_urls', 'connect'), namespace='connect')),
# path('output/', include(('apps.output.api_urls', 'output'), namespace='output')),
#path('player/', include(('apps.player.api_urls', 'player'), namespace='player')),
#path('settings/', include(('apps.settings.api_urls', 'settings'), namespace='settings')),
#path('streams/', include(('apps.streams.api_urls', 'streams'), namespace='streams')),
# OpenAPI Schema and Documentation (drf-spectacular)
path('schema/', SpectacularAPIView.as_view(), name='schema'),
re_path(r'^swagger/?$', SpectacularSwaggerView.as_view(url_name='api:schema'), name='swagger-ui'),
path('redoc/', SpectacularRedocView.as_view(url_name='api:schema'), name='redoc'),
path('swagger.json', SpectacularAPIView.as_view(), name='schema-json'),
]
View File
+18
View File
@@ -0,0 +1,18 @@
from django.urls import path
from . import api_views
app_name = "backups"
urlpatterns = [
path("", api_views.list_backups, name="backup-list"),
path("create/", api_views.create_backup, name="backup-create"),
path("upload/", api_views.upload_backup, name="backup-upload"),
path("schedule/", api_views.get_schedule, name="backup-schedule-get"),
path("schedule/update/", api_views.update_schedule, name="backup-schedule-update"),
path("status/<str:task_id>/", api_views.backup_status, name="backup-status"),
path("<str:filename>/download-token/", api_views.get_download_token, name="backup-download-token"),
path("<str:filename>/download/", api_views.download_backup, name="backup-download"),
path("<str:filename>/delete/", api_views.delete_backup, name="backup-delete"),
path("<str:filename>/restore/", api_views.restore_backup, name="backup-restore"),
]
+374
View File
@@ -0,0 +1,374 @@
import hashlib
import hmac
import logging
import os
from pathlib import Path
from celery.result import AsyncResult
from django.conf import settings
from django.http import HttpResponse, StreamingHttpResponse, Http404
from rest_framework import status
from rest_framework.decorators import api_view, permission_classes, parser_classes
from rest_framework.permissions import AllowAny
from apps.accounts.permissions import IsAdmin
from rest_framework.parsers import MultiPartParser, FormParser
from rest_framework.response import Response
from core.utils import safe_upload_path
from . import services
from .tasks import create_backup_task, restore_backup_task
from .scheduler import get_schedule_settings, update_schedule_settings
logger = logging.getLogger(__name__)
def _generate_task_token(task_id: str) -> str:
"""Generate a signed token for task status access without auth."""
secret = settings.SECRET_KEY.encode()
return hmac.new(secret, task_id.encode(), hashlib.sha256).hexdigest()[:32]
def _verify_task_token(task_id: str, token: str) -> bool:
"""Verify a task token is valid."""
expected = _generate_task_token(task_id)
return hmac.compare_digest(expected, token)
@api_view(["GET"])
@permission_classes([IsAdmin])
def list_backups(request):
"""List all available backup files."""
try:
backups = services.list_backups()
return Response(backups, status=status.HTTP_200_OK)
except Exception as e:
return Response(
{"detail": f"Failed to list backups: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["POST"])
@permission_classes([IsAdmin])
def create_backup(request):
"""Create a new backup (async via Celery)."""
try:
task = create_backup_task.delay()
return Response(
{
"detail": "Backup started",
"task_id": task.id,
"task_token": _generate_task_token(task.id),
},
status=status.HTTP_202_ACCEPTED,
)
except Exception as e:
return Response(
{"detail": f"Failed to start backup: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["GET"])
@permission_classes([AllowAny])
def backup_status(request, task_id):
"""Check the status of a backup/restore task.
Requires either:
- Valid admin authentication, OR
- Valid task_token query parameter
"""
# Check for token-based auth (for restore when session is invalidated)
token = request.query_params.get("token")
if token:
if not _verify_task_token(task_id, token):
return Response(
{"detail": "Invalid task token"},
status=status.HTTP_403_FORBIDDEN,
)
else:
# Fall back to admin auth check
if not request.user.is_authenticated or getattr(request.user, 'user_level', 0) < 10:
return Response(
{"detail": "Authentication required"},
status=status.HTTP_401_UNAUTHORIZED,
)
try:
result = AsyncResult(task_id)
if result.ready():
task_result = result.get()
if task_result.get("status") == "completed":
return Response({
"state": "completed",
"result": task_result,
})
else:
return Response({
"state": "failed",
"error": task_result.get("error", "Unknown error"),
})
elif result.failed():
return Response({
"state": "failed",
"error": str(result.result),
})
else:
return Response({
"state": result.state.lower(),
})
except Exception as e:
return Response(
{"detail": f"Failed to get task status: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["GET"])
@permission_classes([IsAdmin])
def get_download_token(request, filename):
"""Get a signed token for downloading a backup file."""
try:
# Security: prevent path traversal
if ".." in filename or "/" in filename or "\\" in filename:
raise Http404("Invalid filename")
backup_dir = services.get_backup_dir()
backup_file = backup_dir / filename
if not backup_file.exists():
raise Http404("Backup file not found")
token = _generate_task_token(filename)
return Response({"token": token})
except Http404:
raise
except Exception as e:
return Response(
{"detail": f"Failed to generate token: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["GET"])
@permission_classes([AllowAny])
def download_backup(request, filename):
"""Download a backup file.
Requires either:
- Valid admin authentication, OR
- Valid download_token query parameter
"""
# Check for token-based auth (avoids CORS preflight issues)
token = request.query_params.get("token")
if token:
if not _verify_task_token(filename, token):
return Response(
{"detail": "Invalid download token"},
status=status.HTTP_403_FORBIDDEN,
)
else:
# Fall back to admin auth check
if not request.user.is_authenticated or getattr(request.user, 'user_level', 0) < 10:
return Response(
{"detail": "Authentication required"},
status=status.HTTP_401_UNAUTHORIZED,
)
try:
# Security: prevent path traversal by checking for suspicious characters
if ".." in filename or "/" in filename or "\\" in filename:
raise Http404("Invalid filename")
backup_dir = services.get_backup_dir()
backup_file = (backup_dir / filename).resolve()
# Security: ensure the resolved path is still within backup_dir
if not str(backup_file).startswith(str(backup_dir.resolve())):
raise Http404("Invalid filename")
if not backup_file.exists() or not backup_file.is_file():
raise Http404("Backup file not found")
file_size = backup_file.stat().st_size
# Use X-Accel-Redirect for nginx (AIO container) - nginx serves file directly
# Fall back to streaming for non-nginx deployments
use_nginx_accel = os.environ.get("USE_NGINX_ACCEL", "").lower() == "true"
logger.info(f"[DOWNLOAD] File: {filename}, Size: {file_size}, USE_NGINX_ACCEL: {use_nginx_accel}")
if use_nginx_accel:
# X-Accel-Redirect: Django returns immediately, nginx serves file
logger.info(f"[DOWNLOAD] Using X-Accel-Redirect: /protected-backups/{filename}")
response = HttpResponse()
response["X-Accel-Redirect"] = f"/protected-backups/{filename}"
response["Content-Type"] = "application/zip"
response["Content-Length"] = file_size
response["Content-Disposition"] = f'attachment; filename="{filename}"'
return response
else:
# Streaming fallback for non-nginx deployments
logger.info(f"[DOWNLOAD] Using streaming fallback (no nginx)")
def file_iterator(file_path, chunk_size=2 * 1024 * 1024):
with open(file_path, "rb") as f:
while chunk := f.read(chunk_size):
yield chunk
response = StreamingHttpResponse(
file_iterator(backup_file),
content_type="application/zip",
)
response["Content-Length"] = file_size
response["Content-Disposition"] = f'attachment; filename="{filename}"'
return response
except Http404:
raise
except Exception as e:
return Response(
{"detail": f"Download failed: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["DELETE"])
@permission_classes([IsAdmin])
def delete_backup(request, filename):
"""Delete a backup file."""
try:
# Security: prevent path traversal
if ".." in filename or "/" in filename or "\\" in filename:
raise Http404("Invalid filename")
services.delete_backup(filename)
return Response(
{"detail": "Backup deleted successfully"},
status=status.HTTP_204_NO_CONTENT,
)
except FileNotFoundError:
raise Http404("Backup file not found")
except Exception as e:
return Response(
{"detail": f"Delete failed: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["POST"])
@permission_classes([IsAdmin])
@parser_classes([MultiPartParser, FormParser])
def upload_backup(request):
"""Upload a backup file for restoration."""
uploaded = request.FILES.get("file")
if not uploaded:
return Response(
{"detail": "No file uploaded"},
status=status.HTTP_400_BAD_REQUEST,
)
try:
backup_dir = services.get_backup_dir()
# Sanitize filename: strip directory components to prevent path traversal
filename = Path(uploaded.name or "uploaded-backup.zip").name
if not filename:
filename = "uploaded-backup.zip"
try:
safe_upload_path(filename, str(backup_dir))
except ValueError:
return Response({"detail": "Invalid filename."}, status=status.HTTP_400_BAD_REQUEST)
# Ensure unique filename
backup_file = (backup_dir / filename).resolve()
counter = 1
while backup_file.exists():
name_parts = filename.rsplit(".", 1)
if len(name_parts) == 2:
backup_file = backup_dir / f"{name_parts[0]}-{counter}.{name_parts[1]}"
else:
backup_file = backup_dir / f"{filename}-{counter}"
counter += 1
# Save uploaded file
with backup_file.open("wb") as f:
for chunk in uploaded.chunks():
f.write(chunk)
return Response(
{
"detail": "Backup uploaded successfully",
"filename": backup_file.name,
},
status=status.HTTP_201_CREATED,
)
except Exception as e:
return Response(
{"detail": f"Upload failed: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["POST"])
@permission_classes([IsAdmin])
def restore_backup(request, filename):
"""Restore from a backup file (async via Celery). WARNING: This will flush the database!"""
try:
# Security: prevent path traversal
if ".." in filename or "/" in filename or "\\" in filename:
raise Http404("Invalid filename")
backup_dir = services.get_backup_dir()
backup_file = backup_dir / filename
if not backup_file.exists():
raise Http404("Backup file not found")
task = restore_backup_task.delay(filename)
return Response(
{
"detail": "Restore started",
"task_id": task.id,
"task_token": _generate_task_token(task.id),
},
status=status.HTTP_202_ACCEPTED,
)
except Http404:
raise
except Exception as e:
return Response(
{"detail": f"Failed to start restore: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["GET"])
@permission_classes([IsAdmin])
def get_schedule(request):
"""Get backup schedule settings."""
try:
settings = get_schedule_settings()
return Response(settings)
except Exception as e:
return Response(
{"detail": f"Failed to get schedule: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@api_view(["PUT"])
@permission_classes([IsAdmin])
def update_schedule(request):
"""Update backup schedule settings."""
try:
settings = update_schedule_settings(request.data)
return Response(settings)
except ValueError as e:
return Response(
{"detail": str(e)},
status=status.HTTP_400_BAD_REQUEST,
)
except Exception as e:
return Response(
{"detail": f"Failed to update schedule: {str(e)}"},
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
+40
View File
@@ -0,0 +1,40 @@
import logging
from django.apps import AppConfig
logger = logging.getLogger(__name__)
class BackupsConfig(AppConfig):
default_auto_field = "django.db.models.BigAutoField"
name = "apps.backups"
verbose_name = "Backups"
def ready(self):
"""Initialize backup scheduler on app startup."""
from dispatcharr.app_initialization import should_skip_initialization
# Skip if this is a management command, worker process, or dev server
if should_skip_initialization():
return
logger.debug("Syncing backup scheduler on app startup")
self._sync_backup_scheduler()
def _sync_backup_scheduler(self):
"""Sync backup scheduler task to database."""
from core.models import CoreSettings
from .scheduler import _sync_periodic_task, DEFAULTS
try:
# Ensure settings exist with defaults if this is a new install
CoreSettings.objects.get_or_create(
key="backup_settings",
defaults={"name": "Backup Settings", "value": DEFAULTS.copy()}
)
# Always sync the periodic task (handles new installs, updates, or missing tasks)
logger.debug("Syncing backup scheduler")
_sync_periodic_task()
except Exception as e:
# Log but don't fail startup if there's an issue
logger.warning(f"Failed to initialize backup scheduler: {e}")
View File
View File
+133
View File
@@ -0,0 +1,133 @@
import json
import logging
from django_celery_beat.models import PeriodicTask
from core.models import CoreSettings
from core.scheduling import (
create_or_update_periodic_task,
delete_periodic_task,
)
logger = logging.getLogger(__name__)
BACKUP_SCHEDULE_TASK_NAME = "backup-scheduled-task"
DEFAULTS = {
"schedule_enabled": True,
"schedule_frequency": "daily",
"schedule_time": "03:00",
"schedule_day_of_week": 0, # Sunday
"retention_count": 3,
"schedule_cron_expression": "",
}
def _get_backup_settings():
"""Get all backup settings from CoreSettings grouped JSON."""
try:
settings_obj = CoreSettings.objects.get(key="backup_settings")
return settings_obj.value if isinstance(settings_obj.value, dict) else DEFAULTS.copy()
except CoreSettings.DoesNotExist:
return DEFAULTS.copy()
def _update_backup_settings(updates: dict) -> None:
"""Update backup settings in the grouped JSON."""
obj, created = CoreSettings.objects.get_or_create(
key="backup_settings",
defaults={"name": "Backup Settings", "value": DEFAULTS.copy()}
)
current = obj.value if isinstance(obj.value, dict) else {}
current.update(updates)
obj.value = current
obj.save()
def get_schedule_settings() -> dict:
"""Get all backup schedule settings."""
settings = _get_backup_settings()
return {
"enabled": bool(settings.get("schedule_enabled", DEFAULTS["schedule_enabled"])),
"frequency": str(settings.get("schedule_frequency", DEFAULTS["schedule_frequency"])),
"time": str(settings.get("schedule_time", DEFAULTS["schedule_time"])),
"day_of_week": int(settings.get("schedule_day_of_week", DEFAULTS["schedule_day_of_week"])),
"retention_count": int(settings.get("retention_count", DEFAULTS["retention_count"])),
"cron_expression": str(settings.get("schedule_cron_expression", DEFAULTS["schedule_cron_expression"])),
}
def update_schedule_settings(data: dict) -> dict:
"""Update backup schedule settings and sync the PeriodicTask."""
# Validate
if "frequency" in data and data["frequency"] not in ("daily", "weekly"):
raise ValueError("frequency must be 'daily' or 'weekly'")
if "time" in data:
try:
hour, minute = data["time"].split(":")
int(hour)
int(minute)
except (ValueError, AttributeError):
raise ValueError("time must be in HH:MM format")
if "day_of_week" in data:
day = int(data["day_of_week"])
if day < 0 or day > 6:
raise ValueError("day_of_week must be 0-6 (Sunday-Saturday)")
if "retention_count" in data:
count = int(data["retention_count"])
if count < 0:
raise ValueError("retention_count must be >= 0")
# Update settings with proper key names
updates = {}
if "enabled" in data:
updates["schedule_enabled"] = bool(data["enabled"])
if "frequency" in data:
updates["schedule_frequency"] = str(data["frequency"])
if "time" in data:
updates["schedule_time"] = str(data["time"])
if "day_of_week" in data:
updates["schedule_day_of_week"] = int(data["day_of_week"])
if "retention_count" in data:
updates["retention_count"] = int(data["retention_count"])
if "cron_expression" in data:
updates["schedule_cron_expression"] = str(data["cron_expression"])
_update_backup_settings(updates)
# Sync the periodic task
_sync_periodic_task()
return get_schedule_settings()
def _sync_periodic_task() -> None:
"""Create, update, or delete the scheduled backup task based on settings."""
settings = get_schedule_settings()
if not settings["enabled"]:
delete_periodic_task(BACKUP_SCHEDULE_TASK_NAME)
logger.info("Backup schedule disabled, removed periodic task")
return
# Check if using cron expression (advanced mode)
if settings["cron_expression"]:
cron_expr = settings["cron_expression"]
else:
# Build a cron expression from simple frequency settings
hour, minute = settings["time"].split(":")
if settings["frequency"] == "daily":
cron_expr = f"{minute} {hour} * * *"
else: # weekly
cron_expr = f"{minute} {hour} * * {settings['day_of_week']}"
create_or_update_periodic_task(
task_name=BACKUP_SCHEDULE_TASK_NAME,
celery_task_path="apps.backups.tasks.scheduled_backup_task",
kwargs={"retention_count": settings["retention_count"]},
cron_expression=cron_expr,
enabled=True,
)
+376
View File
@@ -0,0 +1,376 @@
import datetime
import json
import os
import shutil
import subprocess
import tempfile
from pathlib import Path
from zipfile import ZipFile, ZIP_DEFLATED
import logging
import pytz
from django.conf import settings
from core.models import CoreSettings
logger = logging.getLogger(__name__)
def get_backup_dir() -> Path:
"""Get the backup directory, creating it if necessary."""
backup_dir = Path(settings.BACKUP_ROOT)
backup_dir.mkdir(parents=True, exist_ok=True)
return backup_dir
def _is_postgresql() -> bool:
"""Check if we're using PostgreSQL."""
return settings.DATABASES["default"]["ENGINE"] == "django.db.backends.postgresql"
def _get_pg_env() -> dict:
"""Get environment variables for PostgreSQL commands.
Includes PGPASSWORD for password auth and PGSSL* variables for TLS.
Reads TLS config from DATABASES['default']['OPTIONS'], which is
populated by settings.py when POSTGRES_SSL=true.
"""
db_config = settings.DATABASES["default"]
env = os.environ.copy()
password = db_config.get("PASSWORD", "")
if password:
env["PGPASSWORD"] = password
else:
env.pop("PGPASSWORD", None)
# Propagate TLS configuration from Django OPTIONS to libpq env vars.
options = db_config.get("OPTIONS", {})
_ssl_env_map = {
"sslmode": "PGSSLMODE",
"sslrootcert": "PGSSLROOTCERT",
"sslcert": "PGSSLCERT",
"sslkey": "PGSSLKEY",
}
# Always strip inherited PGSSL* vars first, then set only what is explicitly configured
for opt_key, env_key in _ssl_env_map.items():
env.pop(env_key, None)
value = options.get(opt_key)
if value:
env[env_key] = value
return env
def _get_pg_args() -> list[str]:
"""Get common PostgreSQL command arguments."""
db_config = settings.DATABASES["default"]
return [
"-h", db_config.get("HOST", "localhost"),
"-p", str(db_config.get("PORT", 5432)),
"-U", db_config.get("USER", "postgres"),
"-d", db_config.get("NAME", "dispatcharr"),
]
def _dump_postgresql(output_file: Path) -> None:
"""Dump PostgreSQL database using pg_dump."""
logger.info("Dumping PostgreSQL database with pg_dump...")
cmd = [
"pg_dump",
*_get_pg_args(),
"-Fc", # Custom format for pg_restore
"-v", # Verbose
"-f", str(output_file),
]
result = subprocess.run(
cmd,
env=_get_pg_env(),
capture_output=True,
text=True,
)
if result.returncode != 0:
logger.error(f"pg_dump failed: {result.stderr}")
raise RuntimeError(f"pg_dump failed: {result.stderr}")
logger.debug(f"pg_dump output: {result.stderr}")
def _clean_postgresql_schema() -> None:
"""Drop and recreate the public schema to ensure a completely clean restore."""
logger.info("[PG_CLEAN] Dropping and recreating public schema...")
# Commands to drop and recreate schema
sql_commands = "DROP SCHEMA IF EXISTS public CASCADE; CREATE SCHEMA public; GRANT ALL ON SCHEMA public TO public;"
cmd = [
"psql",
*_get_pg_args(),
"-c", sql_commands,
]
result = subprocess.run(
cmd,
env=_get_pg_env(),
capture_output=True,
text=True,
)
if result.returncode != 0:
logger.error(f"[PG_CLEAN] Failed to clean schema: {result.stderr}")
raise RuntimeError(f"Failed to clean PostgreSQL schema: {result.stderr}")
logger.info("[PG_CLEAN] Schema cleaned successfully")
def _restore_postgresql(dump_file: Path) -> None:
"""Restore PostgreSQL database using pg_restore."""
logger.info("[PG_RESTORE] Starting pg_restore...")
logger.info(f"[PG_RESTORE] Dump file: {dump_file}")
# Drop and recreate schema to ensure a completely clean restore
_clean_postgresql_schema()
pg_args = _get_pg_args()
logger.info(f"[PG_RESTORE] Connection args: {pg_args}")
cmd = [
"pg_restore",
"--no-owner", # Skip ownership commands (we already created schema)
*pg_args,
"-v", # Verbose
str(dump_file),
]
logger.info(f"[PG_RESTORE] Running command: {' '.join(cmd)}")
result = subprocess.run(
cmd,
env=_get_pg_env(),
capture_output=True,
text=True,
)
logger.info(f"[PG_RESTORE] Return code: {result.returncode}")
# pg_restore may return non-zero even on partial success
# Check for actual errors vs warnings
if result.returncode != 0:
# Some errors during restore are expected (e.g., "does not exist" when cleaning)
# Only fail on critical errors
stderr = result.stderr.lower()
if "fatal" in stderr or "could not connect" in stderr:
logger.error(f"[PG_RESTORE] Failed critically: {result.stderr}")
raise RuntimeError(f"pg_restore failed: {result.stderr}")
else:
logger.warning(f"[PG_RESTORE] Completed with warnings: {result.stderr[:500]}...")
logger.info("[PG_RESTORE] Completed successfully")
def _dump_sqlite(output_file: Path) -> None:
"""Dump SQLite database using sqlite3 .backup command."""
logger.info("Dumping SQLite database with sqlite3 .backup...")
db_path = Path(settings.DATABASES["default"]["NAME"])
if not db_path.exists():
raise FileNotFoundError(f"SQLite database not found: {db_path}")
# Use sqlite3 .backup command via stdin for reliable execution
result = subprocess.run(
["sqlite3", str(db_path)],
input=f".backup '{output_file}'\n",
capture_output=True,
text=True,
)
if result.returncode != 0:
logger.error(f"sqlite3 backup failed: {result.stderr}")
raise RuntimeError(f"sqlite3 backup failed: {result.stderr}")
# Verify the backup file was created
if not output_file.exists():
raise RuntimeError("sqlite3 backup failed: output file not created")
logger.info(f"sqlite3 backup completed successfully: {output_file}")
def _restore_sqlite(dump_file: Path) -> None:
"""Restore SQLite database by replacing the database file."""
logger.info("Restoring SQLite database...")
db_path = Path(settings.DATABASES["default"]["NAME"])
backup_current = None
# Backup current database before overwriting
if db_path.exists():
backup_current = db_path.with_suffix(".db.bak")
shutil.copy2(db_path, backup_current)
logger.info(f"Backed up current database to {backup_current}")
# Ensure parent directory exists
db_path.parent.mkdir(parents=True, exist_ok=True)
# The backup file from _dump_sqlite is a complete SQLite database file
# We can simply copy it over the existing database
shutil.copy2(dump_file, db_path)
# Verify the restore worked by checking if sqlite3 can read it
result = subprocess.run(
["sqlite3", str(db_path)],
input=".tables\n",
capture_output=True,
text=True,
)
if result.returncode != 0:
logger.error(f"sqlite3 verification failed: {result.stderr}")
# Try to restore from backup
if backup_current and backup_current.exists():
shutil.copy2(backup_current, db_path)
logger.info("Restored original database from backup")
raise RuntimeError(f"sqlite3 restore verification failed: {result.stderr}")
logger.info("sqlite3 restore completed successfully")
def create_backup() -> Path:
"""
Create a backup archive containing database dump and data directories.
Returns the path to the created backup file.
"""
backup_dir = get_backup_dir()
# Use system timezone for filename (user-friendly), but keep internal timestamps as UTC
system_tz_name = CoreSettings.get_system_time_zone()
try:
system_tz = pytz.timezone(system_tz_name)
now_local = datetime.datetime.now(datetime.UTC).astimezone(system_tz)
timestamp = now_local.strftime("%Y.%m.%d.%H.%M.%S")
except Exception as e:
logger.warning(f"Failed to use system timezone {system_tz_name}: {e}, falling back to UTC")
timestamp = datetime.datetime.now(datetime.UTC).strftime("%Y.%m.%d.%H.%M.%S")
backup_name = f"dispatcharr-backup-{timestamp}.zip"
backup_file = backup_dir / backup_name
logger.info(f"Creating backup: {backup_name}")
with tempfile.TemporaryDirectory(prefix="dispatcharr-backup-") as temp_dir:
temp_path = Path(temp_dir)
# Determine database type and dump accordingly
if _is_postgresql():
db_dump_file = temp_path / "database.dump"
_dump_postgresql(db_dump_file)
db_type = "postgresql"
else:
db_dump_file = temp_path / "database.sqlite3"
_dump_sqlite(db_dump_file)
db_type = "sqlite"
# Create ZIP archive with compression and ZIP64 support for large files
with ZipFile(backup_file, "w", compression=ZIP_DEFLATED, allowZip64=True) as zip_file:
# Add database dump
zip_file.write(db_dump_file, db_dump_file.name)
# Add metadata
metadata = {
"format": "dispatcharr-backup",
"version": 2,
"database_type": db_type,
"database_file": db_dump_file.name,
"created_at": datetime.datetime.now(datetime.UTC).isoformat(),
}
zip_file.writestr("metadata.json", json.dumps(metadata, indent=2))
logger.info(f"Backup created successfully: {backup_file}")
return backup_file
def restore_backup(backup_file: Path) -> None:
"""
Restore from a backup archive.
WARNING: This will overwrite the database!
"""
if not backup_file.exists():
raise FileNotFoundError(f"Backup file not found: {backup_file}")
logger.info(f"Restoring from backup: {backup_file}")
with tempfile.TemporaryDirectory(prefix="dispatcharr-restore-") as temp_dir:
temp_path = Path(temp_dir)
# Extract backup
logger.debug("Extracting backup archive...")
with ZipFile(backup_file, "r") as zip_file:
zip_file.extractall(temp_path)
# Read metadata
metadata_file = temp_path / "metadata.json"
if not metadata_file.exists():
raise ValueError("Invalid backup: missing metadata.json")
with open(metadata_file) as f:
metadata = json.load(f)
# Restore database
_restore_database(temp_path, metadata)
logger.info("Restore completed successfully")
def _restore_database(temp_path: Path, metadata: dict) -> None:
"""Restore database from backup."""
db_type = metadata.get("database_type", "postgresql")
db_file = metadata.get("database_file", "database.dump")
dump_file = temp_path / db_file
if not dump_file.exists():
raise ValueError(f"Invalid backup: missing {db_file}")
current_db_type = "postgresql" if _is_postgresql() else "sqlite"
if db_type != current_db_type:
raise ValueError(
f"Database type mismatch: backup is {db_type}, "
f"but current database is {current_db_type}"
)
if db_type == "postgresql":
_restore_postgresql(dump_file)
else:
_restore_sqlite(dump_file)
def list_backups() -> list[dict]:
"""List all available backup files with metadata."""
backup_dir = get_backup_dir()
backups = []
for backup_file in sorted(backup_dir.glob("dispatcharr-backup-*.zip"), reverse=True):
# Use UTC timezone so frontend can convert to user's local time
created_time = datetime.datetime.fromtimestamp(backup_file.stat().st_mtime, datetime.UTC)
backups.append({
"name": backup_file.name,
"size": backup_file.stat().st_size,
"created": created_time.isoformat(),
})
return backups
def delete_backup(filename: str) -> None:
"""Delete a backup file."""
backup_dir = get_backup_dir()
backup_file = backup_dir / filename
if not backup_file.exists():
raise FileNotFoundError(f"Backup file not found: {filename}")
if not backup_file.is_file():
raise ValueError(f"Invalid backup file: {filename}")
backup_file.unlink()
logger.info(f"Deleted backup: {filename}")
+106
View File
@@ -0,0 +1,106 @@
import logging
import traceback
from celery import shared_task
from . import services
logger = logging.getLogger(__name__)
def _cleanup_old_backups(retention_count: int) -> int:
"""Delete old backups, keeping only the most recent N. Returns count deleted."""
if retention_count <= 0:
return 0
backups = services.list_backups()
if len(backups) <= retention_count:
return 0
# Backups are sorted newest first, so delete from the end
to_delete = backups[retention_count:]
deleted = 0
for backup in to_delete:
try:
services.delete_backup(backup["name"])
deleted += 1
logger.info(f"[CLEANUP] Deleted old backup: {backup['name']}")
except Exception as e:
logger.error(f"[CLEANUP] Failed to delete {backup['name']}: {e}")
return deleted
@shared_task(bind=True)
def create_backup_task(self):
"""Celery task to create a backup asynchronously."""
try:
logger.info(f"[BACKUP] Starting backup task {self.request.id}")
backup_file = services.create_backup()
logger.info(f"[BACKUP] Task {self.request.id} completed: {backup_file.name}")
return {
"status": "completed",
"filename": backup_file.name,
"size": backup_file.stat().st_size,
}
except Exception as e:
logger.error(f"[BACKUP] Task {self.request.id} failed: {str(e)}")
logger.error(f"[BACKUP] Traceback: {traceback.format_exc()}")
return {
"status": "failed",
"error": str(e),
}
@shared_task(bind=True)
def restore_backup_task(self, filename: str):
"""Celery task to restore a backup asynchronously."""
try:
logger.info(f"[RESTORE] Starting restore task {self.request.id} for {filename}")
backup_dir = services.get_backup_dir()
backup_file = backup_dir / filename
logger.info(f"[RESTORE] Backup file path: {backup_file}")
services.restore_backup(backup_file)
logger.info(f"[RESTORE] Task {self.request.id} completed successfully")
return {
"status": "completed",
"filename": filename,
}
except Exception as e:
logger.error(f"[RESTORE] Task {self.request.id} failed: {str(e)}")
logger.error(f"[RESTORE] Traceback: {traceback.format_exc()}")
return {
"status": "failed",
"error": str(e),
}
@shared_task(bind=True)
def scheduled_backup_task(self, retention_count: int = 0):
"""Celery task for scheduled backups with optional retention cleanup."""
try:
logger.info(f"[SCHEDULED] Starting scheduled backup task {self.request.id}")
# Create backup
backup_file = services.create_backup()
logger.info(f"[SCHEDULED] Backup created: {backup_file.name}")
# Cleanup old backups if retention is set
deleted = 0
if retention_count > 0:
deleted = _cleanup_old_backups(retention_count)
logger.info(f"[SCHEDULED] Cleanup complete, deleted {deleted} old backup(s)")
return {
"status": "completed",
"filename": backup_file.name,
"size": backup_file.stat().st_size,
"deleted_count": deleted,
}
except Exception as e:
logger.error(f"[SCHEDULED] Task {self.request.id} failed: {str(e)}")
logger.error(f"[SCHEDULED] Traceback: {traceback.format_exc()}")
return {
"status": "failed",
"error": str(e),
}
File diff suppressed because it is too large Load Diff
View File
+38
View File
@@ -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'
+55
View File
@@ -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
+11
View File
@@ -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
+53
View File
@@ -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',
]
+77
View File
@@ -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),
),
]
@@ -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'),
),
]
@@ -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),
),
]
+988
View File
@@ -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 []
+958
View File
@@ -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 []
+537
View File
@@ -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)
+236
View File
@@ -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)
+3973
View File
File diff suppressed because it is too large Load Diff
View File
+211
View File
@@ -0,0 +1,211 @@
from django.test import TestCase
from django.contrib.auth import get_user_model
from rest_framework.test import APIClient
from rest_framework import status
from apps.channels.models import Channel, ChannelGroup
User = get_user_model()
class ChannelBulkEditAPITests(TestCase):
def setUp(self):
# Create a test admin user (user_level >= 10) and authenticate
self.user = User.objects.create_user(username="testuser", password="testpass123")
self.user.user_level = 10 # Set admin level
self.user.save()
self.client = APIClient()
self.client.force_authenticate(user=self.user)
self.bulk_edit_url = "/api/channels/channels/edit/bulk/"
# Create test channel group
self.group1 = ChannelGroup.objects.create(name="Test Group 1")
self.group2 = ChannelGroup.objects.create(name="Test Group 2")
# Create test channels
self.channel1 = Channel.objects.create(
channel_number=1.0,
name="Channel 1",
tvg_id="channel1",
channel_group=self.group1
)
self.channel2 = Channel.objects.create(
channel_number=2.0,
name="Channel 2",
tvg_id="channel2",
channel_group=self.group1
)
self.channel3 = Channel.objects.create(
channel_number=3.0,
name="Channel 3",
tvg_id="channel3"
)
def test_bulk_edit_success(self):
"""Test successful bulk update of multiple channels"""
data = [
{"id": self.channel1.id, "name": "Updated Channel 1"},
{"id": self.channel2.id, "name": "Updated Channel 2", "channel_number": 22.0},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
self.assertEqual(len(response.data["channels"]), 2)
# Verify database changes
self.channel1.refresh_from_db()
self.channel2.refresh_from_db()
self.assertEqual(self.channel1.name, "Updated Channel 1")
self.assertEqual(self.channel2.name, "Updated Channel 2")
self.assertEqual(self.channel2.channel_number, 22.0)
def test_bulk_edit_with_empty_validated_data_first(self):
"""
Test the bug fix: when first channel has empty validated_data.
This was causing: ValueError: Field names must be given to bulk_update()
"""
# Create a channel with data that will be "unchanged" (empty validated_data)
# We'll send the same data it already has
data = [
# First channel: no actual changes (this would create empty validated_data)
{"id": self.channel1.id},
# Second channel: has changes
{"id": self.channel2.id, "name": "Updated Channel 2"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should not crash with ValueError
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
# Verify the channel with changes was updated
self.channel2.refresh_from_db()
self.assertEqual(self.channel2.name, "Updated Channel 2")
def test_bulk_edit_all_empty_updates(self):
"""Test when all channels have empty updates (no actual changes)"""
data = [
{"id": self.channel1.id},
{"id": self.channel2.id},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should succeed without calling bulk_update
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
def test_bulk_edit_mixed_fields(self):
"""Test bulk update where different channels update different fields"""
data = [
{"id": self.channel1.id, "name": "New Name 1"},
{"id": self.channel2.id, "channel_number": 99.0},
{"id": self.channel3.id, "tvg_id": "new_tvg_id", "name": "New Name 3"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 3 channels")
# Verify all updates
self.channel1.refresh_from_db()
self.channel2.refresh_from_db()
self.channel3.refresh_from_db()
self.assertEqual(self.channel1.name, "New Name 1")
self.assertEqual(self.channel2.channel_number, 99.0)
self.assertEqual(self.channel3.tvg_id, "new_tvg_id")
self.assertEqual(self.channel3.name, "New Name 3")
def test_bulk_edit_with_channel_group(self):
"""Test bulk update with channel_group_id changes"""
data = [
{"id": self.channel1.id, "channel_group_id": self.group2.id},
{"id": self.channel3.id, "channel_group_id": self.group1.id},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
# Verify group changes
self.channel1.refresh_from_db()
self.channel3.refresh_from_db()
self.assertEqual(self.channel1.channel_group, self.group2)
self.assertEqual(self.channel3.channel_group, self.group1)
def test_bulk_edit_nonexistent_channel(self):
"""Test bulk update with a channel that doesn't exist"""
nonexistent_id = 99999
data = [
{"id": nonexistent_id, "name": "Should Fail"},
{"id": self.channel1.id, "name": "Should Still Update"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should return 400 with errors
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn("errors", response.data)
self.assertEqual(len(response.data["errors"]), 1)
self.assertEqual(response.data["errors"][0]["channel_id"], nonexistent_id)
self.assertEqual(response.data["errors"][0]["error"], "Channel not found")
# The valid channel should still be updated
self.assertEqual(response.data["updated_count"], 1)
def test_bulk_edit_validation_error(self):
"""Test bulk update with invalid data (validation error)"""
data = [
{"id": self.channel1.id, "channel_number": "invalid_number"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Should return 400 with validation errors
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn("errors", response.data)
self.assertEqual(len(response.data["errors"]), 1)
self.assertIn("channel_number", response.data["errors"][0]["errors"])
def test_bulk_edit_empty_channel_updates(self):
"""Test bulk update with empty list"""
data = []
response = self.client.patch(self.bulk_edit_url, data, format="json")
# Empty list is accepted and returns success with 0 updates
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["message"], "Successfully updated 0 channels")
def test_bulk_edit_missing_channel_updates(self):
"""Test bulk update without proper format (dict instead of list)"""
data = {"channel_updates": {}}
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertEqual(response.data["error"], "Expected a list of channel updates")
def test_bulk_edit_preserves_other_fields(self):
"""Test that bulk update only changes specified fields"""
original_channel_number = self.channel1.channel_number
original_tvg_id = self.channel1.tvg_id
data = [
{"id": self.channel1.id, "name": "Only Name Changed"},
]
response = self.client.patch(self.bulk_edit_url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
# Verify only name changed, other fields preserved
self.channel1.refresh_from_db()
self.assertEqual(self.channel1.name, "Only Name Changed")
self.assertEqual(self.channel1.channel_number, original_channel_number)
self.assertEqual(self.channel1.tvg_id, original_tvg_id)
+278
View File
@@ -0,0 +1,278 @@
"""Tests for DVR retry logic.
Covers:
- _db_retry(): exponential backoff, max retries, connection reset
- Final metadata save retry in run_recording post-processing
- Initial TS proxy connection retry (per-base retry on retriable errors)
- recover_recordings_on_startup DB retry wrappers
"""
from datetime import timedelta
from unittest.mock import MagicMock, patch, call
from django.db import OperationalError
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel, Recording
from apps.channels.tasks import _db_retry
# ---------------------------------------------------------------------------
# _db_retry unit tests
# ---------------------------------------------------------------------------
class DbRetryTests(TestCase):
"""Tests for the _db_retry() exponential backoff helper."""
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_succeeds_on_first_attempt(self, _close, _sleep):
"""No retry needed when fn succeeds immediately."""
result = _db_retry(lambda: "ok", max_retries=3)
self.assertEqual(result, "ok")
_sleep.assert_not_called()
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_retries_on_operational_error_then_succeeds(self, mock_close, mock_sleep):
"""Retry succeeds on second attempt after OperationalError."""
call_count = {"n": 0}
def flaky():
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("connection reset")
return "recovered"
result = _db_retry(flaky, max_retries=3, base_interval=1)
self.assertEqual(result, "recovered")
self.assertEqual(call_count["n"], 2)
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_raises_after_max_retries_exhausted(self, mock_close, mock_sleep):
"""Raises OperationalError after all retries fail."""
def always_fail():
raise OperationalError("db gone")
with self.assertRaises(OperationalError):
_db_retry(always_fail, max_retries=3, base_interval=1)
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_exponential_backoff_timing(self, mock_close, mock_sleep):
"""Sleep durations follow exponential backoff: 1s, 2s, 4s."""
call_count = {"n": 0}
def fail_twice():
call_count["n"] += 1
if call_count["n"] <= 2:
raise OperationalError("retry me")
return "done"
_db_retry(fail_twice, max_retries=3, base_interval=1)
mock_sleep.assert_has_calls([call(1), call(2)])
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_close_old_connections_called_between_retries(self, mock_close, mock_sleep):
"""Stale DB connections are reset before each retry attempt."""
call_count = {"n": 0}
def fail_once():
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("stale conn")
return "ok"
_db_retry(fail_once, max_retries=3)
mock_close.assert_called_once()
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_non_operational_error_not_retried(self, mock_close, mock_sleep):
"""Non-OperationalError exceptions propagate immediately."""
def raise_value_error():
raise ValueError("not a DB error")
with self.assertRaises(ValueError):
_db_retry(raise_value_error, max_retries=3)
mock_sleep.assert_not_called()
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_returns_fn_return_value(self, mock_close, mock_sleep):
"""Return value of fn() is passed through."""
result = _db_retry(lambda: {"key": "value"}, max_retries=3)
self.assertEqual(result, {"key": "value"})
@patch("apps.channels.tasks.time.sleep")
@patch("apps.channels.tasks.close_old_connections")
def test_single_retry_allowed(self, mock_close, mock_sleep):
"""max_retries=1 means no retry — fail immediately."""
with self.assertRaises(OperationalError):
_db_retry(
lambda: (_ for _ in ()).throw(OperationalError("fail")),
max_retries=1,
)
mock_sleep.assert_not_called()
# ---------------------------------------------------------------------------
# Final metadata save retry integration tests
# ---------------------------------------------------------------------------
class FinalMetadataSaveRetryTests(TestCase):
"""The final recording metadata save must retry on transient DB errors."""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=95, name="Retry Test Channel"
)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_metadata_save_uses_db_retry(self, _ws):
"""Verify recording metadata is saved via _db_retry (retries on OperationalError)."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
# Directly call _db_retry to save metadata as run_recording does
cp = rec.custom_properties.copy()
cp["status"] = "completed"
cp["ended_at"] = str(now)
cp["bytes_written"] = 1024
def _save():
rec.custom_properties = cp
rec.save(update_fields=["custom_properties"])
_db_retry(_save, max_retries=3, base_interval=1, label="test save")
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], "completed")
self.assertEqual(rec.custom_properties["bytes_written"], 1024)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_metadata_survives_transient_save_failure(self, _ws):
"""Simulate OperationalError on first save, success on retry."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
cp = {"status": "completed", "bytes_written": 2048}
call_count = {"n": 0}
_real_save = rec.save
def patched_save(**kwargs):
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("connection reset by peer")
return _real_save(**kwargs)
with patch.object(rec, "save", side_effect=patched_save):
with patch("apps.channels.tasks.time.sleep"):
with patch("apps.channels.tasks.close_old_connections"):
def _save():
rec.custom_properties = cp
rec.save(update_fields=["custom_properties"])
_db_retry(_save, max_retries=3, base_interval=1, label="test")
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], "completed")
# ---------------------------------------------------------------------------
# Initial connection retry tests
# ---------------------------------------------------------------------------
class InitialConnectionRetryTests(TestCase):
"""Verify that the DVR task's reconnection logic retries the same
base URL before falling back to the next candidate."""
def test_reconnect_max_constant_exists_in_run_recording(self):
"""run_recording must define a max-reconnect limit to prevent
infinite retries on the same broken base URL."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
# The reconnection counter pattern must be present
self.assertIn("reconnect", source.lower(),
"run_recording must contain reconnection logic")
# ---------------------------------------------------------------------------
# recover_recordings_on_startup retry tests
# ---------------------------------------------------------------------------
class RecoveryRetryTests(TestCase):
"""DB operations in recover_recordings_on_startup must use _db_retry."""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=97, name="Recovery Retry Channel"
)
@patch("apps.channels.tasks.run_recording.apply_async")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_recovery_save_retries_on_operational_error(self, _ws, mock_async):
"""Recovery status update uses _db_retry — survives one OperationalError."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=30),
end_time=now + timedelta(minutes=30),
custom_properties={},
)
# Simulate what recovery does: mark interrupted, then save with retry
cp = rec.custom_properties or {}
cp["status"] = "interrupted"
cp["interrupted_reason"] = "server_restarted"
rec.custom_properties = cp
call_count = {"n": 0}
_real_save = Recording.save
def patched_save(self_rec, **kwargs):
call_count["n"] += 1
if call_count["n"] == 1:
raise OperationalError("db temporarily unavailable")
return _real_save(self_rec, **kwargs)
with patch.object(Recording, "save", patched_save):
with patch("apps.channels.tasks.time.sleep"):
with patch("apps.channels.tasks.close_old_connections"):
_db_retry(
lambda: rec.save(update_fields=["custom_properties"]),
max_retries=3,
label="test recovery",
)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("status"), "interrupted")
self.assertEqual(rec.custom_properties.get("interrupted_reason"), "server_restarted")
def test_db_retry_fetches_recording_list(self):
"""_db_retry correctly returns query results for recording list fetch."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=30),
end_time=now + timedelta(minutes=30),
custom_properties={},
)
result = _db_retry(
lambda: list(Recording.objects.filter(
start_time__lte=now, end_time__gt=now
)),
label="test query",
)
self.assertGreaterEqual(len(result), 1)
ids = [r.id for r in result]
self.assertIn(rec.id, ids)
@@ -0,0 +1,59 @@
import os
from django.test import SimpleTestCase
from unittest.mock import patch
from apps.channels.tasks import build_dvr_candidates
class DVRPortResolutionTests(SimpleTestCase):
"""
Tests that DVR recording candidate URLs respect the DISPATCHARR_PORT
environment variable instead of hardcoding port 9191.
"""
@patch.dict(os.environ, {'REDIS_HOST': 'redis'}, clear=True)
def test_default_port_uses_9191(self):
"""Without DISPATCHARR_PORT set, candidates default to 9191."""
candidates = build_dvr_candidates()
self.assertIn('http://web:9191', candidates)
self.assertIn('http://localhost:9191', candidates)
@patch.dict(os.environ, {'DISPATCHARR_PORT': '8080', 'REDIS_HOST': 'redis'}, clear=True)
def test_custom_port_reflected_in_candidates(self):
"""DISPATCHARR_PORT=8080 replaces all hardcoded 9191 references."""
candidates = build_dvr_candidates()
self.assertIn('http://web:8080', candidates)
self.assertIn('http://localhost:8080', candidates)
self.assertNotIn('http://web:9191', candidates)
self.assertNotIn('http://localhost:9191', candidates)
@patch.dict(os.environ, {
'DISPATCHARR_PORT': '7777',
'DISPATCHARR_ENV': 'dev',
'REDIS_HOST': 'redis',
}, clear=True)
def test_dev_mode_includes_5656_and_custom_port(self):
"""Dev mode includes both uwsgi internal port (5656) and custom port."""
candidates = build_dvr_candidates()
self.assertIn('http://127.0.0.1:5656', candidates)
self.assertIn('http://127.0.0.1:7777', candidates)
@patch.dict(os.environ, {
'DISPATCHARR_INTERNAL_TS_BASE_URL': 'http://custom:1234',
'REDIS_HOST': 'redis',
}, clear=True)
def test_explicit_override_is_first(self):
"""DISPATCHARR_INTERNAL_TS_BASE_URL should be the first candidate."""
candidates = build_dvr_candidates()
self.assertEqual(candidates[0], 'http://custom:1234')
@patch.dict(os.environ, {
'DISPATCHARR_PORT': '3000',
'DISPATCHARR_INTERNAL_API_BASE': 'http://myhost:4000',
'REDIS_HOST': 'redis',
}, clear=True)
def test_internal_api_base_overrides_web_fallback(self):
"""DISPATCHARR_INTERNAL_API_BASE replaces the http://web:{port} default."""
candidates = build_dvr_candidates()
self.assertIn('http://myhost:4000', candidates)
self.assertNotIn('http://web:3000', candidates)
+180
View File
@@ -0,0 +1,180 @@
"""Tests for the _match_epg_program_by_timeslot() helper in tasks.py.
Covers:
- Exact time-slot match returns program dict
- 80% overlap threshold: at boundary, above, and below
- Multiple overlapping programs: dominant vs. evenly split
- Edge cases: None inputs, zero-duration recording, no EPG data
- Returned dict structure (id, title, sub_title, description)
"""
from datetime import timedelta
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel
from apps.epg.models import EPGSource, EPGData, ProgramData
from apps.channels.tasks import _match_epg_program_by_timeslot
class EpgMatchingSetupMixin:
"""Shared setup for EPG matching tests."""
def setUp(self):
self.source = EPGSource.objects.create(name="Test Source")
self.epg = EPGData.objects.create(
tvg_id="test.channel", name="Test Channel EPG", epg_source=self.source,
)
self.channel = Channel.objects.create(
channel_number=50, name="EPG Match Channel", epg_data=self.epg,
)
self.base = timezone.now().replace(second=0, microsecond=0)
def _prog(self, offset_min, duration_min, title="Test Show", **kwargs):
"""Create a ProgramData starting offset_min from self.base."""
start = self.base + timedelta(minutes=offset_min)
end = start + timedelta(minutes=duration_min)
return ProgramData.objects.create(
epg=self.epg, start_time=start, end_time=end, title=title, **kwargs,
)
class ExactMatchTests(EpgMatchingSetupMixin, TestCase):
"""Recording window exactly matches an EPG program."""
def test_exact_match_returns_program_dict(self):
prog = self._prog(0, 60, title="News at 9", sub_title="Top Stories",
description="Evening news broadcast")
result = _match_epg_program_by_timeslot(
self.epg, prog.start_time, prog.end_time,
)
self.assertIsNotNone(result)
self.assertEqual(result["id"], prog.id)
self.assertEqual(result["title"], "News at 9")
self.assertEqual(result["sub_title"], "Top Stories")
self.assertEqual(result["description"], "Evening news broadcast")
def test_missing_optional_fields_returned_as_empty_strings(self):
prog = self._prog(0, 30, title="Minimal Show")
result = _match_epg_program_by_timeslot(
self.epg, prog.start_time, prog.end_time,
)
self.assertIsNotNone(result)
self.assertEqual(result["sub_title"], "")
self.assertEqual(result["description"], "")
class OverlapThresholdTests(EpgMatchingSetupMixin, TestCase):
"""80% overlap threshold boundary tests."""
def test_exactly_80_percent_overlap_returns_match(self):
"""Program covers exactly 80% of the recording window."""
# Program: 0-60min, Recording: 0-75min → overlap = 60/75 = 80%
prog = self._prog(0, 60, title="Borderline Show")
rec_start = self.base
rec_end = self.base + timedelta(minutes=75)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Borderline Show")
def test_below_80_percent_returns_none(self):
"""Program covers 79% of the recording — below threshold."""
# Program: 0-60min, Recording: 0-76min → overlap = 60/76 ≈ 78.9%
prog = self._prog(0, 60, title="Too Short")
rec_start = self.base
rec_end = self.base + timedelta(minutes=76)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNone(result)
def test_above_80_percent_returns_match(self):
"""Program covers 90% of the recording."""
# Program: 0-60min, Recording: 0-66min → overlap = 60/66 ≈ 90.9%
prog = self._prog(0, 60, title="Good Match")
rec_start = self.base
rec_end = self.base + timedelta(minutes=66)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Good Match")
class MultipleProgramTests(EpgMatchingSetupMixin, TestCase):
"""Recording spans multiple EPG programs."""
def test_dominant_program_returned(self):
"""Recording spans 2 programs; one covers 85%, the other 15%."""
# Show A: 0-60min, Show B: 60-120min
# Recording: 9-69min → A overlap=51/60=85%, B overlap=9/60=15%
self._prog(0, 60, title="Show A")
self._prog(60, 60, title="Show B")
rec_start = self.base + timedelta(minutes=9)
rec_end = self.base + timedelta(minutes=69)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Show A")
def test_evenly_split_returns_none(self):
"""Recording spans 2 equal programs — neither reaches 80%."""
# Show A: 0-60min, Show B: 60-120min
# Recording: 30-90min → each covers 50%
self._prog(0, 60, title="Show A")
self._prog(60, 60, title="Show B")
rec_start = self.base + timedelta(minutes=30)
rec_end = self.base + timedelta(minutes=90)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNone(result)
def test_three_programs_one_dominant(self):
"""Recording spans 3 programs; middle one is dominant."""
# A: 0-30min, B: 30-90min, C: 90-120min
# Recording: 25-95min (70min window) → B overlap=60/70≈85.7%
self._prog(0, 30, title="Show A")
self._prog(30, 60, title="Show B")
self._prog(90, 30, title="Show C")
rec_start = self.base + timedelta(minutes=25)
rec_end = self.base + timedelta(minutes=95)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNotNone(result)
self.assertEqual(result["title"], "Show B")
class EdgeCaseTests(EpgMatchingSetupMixin, TestCase):
"""Edge cases and error handling."""
def test_none_epg_data_returns_none(self):
result = _match_epg_program_by_timeslot(None, self.base, self.base + timedelta(hours=1))
self.assertIsNone(result)
def test_none_start_time_returns_none(self):
result = _match_epg_program_by_timeslot(self.epg, None, self.base + timedelta(hours=1))
self.assertIsNone(result)
def test_none_end_time_returns_none(self):
result = _match_epg_program_by_timeslot(self.epg, self.base, None)
self.assertIsNone(result)
def test_zero_duration_returns_none(self):
"""Recording with start == end should return None."""
result = _match_epg_program_by_timeslot(self.epg, self.base, self.base)
self.assertIsNone(result)
def test_negative_duration_returns_none(self):
"""Recording with end before start should return None."""
result = _match_epg_program_by_timeslot(
self.epg, self.base + timedelta(hours=1), self.base,
)
self.assertIsNone(result)
def test_no_overlapping_programs_returns_none(self):
"""No EPG programs in the recording window."""
self._prog(0, 60, title="Earlier Show")
rec_start = self.base + timedelta(hours=5)
rec_end = rec_start + timedelta(hours=1)
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
self.assertIsNone(result)
def test_empty_epg_no_programs_returns_none(self):
"""EPGData exists but has no programs."""
result = _match_epg_program_by_timeslot(
self.epg, self.base, self.base + timedelta(hours=1),
)
self.assertIsNone(result)
@@ -0,0 +1,235 @@
"""Tests for the Extend In-Progress Recording feature.
Covers:
- extend() API endpoint (happy path and validation)
- pre_save signal guard: end_time change must NOT revoke a live recording
- pre_save signal guard: end_time change MUST still revoke an upcoming recording
- TOCTOU edge cases (extend on a completed/stopped/nonexistent recording)
"""
from datetime import timedelta
from unittest.mock import MagicMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording
from apps.channels.api_views import RecordingViewSet
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="extend_test_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
# ---------------------------------------------------------------------------
# Extend endpoint tests
# ---------------------------------------------------------------------------
class ExtendEndpointTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/extend/"""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=88, name="Extend Test Channel"
)
self.user = _make_admin()
self.factory = APIRequestFactory()
def _extend(self, rec, extra_minutes):
request = self.factory.post(
f"/api/channels/recordings/{rec.id}/extend/",
{"extra_minutes": extra_minutes},
format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "extend"})
return view(request, pk=rec.id)
def _make_rec(self, status="recording"):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": status},
)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_updates_end_time_in_db(self, _ws):
rec = self._make_rec()
original_end = rec.end_time
response = self._extend(rec, 30)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data.get("success"))
rec.refresh_from_db()
expected = original_end + timedelta(minutes=30)
delta = abs((rec.end_time - expected).total_seconds())
self.assertLess(delta, 1, "end_time was not extended by the correct amount")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_stacks_multiple_extensions(self, _ws):
"""Calling extend() twice adds both increments."""
rec = self._make_rec()
original_end = rec.end_time
self._extend(rec, 15)
self._extend(rec, 30)
rec.refresh_from_db()
expected = original_end + timedelta(minutes=45)
delta = abs((rec.end_time - expected).total_seconds())
self.assertLess(delta, 1)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_does_not_clear_task_id(self, _ws):
"""The running Celery task must survive the DB save."""
rec = self._make_rec()
rec.task_id = "dvr-recording-999"
rec.save(update_fields=["task_id"])
self._extend(rec, 30)
rec.refresh_from_db()
self.assertEqual(rec.task_id, "dvr-recording-999")
def test_extend_returns_400_if_finished(self):
"""Cannot extend a completed, stopped, or interrupted recording."""
for bad_status in ("completed", "stopped", "interrupted"):
with self.subTest(status=bad_status):
rec = self._make_rec(status=bad_status)
response = self._extend(rec, 30)
self.assertEqual(response.status_code, 400)
self.assertFalse(response.data.get("success"))
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_succeeds_before_task_sets_status(self, _ws):
"""Extend must work when status is empty (task hasn't started yet)."""
rec = self._make_rec(status="")
response = self._extend(rec, 15)
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
expected = rec.end_time # already extended
self.assertTrue(response.data.get("success"))
@patch("apps.channels.signals.revoke_task")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_extend_bypasses_signals_no_revoke(self, _ws, mock_revoke):
"""Extend uses .update() to bypass pre_save — revoke_task must never fire."""
rec = self._make_rec(status="")
rec.task_id = "dvr-recording-500"
rec.save(update_fields=["task_id"])
self._extend(rec, 15)
self._extend(rec, 30)
mock_revoke.assert_not_called()
rec.refresh_from_db()
self.assertEqual(rec.task_id, "dvr-recording-500")
def test_extend_returns_400_for_zero_minutes(self):
response = self._extend(self._make_rec(), 0)
self.assertEqual(response.status_code, 400)
def test_extend_returns_400_for_negative_minutes(self):
response = self._extend(self._make_rec(), -15)
self.assertEqual(response.status_code, 400)
def test_extend_returns_400_for_non_numeric_minutes(self):
rec = self._make_rec()
request = self.factory.post(
f"/api/channels/recordings/{rec.id}/extend/",
{"extra_minutes": "lots"},
format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "extend"})
response = view(request, pk=rec.id)
self.assertEqual(response.status_code, 400)
def test_extend_returns_404_for_nonexistent_recording(self):
request = self.factory.post(
"/api/channels/recordings/999999/extend/",
{"extra_minutes": 30},
format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "extend"})
response = view(request, pk=999999)
self.assertEqual(response.status_code, 404)
# ---------------------------------------------------------------------------
# pre_save signal guard tests
# ---------------------------------------------------------------------------
class PreSaveExtendGuardTests(TestCase):
"""The pre_save signal must NOT revoke a live recording when end_time changes,
but MUST still revoke a scheduled (upcoming) recording as before."""
def setUp(self):
self.channel = Channel.objects.create(
channel_number=77, name="Signal Guard Channel"
)
def _make_rec(self, status="", task_id="dvr-recording-42"):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now + timedelta(hours=1),
end_time=now + timedelta(hours=2),
task_id=task_id,
custom_properties={"status": status} if status else {},
)
@patch("apps.channels.signals.revoke_task")
def test_end_time_change_does_not_revoke_live_recording(self, mock_revoke):
"""When status='recording', extending end_time must not call revoke_task."""
rec = self._make_rec(status="recording", task_id="dvr-recording-42")
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
mock_revoke.assert_not_called()
@patch("apps.channels.signals.revoke_task")
def test_task_id_preserved_after_extend_on_live_recording(self, mock_revoke):
"""task_id must not be cleared for a live recording's end_time change."""
rec = self._make_rec(status="recording", task_id="dvr-recording-42")
original_task_id = rec.task_id
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
rec.refresh_from_db()
self.assertEqual(rec.task_id, original_task_id)
@patch("apps.channels.signals.revoke_task")
def test_end_time_change_still_revokes_upcoming_recording(self, mock_revoke):
"""The guard must NOT apply to upcoming recordings — existing behavior preserved."""
rec = self._make_rec(status="", task_id="dvr-recording-77")
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
mock_revoke.assert_called_once_with("dvr-recording-77")
@patch("apps.channels.signals.revoke_task")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_pre_save_guard_reads_db_status_not_memory_status(self, _ws, mock_revoke):
"""pre_save reads status from DB (old object), not from the instance being saved."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
task_id="dvr-recording-66",
custom_properties={"status": "recording"},
)
# Simulate: DB status changes to 'completed' behind the instance's back
Recording.objects.filter(pk=rec.pk).update(
custom_properties={"status": "completed"}
)
rec.end_time = rec.end_time + timedelta(minutes=30)
rec.save(update_fields=["end_time"])
# revoke_task should be called because DB status is "completed", not "recording"
mock_revoke.assert_called_once_with("dvr-recording-66")
@@ -0,0 +1,379 @@
"""Tests for recording metadata endpoints and logo proxy negative cache.
Covers:
- update_metadata endpoint: title/description, user_edited flag, validation
- refresh_artwork endpoint: returns immediately, background thread behavior
- Logo proxy negative cache: cache hit/miss, expiry, eviction, success clears
"""
import time as time_mod
from datetime import timedelta
from unittest.mock import MagicMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording, Logo
from apps.channels.api_views import RecordingViewSet, LogoViewSet
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="metadata_test_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
# ---------------------------------------------------------------------------
# update_metadata endpoint
# ---------------------------------------------------------------------------
class UpdateMetadataTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/update-metadata/"""
def setUp(self):
self.channel = Channel.objects.create(channel_number=70, name="Meta Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _update(self, rec, data):
request = self.factory.post(
f"/api/channels/recordings/{rec.id}/update-metadata/",
data, format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "update_metadata"})
return view(request, pk=rec.id)
def _make_rec(self, custom_properties=None):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties=custom_properties or {},
)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_update_title_only(self, _ws):
rec = self._make_rec()
response = self._update(rec, {"title": "My Show"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
program = rec.custom_properties["program"]
self.assertEqual(program["title"], "My Show")
self.assertTrue(program["user_edited"])
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_update_description_only(self, _ws):
rec = self._make_rec({"program": {"title": "Existing Title"}})
response = self._update(rec, {"description": "A great episode"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
program = rec.custom_properties["program"]
self.assertEqual(program["description"], "A great episode")
self.assertEqual(program["title"], "Existing Title")
self.assertTrue(program["user_edited"])
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_update_both_fields(self, _ws):
rec = self._make_rec()
response = self._update(rec, {"title": "New Title", "description": "New Desc"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
program = rec.custom_properties["program"]
self.assertEqual(program["title"], "New Title")
self.assertEqual(program["description"], "New Desc")
self.assertTrue(program["user_edited"])
def test_no_fields_returns_400(self):
rec = self._make_rec()
response = self._update(rec, {})
self.assertEqual(response.status_code, 400)
self.assertFalse(response.data.get("success"))
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_whitespace_trimmed(self, _ws):
rec = self._make_rec()
response = self._update(rec, {"title": " Padded Title "})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["program"]["title"], "Padded Title")
def test_whitespace_only_title_returns_400(self):
"""Whitespace-only title and description should be rejected."""
rec = self._make_rec()
response = self._update(rec, {"title": " ", "description": " "})
self.assertEqual(response.status_code, 400)
self.assertFalse(response.data.get("success"))
def test_whitespace_only_title_with_valid_description(self):
"""Whitespace-only title is ignored; valid description is accepted."""
rec = self._make_rec({"program": {"title": "Original"}})
with patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None):
response = self._update(rec, {"title": " ", "description": "Valid desc"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
# Title should remain unchanged since the whitespace-only value is not applied
self.assertEqual(rec.custom_properties["program"]["title"], "Original")
self.assertEqual(rec.custom_properties["program"]["description"], "Valid desc")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_creates_program_dict_when_absent(self, _ws):
"""Recording with no program dict gets one created."""
rec = self._make_rec({"status": "completed"})
response = self._update(rec, {"title": "Brand New"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
self.assertIn("program", rec.custom_properties)
self.assertEqual(rec.custom_properties["program"]["title"], "Brand New")
def test_returns_404_for_nonexistent(self):
request = self.factory.post(
"/api/channels/recordings/99999/update-metadata/",
{"title": "Ghost"}, format="json",
)
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "update_metadata"})
self.assertEqual(view(request, pk=99999).status_code, 404)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_sends_websocket_event(self, mock_ws):
rec = self._make_rec()
self._update(rec, {"title": "WS Test"})
mock_ws.assert_called_once()
payload = mock_ws.call_args[0][2]
self.assertEqual(payload["type"], "recording_updated")
self.assertEqual(payload["recording_id"], rec.id)
@patch("core.utils.send_websocket_update", side_effect=Exception("WS down"))
def test_ws_failure_does_not_fail_request(self, _ws):
"""WebSocket errors are silenced — the save still succeeds."""
rec = self._make_rec()
response = self._update(rec, {"title": "Resilient"})
self.assertEqual(response.status_code, 200)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["program"]["title"], "Resilient")
# ---------------------------------------------------------------------------
# refresh_artwork endpoint
# ---------------------------------------------------------------------------
class RefreshArtworkTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/refresh-artwork/"""
def setUp(self):
self.channel = Channel.objects.create(channel_number=71, name="Artwork Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _refresh(self, rec):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
return view(request, pk=rec.id)
def _make_rec(self, custom_properties=None):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties=custom_properties or {},
)
@patch("threading.Thread")
def test_returns_200_immediately(self, mock_thread):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
response = self._refresh(rec)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data.get("success"))
@patch("threading.Thread")
def test_spawns_background_thread(self, mock_thread):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
self._refresh(rec)
mock_thread.assert_called_once()
self.assertTrue(mock_thread.call_args[1].get("daemon", False))
mock_thread.return_value.start.assert_called_once()
def test_returns_404_for_nonexistent(self):
request = self.factory.post("/api/channels/recordings/99999/refresh-artwork/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
self.assertEqual(view(request, pk=99999).status_code, 404)
@patch("django.db.close_old_connections")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_no_downgrade_to_channel_logo(self, _ws, _close):
"""When the pipeline returns the channel's own logo, existing poster is preserved."""
logo = Logo.objects.create(name="Channel Logo", url="https://example.com/ch.png")
self.channel.logo = logo
self.channel.save()
rec = self._make_rec({
"poster_logo_id": 999, # existing real poster
"poster_url": "https://tmdb.com/real-poster.jpg",
})
with patch("apps.channels.tasks._resolve_poster_for_program",
return_value=(logo.id, None)):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
force_authenticate(request, user=self.user)
# Run synchronously by intercepting the thread
captured_fn = None
def capture_thread(*args, **kwargs):
nonlocal captured_fn
captured_fn = kwargs.get("target") or args[0]
mock = MagicMock()
mock.start = lambda: captured_fn(*kwargs.get("args", ()))
return mock
with patch("threading.Thread", side_effect=capture_thread):
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
view(request, pk=rec.id)
rec.refresh_from_db()
# Existing poster should be preserved — not downgraded to channel logo
self.assertEqual(rec.custom_properties.get("poster_logo_id"), 999)
self.assertEqual(rec.custom_properties.get("poster_url"), "https://tmdb.com/real-poster.jpg")
@patch("django.db.close_old_connections")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_upgrade_from_no_poster(self, _ws, _close):
"""When a recording has no poster and the pipeline finds one, it gets updated."""
rec = self._make_rec({"program": {"title": "Some Show", "id": 42}})
with patch("apps.channels.tasks._resolve_poster_for_program",
return_value=(555, "https://tmdb.com/new-poster.jpg")):
captured_fn = None
def capture_thread(*args, **kwargs):
nonlocal captured_fn
captured_fn = kwargs.get("target") or args[0]
mock = MagicMock()
mock.start = lambda: captured_fn(*kwargs.get("args", ()))
return mock
with patch("threading.Thread", side_effect=capture_thread):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
view(request, pk=rec.id)
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("poster_logo_id"), 555)
self.assertEqual(rec.custom_properties.get("poster_url"), "https://tmdb.com/new-poster.jpg")
# ---------------------------------------------------------------------------
# Logo proxy negative cache
# ---------------------------------------------------------------------------
class LogoNegativeCacheTests(TestCase):
"""Tests for the _logo_fetch_failures negative cache in LogoViewSet.cache()."""
def setUp(self):
from apps.channels import api_views
self._failures = api_views._logo_fetch_failures
self._failures.clear()
self.factory = APIRequestFactory()
self.user = _make_admin()
def _fetch_logo(self, logo):
request = self.factory.get(f"/api/channels/logos/{logo.id}/cache/")
force_authenticate(request, user=self.user)
view = LogoViewSet.as_view({"get": "cache"})
return view(request, pk=logo.id)
def test_failed_url_cached_on_non_200(self):
"""Non-200 response adds URL to negative cache."""
logo = Logo.objects.create(name="Dead Logo", url="https://dead-cdn.com/logo.png")
mock_resp = MagicMock(status_code=404)
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 404)
self.assertIn("https://dead-cdn.com/logo.png", self._failures)
def test_cached_failure_returns_404_immediately(self):
"""Subsequent request for a cached-failed URL returns 404 without making a request."""
logo = Logo.objects.create(name="Cached Fail", url="https://cached-fail.com/logo.png")
self._failures["https://cached-fail.com/logo.png"] = time_mod.monotonic() + 300
with patch("apps.channels.api_views.requests.get") as mock_get:
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 404)
mock_get.assert_not_called()
def test_expired_cache_entry_allows_retry(self):
"""After TTL expires, a new request is made."""
logo = Logo.objects.create(name="Expired", url="https://expired.com/logo.png")
self._failures["https://expired.com/logo.png"] = time_mod.monotonic() - 1 # already expired
mock_resp = MagicMock(status_code=200)
mock_resp.headers = {"Content-Type": "image/png"}
mock_resp.iter_content = MagicMock(return_value=[b"img"])
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 200)
def test_success_clears_previous_failure(self):
"""A successful fetch removes the URL from the failure cache."""
url = "https://recovered.com/logo.png"
logo = Logo.objects.create(name="Recovered", url=url)
self._failures[url] = time_mod.monotonic() - 1 # expired
mock_resp = MagicMock(status_code=200)
mock_resp.headers = {"Content-Type": "image/png"}
mock_resp.iter_content = MagicMock(return_value=[b"img"])
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
self._fetch_logo(logo)
self.assertNotIn(url, self._failures)
def test_request_exception_cached(self):
"""Network errors are cached the same as non-200 responses."""
import requests
logo = Logo.objects.create(name="Timeout", url="https://timeout.com/logo.png")
with patch("apps.channels.api_views.requests.get", side_effect=requests.Timeout("timed out")), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
response = self._fetch_logo(logo)
self.assertEqual(response.status_code, 404)
self.assertIn("https://timeout.com/logo.png", self._failures)
def test_eviction_when_cache_exceeds_256(self):
"""Stale entries are evicted when the cache grows past 256."""
now = time_mod.monotonic()
# Fill with 257 expired entries
for i in range(257):
self._failures[f"https://old-{i}.com/x.png"] = now - 1 # already expired
logo = Logo.objects.create(name="Trigger", url="https://trigger-evict.com/logo.png")
import requests
with patch("apps.channels.api_views.requests.get", side_effect=requests.ConnectionError("fail")), \
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
self._fetch_logo(logo)
# Expired entries should be evicted
old_entries = [k for k in self._failures if k.startswith("https://old-")]
self.assertEqual(len(old_entries), 0)
# New failure entry should exist
self.assertIn("https://trigger-evict.com/logo.png", self._failures)
@@ -0,0 +1,524 @@
"""Tests for recent DVR fixes.
Covers:
1. Collision avoidance: _build_output_paths checks both .mkv and .ts files
2. Logo guard: _resolve_poster_for_program skips external APIs when title ≈ channel name
3. Recording status lifecycle: status transitions visible via API
4. Concat flags: error-tolerant ffmpeg flags used for segment concatenation
5. Recovery skip-list: "recording" status NOT in terminal skip list
"""
import os
from datetime import timedelta
from unittest.mock import MagicMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="dvr_fixes_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
def _make_channel(name="Test Channel", number=100):
return Channel.objects.create(channel_number=number, name=name)
def _make_recording(channel, **overrides):
now = timezone.now()
defaults = {
"channel": channel,
"start_time": now - timedelta(hours=1),
"end_time": now + timedelta(hours=1),
"custom_properties": {},
}
defaults.update(overrides)
return Recording.objects.create(**defaults)
# =========================================================================
# 1. Collision avoidance — _build_output_paths
# =========================================================================
class CollisionAvoidanceTests(TestCase):
"""_build_output_paths must increment the filename counter when
EITHER the .mkv OR the .ts file already exists with size > 0."""
def _call(self, channel, program, start, end):
from apps.channels.tasks import _build_output_paths
return _build_output_paths(channel, program, start, end)
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_no_collision_when_nothing_exists(self, _tv, _fb):
"""Fresh path — no files exist, counter stays at 1."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
raise OSError("No such file")
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
# Should NOT have a _2 suffix
self.assertNotIn("_2", final)
self.assertTrue(final.endswith(".mkv"))
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_collision_when_ts_exists_but_mkv_is_zero_bytes(self, _tv, _fb):
"""Pre-restart scenario: MKV is 0-byte placeholder, TS has real data.
The old code only checked MKV size, so it would reuse the path.
The fix also checks TS, so it must increment."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
if "_2" in path:
raise OSError("No such file")
result = MagicMock()
if path.endswith('.mkv'):
result.st_size = 0 # MKV is 0-byte placeholder
elif path.endswith('.ts'):
result.st_size = 5000000 # TS has real data from pre-restart
else:
result.st_size = 0
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
# Must have incremented to _2
self.assertIn("_2", final, "Should increment counter when TS file has data")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_collision_when_mkv_has_data(self, _tv, _fb):
"""Standard collision: MKV file has data, should increment."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
if "_2" in path:
raise OSError("No such file")
result = MagicMock()
if path.endswith('.mkv'):
result.st_size = 1000000 # MKV has data
else:
result.st_size = 0
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
self.assertIn("_2", final, "Should increment counter when MKV file has data")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_no_collision_when_both_zero_bytes(self, _tv, _fb):
"""Both MKV and TS exist but are 0 bytes — no collision."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
result = MagicMock()
result.st_size = 0 # All files empty
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
self.assertNotIn("_2", final, "Should NOT increment when all files are empty")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
return_value="TV/{show}/{start}.mkv")
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
def test_collision_increments_to_3_when_2_also_occupied(self, _tv, _fb):
"""When both base and _2 are occupied, should go to _3."""
ch = MagicMock(name="TestCh")
ch.name = "TestCh"
program = {"title": "My Show"}
now = timezone.now()
def mock_stat(path):
if "_3" in path:
raise OSError("No such file")
result = MagicMock()
if path.endswith('.ts'):
result.st_size = 5000000
else:
result.st_size = 0
return result
with patch("os.stat", side_effect=mock_stat), \
patch("os.makedirs"):
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
self.assertIn("_3", final, "Should increment to _3 when base and _2 are occupied")
# =========================================================================
# 2. Logo guard — _resolve_poster_for_program
# =========================================================================
class LogoGuardTests(TestCase):
"""When the program title matches the channel name, external API
searches (VOD, TMDB, OMDb, TVMaze, iTunes) must be skipped."""
def _call(self, channel_name, program, channel_logo_id=None):
from apps.channels.tasks import _resolve_poster_for_program
return _resolve_poster_for_program(channel_name, program, channel_logo_id)
@patch("apps.channels.tasks.requests.get")
def test_channel_name_as_title_skips_external_apis(self, mock_get):
"""Title = 'USA A&E SD*', channel = 'USA A&E SD*' → no external calls."""
program = {"title": "USA A&E SD*"}
logo_id, url = self._call("USA A&E SD*", program, channel_logo_id=42)
# Should NOT have called any external APIs
mock_get.assert_not_called()
# Should fall back to channel logo
self.assertEqual(logo_id, 42)
self.assertIsNone(url)
@patch("apps.channels.tasks.requests.get")
def test_channel_name_normalized_match(self, mock_get):
"""Title = 'fox news', channel = 'FOX-News*' → normalized match, skip APIs."""
program = {"title": "fox news"}
logo_id, url = self._call("FOX-News*", program, channel_logo_id=99)
mock_get.assert_not_called()
self.assertEqual(logo_id, 99)
@patch("apps.channels.tasks.requests.get")
def test_real_title_still_searched(self, mock_get):
"""Title = 'Breaking Bad' on channel 'AMC' → should try external APIs."""
# Mock TVMaze returning a result
mock_resp = MagicMock(ok=True, status_code=200)
mock_resp.json.return_value = {
"image": {"original": "https://tvmaze.com/breaking-bad.jpg"}
}
mock_get.return_value = mock_resp
program = {"title": "Breaking Bad"}
logo_id, url = self._call("AMC", program)
# Should have made at least one external API call
self.assertTrue(mock_get.called, "Should search external APIs for real titles")
self.assertIsNotNone(url)
@patch("apps.channels.tasks.requests.get")
def test_no_title_skips_to_channel_logo(self, mock_get):
"""No title at all → falls through to channel logo, no API calls."""
program = {}
logo_id, url = self._call("SomeChannel", program, channel_logo_id=55)
mock_get.assert_not_called()
self.assertEqual(logo_id, 55)
@patch("apps.channels.tasks.requests.get")
def test_epg_image_still_used_even_when_title_is_channel_name(self, mock_get):
"""Even when title = channel name, Stage 1 (EPG images) should still work."""
from apps.epg.models import ProgramData, EPGSource, EPGData
# Create an EPG source + EPGData entry + program with an icon URL
epg_source = EPGSource.objects.create(source_type="xmltv", name="Test EPG")
epg_data = EPGData.objects.create(tvg_id="test.ch", epg_source=epg_source)
prog = ProgramData.objects.create(
epg=epg_data,
title="Test Channel HD",
start_time=timezone.now() - timedelta(hours=1),
end_time=timezone.now() + timedelta(hours=1),
custom_properties={"icon": "https://epg-cdn.com/test-icon.png"},
)
program = {"title": "Test Channel HD", "id": prog.id}
# Mock _validate_url to return True for the icon URL
with patch("apps.channels.tasks._validate_url", return_value=True):
logo_id, url = self._call("Test Channel HD", program, channel_logo_id=10)
# EPG icon should still be used (Stage 1 doesn't depend on title guard)
self.assertEqual(url, "https://epg-cdn.com/test-icon.png")
mock_get.assert_not_called()
# =========================================================================
# 3. Recording status lifecycle via API
# =========================================================================
class RecordingStatusLifecycleTests(TestCase):
"""Verify recording status transitions and that terminal recordings
are properly filterable (supports the red-dot fix in guideUtils)."""
def setUp(self):
self.channel = _make_channel("Status Test Channel", 200)
self.user = _make_admin()
self.factory = APIRequestFactory()
def _list_recordings(self):
from apps.channels.api_views import RecordingViewSet
request = self.factory.get("/api/channels/recordings/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"get": "list"})
return view(request)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_stopped_recording_has_terminal_status(self, _ws):
"""After stop, custom_properties.status = 'stopped'."""
from apps.channels.api_views import RecordingViewSet
rec = _make_recording(self.channel, custom_properties={
"status": "recording",
"program": {"id": 1, "title": "Live Show"},
})
request = self.factory.post(f"/api/channels/recordings/{rec.id}/stop/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "stop"})
with patch("apps.channels.signals.revoke_task"):
response = view(request, pk=rec.id)
self.assertIn(response.status_code, [200, 204])
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("status"), "stopped")
def test_listing_includes_status_in_custom_properties(self):
"""API listing returns custom_properties with status field."""
_make_recording(self.channel, custom_properties={
"status": "recording",
"program": {"id": 1, "title": "Recording Show"},
})
_make_recording(self.channel, custom_properties={
"status": "stopped",
"program": {"id": 2, "title": "Stopped Show"},
})
response = self._list_recordings()
self.assertEqual(response.status_code, 200)
statuses = [r["custom_properties"].get("status") for r in response.data]
self.assertIn("recording", statuses)
self.assertIn("stopped", statuses)
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_delete_recording_removes_from_listing(self, _ws):
"""Deleting a recording removes it from the listing entirely."""
from apps.channels.api_views import RecordingViewSet
rec = _make_recording(self.channel, custom_properties={
"status": "stopped",
"program": {"id": 3, "title": "To Delete"},
})
rec_id = rec.id
request = self.factory.delete(f"/api/channels/recordings/{rec_id}/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"delete": "destroy"})
with patch("apps.channels.signals.revoke_task"):
response = view(request, pk=rec_id)
self.assertIn(response.status_code, [200, 204])
self.assertFalse(Recording.objects.filter(id=rec_id).exists())
# =========================================================================
# 4. Concat flags — error-tolerant ffmpeg
# =========================================================================
class ConcatFlagsTests(TestCase):
"""Verify that the finalize phase uses error-tolerant ffmpeg flags
when concatenating pre-restart segments."""
def test_concat_command_includes_error_tolerant_flags(self):
"""Inspect the source code to confirm error-tolerant flags are present.
This is a static analysis test — no ffmpeg execution needed."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
# The concat subprocess.run call must include these flags
self.assertIn("+genpts+igndts+discardcorrupt", source,
"Concat must use +genpts+igndts+discardcorrupt fflags")
self.assertIn("ignore_err", source,
"Concat must use -err_detect ignore_err")
self.assertIn("-f", source)
self.assertIn("concat", source)
def test_concat_goes_directly_to_mkv(self):
"""Concat must produce MKV directly (not intermediate .ts) to
preserve timestamp boundaries and avoid playback freeze at splice."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
# Must contain reset_timestamps for proper segment boundary handling
self.assertIn("reset_timestamps", source,
"Concat must use -reset_timestamps 1 for seamless seeking")
# Must write directly to final_path (MKV), not an intermediate .ts
self.assertIn("_concat_did_remux", source,
"Concat path must set flag to skip separate remux step")
def test_segment_time_metadata_present(self):
"""Verify concat uses -segment_time_metadata for boundary awareness."""
import inspect
from apps.channels.tasks import run_recording
source = inspect.getsource(run_recording)
self.assertIn("segment_time_metadata", source,
"Concat must use -segment_time_metadata 1 for segment boundary handling")
# =========================================================================
# 5. Recovery skip-list
# =========================================================================
class RecoverySkipListTests(TestCase):
"""Verify that the recovery function does NOT skip 'recording' status,
since that's the exact status recordings have when the server crashes."""
def test_recording_status_not_in_skip_list(self):
"""Inspect recover_recordings_on_startup to ensure 'recording' is
NOT treated as a terminal/skip state."""
import inspect
from apps.channels.tasks import recover_recordings_on_startup
source = inspect.getsource(recover_recordings_on_startup)
# Find the skip condition line
# It should be: if current_status in ("completed", "stopped"):
# NOT: if current_status in ("completed", "stopped", "recording"):
lines = source.split('\n')
skip_line = None
for line in lines:
if 'current_status in' in line and ('completed' in line or 'stopped' in line):
skip_line = line.strip()
break
self.assertIsNotNone(skip_line, "Should find the skip-list condition")
self.assertNotIn('"recording"', skip_line,
"Skip list must NOT contain 'recording' — "
"that's the status of crashed mid-stream recordings that need recovery")
@patch("core.utils.RedisClient")
@patch("apps.channels.tasks.run_recording")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_recovery_processes_recording_status(self, _ws, mock_run, mock_redis_cls):
"""A recording with status='recording' should be recovered, not skipped."""
mock_redis_conn = MagicMock()
mock_redis_conn.set.return_value = True # Acquire lock
mock_redis_cls.get_client.return_value = mock_redis_conn
channel = _make_channel("Recovery Test", 300)
now = timezone.now()
rec = _make_recording(channel, custom_properties={
"status": "recording",
"program": {"title": "Crashed Show"},
}, end_time=now + timedelta(hours=2))
from apps.channels.tasks import recover_recordings_on_startup
with patch("apps.channels.signals.revoke_task"):
result = recover_recordings_on_startup()
# The recording should have been dispatched for recovery
self.assertTrue(mock_run.apply_async.called,
"Recording with status='recording' should be dispatched for recovery")
@patch("core.utils.RedisClient")
@patch("apps.channels.tasks.run_recording")
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
def test_recovery_skips_stopped_recordings(self, _ws, mock_run, mock_redis_cls):
"""A recording with status='stopped' should be skipped by recovery."""
mock_redis_conn = MagicMock()
mock_redis_conn.set.return_value = True
mock_redis_cls.get_client.return_value = mock_redis_conn
channel = _make_channel("Recovery Skip Test", 301)
now = timezone.now()
rec = _make_recording(channel, custom_properties={
"status": "stopped",
"program": {"title": "Finished Show"},
}, end_time=now + timedelta(hours=2))
from apps.channels.tasks import recover_recordings_on_startup
with patch("apps.channels.signals.revoke_task"):
recover_recordings_on_startup()
# Should NOT have dispatched a recovery task
mock_run.apply_async.assert_not_called()
# =========================================================================
# 6. Frontend red-dot filter (guideUtils.mapRecordingsByProgramId)
# =========================================================================
class MapRecordingsByProgramIdTests(TestCase):
"""These test the BACKEND side — confirming that recording status
is preserved in the API response so the frontend can filter on it.
The actual frontend filtering is covered by frontend/src/pages/__tests__/DVR.test.jsx
and the guideUtils code, but we verify the data contract here."""
def test_recording_custom_properties_status_persisted(self):
"""Recording status in custom_properties survives save/load cycle."""
channel = _make_channel("Red Dot Test", 400)
rec = _make_recording(channel, custom_properties={
"status": "stopped",
"program": {"id": 42, "title": "A Show"},
})
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], "stopped")
def test_terminal_statuses_are_well_defined(self):
"""Verify the terminal status set matches what the frontend uses."""
# These are the statuses that should NOT show a red dot in the Guide
terminal = {"stopped", "completed", "interrupted", "failed"}
channel = _make_channel("Terminal Status Test", 410)
# Verify each status is a valid recording status
for status in terminal:
rec = _make_recording(channel, custom_properties={
"status": status,
"program": {"id": 100, "title": "Test"},
})
rec.refresh_from_db()
self.assertEqual(rec.custom_properties["status"], status)
@@ -0,0 +1,585 @@
"""Tests for DVR recording scheduling with ClockedSchedule.
Uses ClockedSchedule instead of apply_async with countdown because Redis
visibility_timeout (default 3600s) causes task redelivery for long countdowns,
leading to duplicate recordings.
"""
from datetime import timedelta
from unittest.mock import patch, MagicMock
from django.test import TestCase
from django.utils import timezone
from django_celery_beat.models import ClockedSchedule, PeriodicTask
from apps.channels.models import Channel, Recording
from apps.channels.signals import (
schedule_recording_task,
revoke_task,
_dvr_task_name,
)
class ScheduleRecordingTaskTests(TestCase):
"""Tests for schedule_recording_task()."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.run_recording")
def test_future_recording_creates_periodic_task(self, mock_run_recording):
"""Recordings in the future create a ClockedSchedule + PeriodicTask."""
future_time = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future_time,
end_time=future_time + timedelta(hours=1),
)
task_id = schedule_recording_task(rec, eta=future_time)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
pt = PeriodicTask.objects.get(name=expected_name)
self.assertTrue(pt.one_off)
self.assertTrue(pt.enabled)
self.assertEqual(pt.task, "apps.channels.tasks.run_recording")
self.assertIsNotNone(pt.clocked)
# apply_async should not have been called
mock_run_recording.apply_async.assert_not_called()
@patch("apps.channels.signals.run_recording")
def test_immediate_recording_creates_periodic_task(self, mock_run_recording):
"""Recordings starting now also use ClockedSchedule for consistency."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now,
end_time=now + timedelta(hours=1),
)
task_id = schedule_recording_task(rec, eta=now)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
self.assertTrue(PeriodicTask.objects.filter(name=expected_name).exists())
@patch("apps.channels.signals.run_recording")
def test_past_start_time_clamps_to_now(self, mock_run_recording):
"""Recordings with past start_time get clamped to now."""
past_time = timezone.now() - timedelta(minutes=5)
rec = Recording.objects.create(
channel=self.channel,
start_time=past_time,
end_time=timezone.now() + timedelta(hours=1),
)
task_id = schedule_recording_task(rec, eta=past_time)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
pt = PeriodicTask.objects.get(name=expected_name)
# Clocked time should be >= now
self.assertGreaterEqual(pt.clocked.clocked_time, past_time)
@patch("apps.channels.signals.run_recording")
def test_reschedule_updates_existing_periodic_task(self, mock_run_recording):
"""Calling schedule_recording_task twice updates the existing PeriodicTask."""
future_time = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future_time,
end_time=future_time + timedelta(hours=1),
)
schedule_recording_task(rec, eta=future_time)
# Reschedule with a different time
new_eta = future_time + timedelta(hours=1)
schedule_recording_task(rec, eta=new_eta)
# Should still be exactly one PeriodicTask
task_name = f"dvr-recording-{rec.id}"
self.assertEqual(PeriodicTask.objects.filter(name=task_name).count(), 1)
@patch("apps.channels.signals.run_recording")
def test_naive_eta_is_made_aware(self, mock_run_recording):
"""A naive (timezone-unaware) eta is made timezone-aware."""
from datetime import datetime
naive_eta = datetime(2030, 6, 15, 14, 0, 0)
rec = Recording.objects.create(
channel=self.channel,
start_time=timezone.now() + timedelta(hours=1),
end_time=timezone.now() + timedelta(hours=2),
)
task_id = schedule_recording_task(rec, eta=naive_eta)
expected_name = f"dvr-recording-{rec.id}"
self.assertEqual(task_id, expected_name)
pt = PeriodicTask.objects.get(name=expected_name)
self.assertTrue(timezone.is_aware(pt.clocked.clocked_time))
class RevokeTaskTests(TestCase):
"""Tests for revoke_task()."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
def test_revoke_deletes_periodic_task_and_clocked_schedule(self):
"""revoke_task deletes the PeriodicTask and orphaned ClockedSchedule."""
eta = timezone.now() + timedelta(hours=5)
clocked = ClockedSchedule.objects.create(clocked_time=eta)
PeriodicTask.objects.create(
name="dvr-recording-10",
task="apps.channels.tasks.run_recording",
clocked=clocked,
one_off=True,
enabled=True,
)
revoke_task("dvr-recording-10")
self.assertFalse(PeriodicTask.objects.filter(name="dvr-recording-10").exists())
self.assertFalse(ClockedSchedule.objects.filter(id=clocked.id).exists())
def test_revoke_keeps_shared_clocked_schedule(self):
"""ClockedSchedule is kept if another PeriodicTask still references it."""
eta = timezone.now() + timedelta(hours=5)
clocked = ClockedSchedule.objects.create(clocked_time=eta)
PeriodicTask.objects.create(
name="dvr-recording-10",
task="apps.channels.tasks.run_recording",
clocked=clocked,
one_off=True,
)
PeriodicTask.objects.create(
name="dvr-recording-11",
task="apps.channels.tasks.run_recording",
clocked=clocked,
one_off=True,
)
revoke_task("dvr-recording-10")
self.assertFalse(PeriodicTask.objects.filter(name="dvr-recording-10").exists())
self.assertTrue(ClockedSchedule.objects.filter(id=clocked.id).exists())
@patch("apps.channels.signals.AsyncResult")
def test_revoke_falls_back_to_async_result_for_legacy_ids(self, mock_async_result):
"""revoke_task falls back to AsyncResult.revoke() for old-style UUIDs."""
revoke_task("550e8400-e29b-41d4-a716-446655440000")
mock_async_result.assert_called_once_with("550e8400-e29b-41d4-a716-446655440000")
mock_async_result.return_value.revoke.assert_called_once()
def test_revoke_none_is_noop(self):
"""revoke_task(None) does nothing."""
revoke_task(None) # Should not raise
def test_revoke_empty_string_is_noop(self):
"""revoke_task('') does nothing."""
revoke_task("") # Should not raise
class DvrTaskNameTests(TestCase):
"""Tests for the naming convention helper."""
def test_task_name_format(self):
self.assertEqual(_dvr_task_name(42), "dvr-recording-42")
def test_task_name_fits_in_charfield(self):
name = _dvr_task_name(999999999)
self.assertLessEqual(len(name), 255)
class SignalIntegrationTests(TestCase):
"""Integration tests for the post_save / post_delete signal handlers."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_save_creates_periodic_task_for_future_recording(self, mock_artwork):
"""Saving a future Recording creates a PeriodicTask via post_save signal."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
)
rec.refresh_from_db()
task_name = f"dvr-recording-{rec.id}"
self.assertEqual(rec.task_id, task_name)
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_delete_removes_periodic_task(self, mock_artwork):
"""Deleting a Recording removes its PeriodicTask."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
)
rec.refresh_from_db()
task_name = rec.task_id
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
rec.delete()
self.assertFalse(PeriodicTask.objects.filter(name=task_name).exists())
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_bulk_delete_cleans_up_all_periodic_tasks(self, mock_artwork):
"""Bulk deleting recordings cleans up all their PeriodicTasks."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec_ids = []
for i in range(5):
rec = Recording.objects.create(
channel=self.channel,
start_time=future + timedelta(hours=i),
end_time=future + timedelta(hours=i + 1),
)
rec_ids.append(rec.id)
for rid in rec_ids:
self.assertTrue(
PeriodicTask.objects.filter(name=f"dvr-recording-{rid}").exists()
)
Recording.objects.filter(channel=self.channel).delete()
self.assertEqual(
PeriodicTask.objects.filter(name__startswith="dvr-recording-").count(), 0
)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_save_schedules_currently_playing_recording(self, mock_artwork):
"""A recording with past start_time but future end_time schedules immediately."""
mock_artwork.apply_async.return_value = MagicMock()
past_start = timezone.now() - timedelta(minutes=30)
future_end = timezone.now() + timedelta(minutes=30)
rec = Recording.objects.create(
channel=self.channel,
start_time=past_start,
end_time=future_end,
)
rec.refresh_from_db()
task_name = f"dvr-recording-{rec.id}"
self.assertEqual(rec.task_id, task_name)
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_post_save_skips_fully_past_recording(self, mock_artwork):
"""A recording with both start_time and end_time in the past is not scheduled."""
mock_artwork.apply_async.return_value = MagicMock()
past_start = timezone.now() - timedelta(hours=2)
past_end = timezone.now() - timedelta(hours=1)
rec = Recording.objects.create(
channel=self.channel,
start_time=past_start,
end_time=past_end,
)
rec.refresh_from_db()
self.assertIsNone(rec.task_id)
self.assertFalse(
PeriodicTask.objects.filter(name=f"dvr-recording-{rec.id}").exists()
)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_pre_save_revokes_on_time_change(self, mock_artwork):
"""Changing a recording's start_time revokes the old task and creates a new one."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
)
rec.refresh_from_db()
old_task_name = rec.task_id
self.assertTrue(PeriodicTask.objects.filter(name=old_task_name).exists())
# Change the start time — pre_save clears task_id, post_save reschedules
new_future = future + timedelta(hours=3)
rec.start_time = new_future
rec.end_time = new_future + timedelta(hours=1)
rec.save()
rec.refresh_from_db()
# Old PeriodicTask should be deleted; new one should exist
self.assertIsNotNone(rec.task_id)
self.assertTrue(
PeriodicTask.objects.filter(name=f"dvr-recording-{rec.id}").exists()
)
class IdempotencyGuardTests(TestCase):
"""Tests for the idempotency guard in run_recording()."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
@patch("apps.channels.tasks.get_channel_layer")
def test_skips_if_already_recording(self, mock_layer):
"""run_recording returns early if status is already 'recording'."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now,
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording", "started_at": str(now)},
)
from apps.channels.tasks import run_recording as run_rec_task
result = run_rec_task(rec.id, self.channel.id, str(now), str(now + timedelta(hours=1)))
self.assertIsNone(result)
# get_channel_layer should not have been called (returned before)
mock_layer.assert_not_called()
@patch("apps.channels.tasks.get_channel_layer")
def test_skips_if_already_completed(self, mock_layer):
"""run_recording returns early if status is already 'completed'."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=2),
end_time=now - timedelta(hours=1),
custom_properties={"status": "completed"},
)
from apps.channels.tasks import run_recording as run_rec_task
result = run_rec_task(rec.id, self.channel.id, str(rec.start_time), str(rec.end_time))
self.assertIsNone(result)
mock_layer.assert_not_called()
@patch("apps.channels.tasks.get_channel_layer")
def test_skips_if_already_stopped(self, mock_layer):
"""run_recording returns early if status is already 'stopped' (user stopped it early)."""
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": "stopped", "stopped_at": str(now)},
)
from apps.channels.tasks import run_recording as run_rec_task
result = run_rec_task(rec.id, self.channel.id, str(rec.start_time), str(rec.end_time))
self.assertIsNone(result)
mock_layer.assert_not_called()
class ArtworkPrefetchSignalGuardTests(TestCase):
"""Tests that the post_save signal does not schedule artwork prefetch when
the recording is in an active or terminal state."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_not_scheduled_when_status_recording(self, mock_artwork):
"""post_save must NOT schedule artwork prefetch when status='recording'
to prevent a race that overwrites the running task's status updates."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={"status": "recording"},
)
# Simulate a save that run_recording itself might do mid-recording
rec.custom_properties = {"status": "recording", "file_path": "/data/recordings/test.mkv"}
rec.save(update_fields=["custom_properties"])
# apply_async was not called for the "recording" save
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_not_scheduled_when_status_completed(self, mock_artwork):
"""post_save must NOT schedule artwork prefetch when status='completed'."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={"status": "completed"},
)
rec.custom_properties = {"status": "completed"}
rec.save(update_fields=["custom_properties"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_not_scheduled_when_status_stopped(self, mock_artwork):
"""post_save must NOT schedule artwork prefetch when status='stopped'."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={"status": "stopped"},
)
rec.custom_properties = {"status": "stopped"}
rec.save(update_fields=["custom_properties"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_artwork_prefetch_scheduled_for_new_upcoming_recording(self, mock_artwork):
"""post_save SHOULD schedule artwork prefetch for a newly created upcoming recording."""
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={}, # no status yet — should trigger prefetch
)
self.assertTrue(mock_artwork.apply_async.called)
class DestroyDvrClientIsolationTests(TestCase):
"""Tests that deleting a recording only stops DVR clients when the
recording is actively streaming — never for completed/upcoming recordings
that could share a channel with an unrelated in-progress recording."""
def setUp(self):
from django.contrib.auth import get_user_model
from rest_framework.test import APIRequestFactory, force_authenticate
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
User = get_user_model()
self.user = User.objects.create_user(
username="dvr_test_admin", password="pass",
user_level=User.UserLevel.ADMIN,
)
self.factory = APIRequestFactory()
self.force_authenticate = force_authenticate
def _delete_recording(self, rec):
from apps.channels.api_views import RecordingViewSet
request = self.factory.delete(f"/api/channels/recordings/{rec.id}/")
self.force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"delete": "destroy"})
return view(request, pk=rec.id)
@patch("apps.channels.api_views._stop_dvr_clients")
def test_destroy_completed_recording_does_not_stop_dvr_clients(self, mock_stop):
"""Deleting a completed recording must NOT call _stop_dvr_clients."""
rec = Recording.objects.create(
channel=self.channel,
start_time=timezone.now() - timedelta(hours=2),
end_time=timezone.now() - timedelta(hours=1),
custom_properties={"status": "completed", "file_path": "/data/recordings/test.mkv"},
)
self._delete_recording(rec)
mock_stop.assert_not_called()
@patch("apps.channels.api_views._stop_dvr_clients")
def test_destroy_upcoming_recording_does_not_stop_dvr_clients(self, mock_stop):
"""Deleting an upcoming (scheduled) recording must NOT call _stop_dvr_clients."""
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={},
)
self._delete_recording(rec)
mock_stop.assert_not_called()
@patch("apps.channels.api_views._stop_dvr_clients")
def test_destroy_active_recording_does_stop_dvr_clients(self, mock_stop):
"""Deleting an in-progress recording MUST call _stop_dvr_clients."""
mock_stop.return_value = 1
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=5),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
self._delete_recording(rec)
mock_stop.assert_called_once_with(str(self.channel.uuid), recording_id=rec.id)
class PeriodicTaskCleanupOnExecutionTests(TestCase):
"""Tests for PeriodicTask cleanup when run_recording starts."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
def tearDown(self):
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
ClockedSchedule.objects.all().delete()
@patch("apps.channels.signals.prefetch_recording_artwork")
@patch("apps.channels.tasks.get_channel_layer")
def test_periodic_task_cleaned_up_on_execution(self, mock_layer, mock_artwork):
"""When run_recording executes, it deletes its own PeriodicTask."""
mock_layer.return_value = MagicMock()
mock_artwork.apply_async.return_value = MagicMock()
future = timezone.now() + timedelta(hours=2)
rec = Recording.objects.create(
channel=self.channel,
start_time=future,
end_time=future + timedelta(hours=1),
custom_properties={},
)
# post_save signal should have created the PeriodicTask
task_name = f"dvr-recording-{rec.id}"
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
pt = PeriodicTask.objects.get(name=task_name)
clocked_id = pt.clocked_id
from apps.channels.tasks import run_recording as run_rec_task
# This will proceed past guards, clean up the PeriodicTask, then
# eventually fail on the actual stream connection (expected)
try:
run_rec_task(rec.id, self.channel.id, str(future), str(future + timedelta(hours=1)))
except Exception:
pass
self.assertFalse(PeriodicTask.objects.filter(name=task_name).exists())
self.assertFalse(ClockedSchedule.objects.filter(id=clocked_id).exists())
@@ -0,0 +1,356 @@
"""Tests for the DVR Stop/Cancel feature set.
Covers:
- stop() endpoint
- destroy() was_in_progress field in recording_cancelled WebSocket event
- signals.py update_fields re-entrancy guard
- run_recording race guard before status write
- _stop_dvr_clients() DVR client isolation
"""
from datetime import timedelta
from unittest.mock import MagicMock, AsyncMock, patch
from django.test import TestCase
from django.utils import timezone
from rest_framework.test import APIRequestFactory, force_authenticate
from apps.channels.models import Channel, Recording
from apps.channels.api_views import RecordingViewSet, _stop_dvr_clients
def _make_admin():
from django.contrib.auth import get_user_model
User = get_user_model()
u, _ = User.objects.get_or_create(
username="stop_test_admin",
defaults={"user_level": User.UserLevel.ADMIN},
)
u.set_password("pass")
u.save()
return u
def _async_channel_layer_mock():
layer = MagicMock()
layer.group_send = AsyncMock()
return layer
class StopEndpointTests(TestCase):
"""Tests for POST /api/channels/recordings/{id}/stop/"""
def setUp(self):
self.channel = Channel.objects.create(channel_number=99, name="Stop Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _stop(self, rec):
request = self.factory.post(f"/api/channels/recordings/{rec.id}/stop/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "stop"})
return view(request, pk=rec.id)
def _make_rec(self, status="recording"):
now = timezone.now()
return Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(hours=1),
end_time=now + timedelta(hours=1),
custom_properties={"status": status},
)
@patch("core.utils.send_websocket_update")
@patch("threading.Thread")
def test_stop_writes_status_to_db_before_returning(self, mock_thread, mock_ws):
"""DB write is synchronous — run_recording polls for this."""
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
response = self._stop(rec)
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data.get("success"))
rec.refresh_from_db()
self.assertEqual(rec.custom_properties.get("status"), "stopped")
@patch("core.utils.send_websocket_update")
@patch("threading.Thread")
def test_stop_writes_stopped_at_timestamp(self, mock_thread, mock_ws):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec()
self._stop(rec)
rec.refresh_from_db()
self.assertIn("stopped_at", rec.custom_properties)
def test_stop_calls_stop_dvr_clients_in_background(self):
"""stop() spawns a background thread whose target calls _stop_dvr_clients."""
rec = self._make_rec()
with patch("apps.channels.api_views._stop_dvr_clients", return_value=1) as mock_stop, \
patch("core.utils.send_websocket_update"), \
patch("threading.Thread") as mock_thread:
mock_thread.return_value.start = MagicMock()
self._stop(rec)
# Verify a daemon thread was spawned
mock_thread.assert_called_once()
thread_kwargs = mock_thread.call_args[1]
self.assertTrue(thread_kwargs.get("daemon"), "Thread must be daemon")
# Execute the captured target with DB connection close patched out
target = thread_kwargs["target"]
with patch("apps.channels.api_views._stop_dvr_clients", return_value=1) as mock_stop2, \
patch("apps.channels.signals.revoke_task", side_effect=Exception("skip")), \
patch("django.db.connection") as mock_conn:
target()
self.assertTrue(mock_stop2.called)
args, kwargs = mock_stop2.call_args
actual_rec_id = kwargs.get("recording_id") or (args[1] if len(args) > 1 else None)
self.assertEqual(actual_rec_id, rec.id)
def test_stop_returns_404_for_nonexistent(self):
request = self.factory.post("/api/channels/recordings/99999/stop/")
force_authenticate(request, user=self.user)
view = RecordingViewSet.as_view({"post": "stop"})
self.assertEqual(view(request, pk=99999).status_code, 404)
@patch("core.utils.send_websocket_update")
@patch("threading.Thread")
def test_stop_idempotent_on_already_stopped(self, mock_thread, mock_ws):
mock_thread.return_value.start = MagicMock()
rec = self._make_rec(status="stopped")
self.assertEqual(self._stop(rec).status_code, 200)
class CancelDestroyWasInProgressTests(TestCase):
"""was_in_progress field in the recording_cancelled WebSocket event."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=98, name="Cancel Test Channel")
self.user = _make_admin()
self.factory = APIRequestFactory()
def _delete(self, rec):
request = self.factory.delete(f"/api/channels/recordings/{rec.id}/")
force_authenticate(request, user=self.user)
return RecordingViewSet.as_view({"delete": "destroy"})(request, pk=rec.id)
@patch("apps.channels.api_views._stop_dvr_clients", return_value=1)
@patch("core.utils.send_websocket_update")
def test_in_progress_sends_was_in_progress_true(self, mock_ws, _):
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=10),
end_time=now + timedelta(hours=1),
custom_properties={"status": "recording"},
)
self._delete(rec)
payload = mock_ws.call_args[0][2]
self.assertEqual(payload["type"], "recording_cancelled")
self.assertTrue(payload["was_in_progress"])
@patch("core.utils.send_websocket_update")
def test_completed_sends_was_in_progress_false(self, mock_ws):
rec = Recording.objects.create(
channel=self.channel,
start_time=timezone.now() - timedelta(hours=2),
end_time=timezone.now() - timedelta(hours=1),
custom_properties={"status": "completed"},
)
self._delete(rec)
self.assertFalse(mock_ws.call_args[0][2]["was_in_progress"])
class SignalUpdateFieldsReentrancyGuardTests(TestCase):
"""update_fields guard in schedule_task_on_save prevents redundant WS events."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=97, name="Signal Guard Channel")
def _create_upcoming(self):
future = timezone.now() + timedelta(hours=2)
return Recording.objects.create(
channel=self.channel, start_time=future,
end_time=future + timedelta(hours=1), custom_properties={},
)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_custom_properties_save_skips_artwork(self, mock_artwork):
rec = self._create_upcoming()
mock_artwork.reset_mock()
rec.custom_properties = {"poster_url": "https://example.com/p.jpg"}
rec.save(update_fields=["custom_properties"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_task_id_save_skips_artwork(self, mock_artwork):
rec = self._create_upcoming()
mock_artwork.reset_mock()
rec.task_id = "dvr-recording-999"
rec.save(update_fields=["task_id"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_combined_metadata_save_skips_artwork(self, mock_artwork):
rec = self._create_upcoming()
mock_artwork.reset_mock()
rec.task_id = "dvr-recording-1000"
rec.custom_properties = {"poster_url": "x"}
rec.save(update_fields=["custom_properties", "task_id"])
mock_artwork.apply_async.assert_not_called()
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_creation_dispatches_artwork(self, mock_artwork):
mock_artwork.apply_async.return_value = MagicMock()
self._create_upcoming()
self.assertTrue(mock_artwork.apply_async.called)
@patch("apps.channels.signals.prefetch_recording_artwork")
def test_scheduling_field_update_dispatches_artwork(self, mock_artwork):
"""save(update_fields=['start_time']) is not a metadata save — dispatch runs."""
mock_artwork.apply_async.return_value = MagicMock()
rec = self._create_upcoming()
mock_artwork.reset_mock()
future = timezone.now() + timedelta(hours=3)
rec.start_time = future
rec.end_time = future + timedelta(hours=1)
rec.save(update_fields=["start_time", "end_time"])
mock_artwork.apply_async.assert_called()
class RunRecordingRaceGuardTests(TestCase):
"""Race guard: stop() fires between idempotency check and status write."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=96, name="Race Guard Channel")
def test_race_guard_exits_when_stopped_at_db_read(self):
"""If Recording.objects.get() shows 'stopped', the task must exit
without writing 'recording' to the DB."""
from apps.channels.tasks import run_recording as run_rec
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=1),
end_time=now + timedelta(hours=1),
custom_properties={},
)
mock_layer = _async_channel_layer_mock()
original_get = Recording.objects.get
def patched_get(*args, **kwargs):
obj = original_get(*args, **kwargs)
if kwargs.get("id") == rec.id or (args and args[0] == rec.id):
obj.custom_properties = {"status": "stopped"}
return obj
with patch("apps.channels.tasks.get_channel_layer", return_value=mock_layer), \
patch("core.utils.log_system_event", side_effect=Exception("skip")), \
patch.object(Recording.objects, "get", side_effect=patched_get):
result = run_rec(
rec.id, self.channel.id, str(rec.start_time), str(rec.end_time),
)
self.assertIsNone(result)
rec.refresh_from_db()
self.assertNotEqual(
rec.custom_properties.get("status"), "recording",
"Race guard failed: task overwrote 'stopped' with 'recording'",
)
def test_idempotency_guard_catches_stopped_before_channel_layer(self):
"""When status='stopped' at the idempotency check, get_channel_layer is never called."""
from apps.channels.tasks import run_recording as run_rec
now = timezone.now()
rec = Recording.objects.create(
channel=self.channel,
start_time=now - timedelta(minutes=5),
end_time=now + timedelta(hours=1),
custom_properties={"status": "stopped"},
)
with patch("apps.channels.tasks.get_channel_layer") as mock_get_layer:
result = run_rec(
rec.id, self.channel.id, str(rec.start_time), str(rec.end_time),
)
self.assertIsNone(result)
mock_get_layer.assert_not_called()
class StopDvrClientsTests(TestCase):
"""_stop_dvr_clients() DVR client isolation."""
def setUp(self):
self.channel = Channel.objects.create(channel_number=95, name="DVR Clients Channel")
self._redis = "core.utils.RedisClient"
self._sc = "apps.proxy.ts_proxy.services.channel_service.ChannelService.stop_client"
self._sch = "apps.proxy.ts_proxy.services.channel_service.ChannelService.stop_channel"
def _mock_redis(self, client_ids, ua_map):
r = MagicMock()
r.smembers.return_value = {c.encode() for c in client_ids}
def hget_side(key, field):
ks = key if isinstance(key, str) else key.decode("utf-8", errors="replace")
for cid, ua in ua_map.items():
if cid in ks:
return ua.encode() if isinstance(ua, str) else ua
return b""
r.hget.side_effect = hget_side
return r
def test_returns_zero_when_redis_none(self):
with patch(self._redis) as rc:
rc.get_client.return_value = None
self.assertEqual(_stop_dvr_clients(str(self.channel.uuid)), 0)
def test_stops_only_matching_client_when_recording_id_given(self):
r = self._mock_redis(
["client-a", "client-b"],
{"client-a": "Dispatcharr-DVR/recording-42",
"client-b": "Dispatcharr-DVR/recording-99"},
)
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
result = _stop_dvr_clients(str(self.channel.uuid), recording_id=42)
self.assertEqual(result, 1)
stopped = [c[0][1] for c in sc.call_args_list]
self.assertIn("client-a", stopped)
self.assertNotIn("client-b", stopped)
def test_stops_all_dvr_clients_without_recording_id(self):
r = self._mock_redis(
["client-a", "client-b"],
{"client-a": "Dispatcharr-DVR/recording-42",
"client-b": "Dispatcharr-DVR/recording-99"},
)
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
result = _stop_dvr_clients(str(self.channel.uuid))
self.assertEqual(result, 2)
def test_skips_non_dvr_clients(self):
r = self._mock_redis(
["viewer", "dvr-client"],
{"viewer": "Mozilla/5.0", "dvr-client": "Dispatcharr-DVR/recording-1"},
)
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
result = _stop_dvr_clients(str(self.channel.uuid))
self.assertEqual(result, 1)
stopped = [c[0][1] for c in sc.call_args_list]
self.assertNotIn("viewer", stopped)
def test_returns_zero_for_empty_channel(self):
r = MagicMock()
r.smembers.return_value = set()
with patch(self._redis) as rc, patch(self._sc) as sc:
rc.get_client.return_value = r
self.assertEqual(_stop_dvr_clients(str(self.channel.uuid)), 0)
sc.assert_not_called()
def test_never_calls_stop_channel(self):
"""Must not stop the whole channel proxy — only individual clients."""
r = self._mock_redis(["dvr-1"], {"dvr-1": "Dispatcharr-DVR/recording-1"})
with patch(self._redis) as rc, patch(self._sc), patch(self._sch) as sch:
rc.get_client.return_value = r
_stop_dvr_clients(str(self.channel.uuid))
sch.assert_not_called()
@@ -0,0 +1,40 @@
from datetime import datetime, timedelta
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel, RecurringRecordingRule, Recording
from apps.channels.tasks import sync_recurring_rule_impl, purge_recurring_rule_impl
class RecurringRecordingRuleTasksTests(TestCase):
def test_sync_recurring_rule_creates_and_purges_recordings(self):
now = timezone.now()
channel = Channel.objects.create(channel_number=1, name='Test Channel')
start_time = (now + timedelta(minutes=15)).time().replace(second=0, microsecond=0)
end_time = (now + timedelta(minutes=75)).time().replace(second=0, microsecond=0)
rule = RecurringRecordingRule.objects.create(
channel=channel,
days_of_week=[now.weekday()],
start_time=start_time,
end_time=end_time,
)
created = sync_recurring_rule_impl(rule.id, drop_existing=True, horizon_days=1)
self.assertEqual(created, 1)
recording = Recording.objects.filter(custom_properties__rule__id=rule.id).first()
self.assertIsNotNone(recording)
self.assertEqual(recording.channel, channel)
self.assertEqual(recording.custom_properties.get('rule', {}).get('id'), rule.id)
expected_start = timezone.make_aware(
datetime.combine(recording.start_time.date(), start_time),
timezone.get_current_timezone(),
)
self.assertLess(abs((recording.start_time - expected_start).total_seconds()), 60)
removed = purge_recurring_rule_impl(rule.id)
self.assertEqual(removed, 1)
self.assertFalse(Recording.objects.filter(custom_properties__rule__id=rule.id).exists())
@@ -0,0 +1,718 @@
"""Tests for series rule evaluation deduplication.
Unit tests verify the dedup logic in evaluate_series_rules_impl.
Integration tests exercise the full path: EPG refresh → series rule
evaluation → Recording creation → post_save signal chain.
"""
from datetime import timedelta
from unittest.mock import patch, MagicMock
from django.test import TestCase
from django.utils import timezone
from apps.channels.models import Channel, Recording
from apps.epg.models import EPGSource, EPGData, ProgramData
from core.models import CoreSettings
def _set_series_rules(rules):
"""Helper to store series rules in CoreSettings."""
CoreSettings.set_dvr_series_rules(rules)
def _set_dvr_offsets(pre_min=0, post_min=0):
"""Helper to store DVR pre/post offsets."""
CoreSettings._update_group("dvr_settings", "DVR Settings", {
"pre_offset_minutes": pre_min,
"post_offset_minutes": post_min,
})
class SeriesRuleDedupBaseTestCase(TestCase):
"""Shared setup for series rule dedup tests."""
def setUp(self):
self.now = timezone.now()
self.epg_source = EPGSource.objects.create(
name="Test EPG", source_type="xmltv"
)
self.epg = EPGData.objects.create(
tvg_id="test.channel.1",
name="Test Channel EPG",
epg_source=self.epg_source,
)
self.channel = Channel.objects.create(
channel_number=1, name="Test Channel", epg_data=self.epg
)
_set_series_rules([{
"tvg_id": "test.channel.1",
"mode": "all",
"title": "Test Show",
}])
_set_dvr_offsets(pre_min=0, post_min=0)
def _create_program(self, hours_from_now=1, title="Test Show",
sub_title="Episode 1", tvg_id="test.channel.1"):
"""Create a ProgramData at the given offset."""
start = self.now + timedelta(hours=hours_from_now)
end = start + timedelta(hours=1)
return ProgramData.objects.create(
epg=self.epg,
tvg_id=tvg_id,
start_time=start,
end_time=end,
title=title,
sub_title=sub_title,
)
def _simulate_epg_refresh(self, programs_data):
"""Delete all ProgramData and recreate with new IDs (simulates EPG refresh)."""
ProgramData.objects.filter(epg=self.epg).delete()
new_programs = []
for data in programs_data:
prog = ProgramData.objects.create(epg=self.epg, **data)
new_programs.append(prog)
return new_programs
def _program_data_for_refresh(self, prog):
"""Build the dict needed by _simulate_epg_refresh from a ProgramData."""
return {
"tvg_id": prog.tvg_id,
"start_time": prog.start_time,
"end_time": prog.end_time,
"title": prog.title,
"sub_title": prog.sub_title,
}
# ---------------------------------------------------------------------------
# Unit tests: dedup logic in evaluate_series_rules_impl
# ---------------------------------------------------------------------------
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class ProgramIdStabilityTests(SeriesRuleDedupBaseTestCase):
"""Verify dedup works after EPG refresh changes ProgramData IDs."""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_no_duplicate_after_epg_refresh(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Same program should not be recorded twice after EPG refresh."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
old_id = prog.id
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
new_programs = self._simulate_epg_refresh(
[self._program_data_for_refresh(prog)]
)
self.assertNotEqual(old_id, new_programs[0].id)
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self.assertEqual(result2["scheduled"], 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_no_duplicate_with_offsets_after_refresh(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Dedup works when DVR offsets shift Recording times away from program times."""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=5, post_min=5)
prog = self._create_program(hours_from_now=2)
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
rec = Recording.objects.first()
self.assertEqual(rec.start_time, prog.start_time - timedelta(minutes=5))
self.assertEqual(rec.end_time, prog.end_time + timedelta(minutes=5))
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_different_episodes_still_recorded(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Different episodes on the same channel should each get a recording."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2, sub_title="Episode 1")
self._create_program(hours_from_now=4, sub_title="Episode 2")
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 2)
self.assertEqual(Recording.objects.count(), 2)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_new_episode_after_refresh_is_recorded(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""A genuinely new episode appearing after EPG refresh should be recorded."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2, sub_title="Episode 1")
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self._simulate_epg_refresh([
self._program_data_for_refresh(prog),
{
"tvg_id": "test.channel.1",
"start_time": prog.end_time,
"end_time": prog.end_time + timedelta(hours=1),
"title": "Test Show",
"sub_title": "Episode 2",
},
])
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2)
self.assertEqual(result2["scheduled"], 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_multiple_epg_refreshes_no_duplicates(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Multiple consecutive EPG refreshes should not accumulate duplicates."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
for _ in range(5):
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class ConcurrencyGuardTests(SeriesRuleDedupBaseTestCase):
"""Verify the task lock prevents concurrent evaluation."""
def test_lock_acquired_and_released(self, mock_schedule, mock_artwork):
"""evaluate_series_rules_impl acquires and releases the task lock."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock", return_value=True) as mock_lock, \
patch("apps.channels.tasks.release_task_lock") as mock_release:
evaluate_series_rules_impl()
mock_lock.assert_called_once_with('evaluate_series_rules', 'all')
mock_release.assert_called_once_with('evaluate_series_rules', 'all')
def test_skips_when_lock_held(self, mock_schedule, mock_artwork):
"""Returns early with skip reason when lock is already held."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock", return_value=False):
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 0)
self.assertTrue(
any(d.get("reason") == "concurrent evaluation in progress"
for d in result["details"]),
)
self.assertEqual(Recording.objects.count(), 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_lock_released_on_exception(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Lock is released even if the inner implementation raises."""
from apps.channels.tasks import evaluate_series_rules_impl
with patch("apps.channels.tasks._evaluate_series_rules_locked",
side_effect=RuntimeError("test error")):
with self.assertRaises(RuntimeError):
evaluate_series_rules_impl()
mock_release.assert_called_once_with('evaluate_series_rules', 'all')
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class SecondaryGuardTests(SeriesRuleDedupBaseTestCase):
"""Verify the secondary DB guard uses stable program attributes."""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_secondary_guard_catches_duplicate_with_offsets(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Secondary guard works with stale program IDs and DVR offsets."""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=10, post_min=10)
prog = self._create_program(hours_from_now=2)
# Pre-existing recording with a stale program ID (from previous EPG refresh)
Recording.objects.create(
channel=self.channel,
start_time=prog.start_time - timedelta(minutes=10),
end_time=prog.end_time + timedelta(minutes=10),
custom_properties={
"program": {
"id": 99999,
"tvg_id": prog.tvg_id,
"title": prog.title,
"start_time": prog.start_time.isoformat(),
"end_time": prog.end_time.isoformat(),
}
},
)
result = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self.assertEqual(result["scheduled"], 0)
# ---------------------------------------------------------------------------
# Integration tests: full path from EPG refresh through recording creation
# ---------------------------------------------------------------------------
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class IntegrationEPGRefreshTests(SeriesRuleDedupBaseTestCase):
"""End-to-end tests simulating the EPG refresh → evaluate → record flow.
These exercise the full signal chain: evaluate_series_rules_impl creates
a Recording, the post_save signal fires schedule_recording_task, and
subsequent evaluations (after EPG refresh) must not create duplicates.
"""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_single_episode_no_duplicates(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Simulate: create rule → evaluate → EPG refresh → re-evaluate.
The full recording lifecycle must result in exactly 1 recording.
"""
from apps.channels.tasks import evaluate_series_rules_impl
# Initial EPG data
prog = self._create_program(hours_from_now=2, sub_title="Pilot")
# First evaluation creates the recording
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
# Verify the recording was created with correct program metadata
rec = Recording.objects.first()
self.assertEqual(rec.custom_properties["program"]["tvg_id"], "test.channel.1")
self.assertEqual(rec.custom_properties["program"]["title"], "Test Show")
self.assertEqual(
rec.custom_properties["program"]["start_time"],
prog.start_time.isoformat()
)
# Verify the post_save signal scheduled a task
mock_schedule.assert_called()
initial_schedule_count = mock_schedule.call_count
# Simulate EPG refresh (programs get new DB IDs)
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
# Re-evaluate after refresh (this is what EPG refresh triggers)
result2 = evaluate_series_rules_impl()
self.assertEqual(result2["scheduled"], 0)
self.assertEqual(Recording.objects.count(), 1)
# No additional task scheduling should have occurred
self.assertEqual(mock_schedule.call_count, initial_schedule_count)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_with_offsets_no_duplicates(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Full flow with DVR offsets: recording times differ from program times."""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=5, post_min=10)
prog = self._create_program(hours_from_now=3, sub_title="Episode 1")
result1 = evaluate_series_rules_impl()
self.assertEqual(result1["scheduled"], 1)
rec = Recording.objects.first()
# Verify offset-adjusted recording times
self.assertEqual(rec.start_time, prog.start_time - timedelta(minutes=5))
self.assertEqual(rec.end_time, prog.end_time + timedelta(minutes=10))
# Verify original (unadjusted) program times in custom_properties
self.assertEqual(
rec.custom_properties["program"]["start_time"],
prog.start_time.isoformat()
)
self.assertEqual(
rec.custom_properties["program"]["end_time"],
prog.end_time.isoformat()
)
# EPG refresh + re-evaluate
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result2 = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
self.assertEqual(result2["scheduled"], 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_multiple_episodes_across_refreshes(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""New episodes appear across multiple EPG refreshes; each recorded once."""
from apps.channels.tasks import evaluate_series_rules_impl
ep1 = self._create_program(hours_from_now=2, sub_title="Episode 1")
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
# EPG refresh adds episode 2 alongside episode 1
ep1_data = self._program_data_for_refresh(ep1)
ep2_start = ep1.end_time
ep2_data = {
"tvg_id": "test.channel.1",
"start_time": ep2_start,
"end_time": ep2_start + timedelta(hours=1),
"title": "Test Show",
"sub_title": "Episode 2",
}
self._simulate_epg_refresh([ep1_data, ep2_data])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2)
# Another EPG refresh adds episode 3
ep3_start = ep2_start + timedelta(hours=1)
ep3_data = {
"tvg_id": "test.channel.1",
"start_time": ep3_start,
"end_time": ep3_start + timedelta(hours=1),
"title": "Test Show",
"sub_title": "Episode 3",
}
self._simulate_epg_refresh([ep1_data, ep2_data, ep3_data])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 3)
# Final EPG refresh with no new episodes — count must stay at 3
self._simulate_epg_refresh([ep1_data, ep2_data, ep3_data])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 3)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_multiple_series_rules(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Multiple series rules on different channels, each evaluated correctly."""
from apps.channels.tasks import evaluate_series_rules_impl
# Second channel with its own EPG
epg2 = EPGData.objects.create(
tvg_id="test.channel.2",
name="Channel 2 EPG",
epg_source=self.epg_source,
)
channel2 = Channel.objects.create(
channel_number=2, name="Test Channel 2", epg_data=epg2
)
_set_series_rules([
{"tvg_id": "test.channel.1", "mode": "all", "title": "Show A"},
{"tvg_id": "test.channel.2", "mode": "all", "title": "Show B"},
])
# Programs on both channels
start1 = self.now + timedelta(hours=2)
prog1 = ProgramData.objects.create(
epg=self.epg, tvg_id="test.channel.1",
start_time=start1, end_time=start1 + timedelta(hours=1),
title="Show A", sub_title="Episode 1",
)
start2 = self.now + timedelta(hours=3)
prog2 = ProgramData.objects.create(
epg=epg2, tvg_id="test.channel.2",
start_time=start2, end_time=start2 + timedelta(hours=1),
title="Show B", sub_title="Episode 1",
)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2)
self.assertEqual(Recording.objects.filter(channel=self.channel).count(), 1)
self.assertEqual(Recording.objects.filter(channel=channel2).count(), 1)
# EPG refresh for both channels
ProgramData.objects.filter(epg=self.epg).delete()
ProgramData.objects.filter(epg=epg2).delete()
ProgramData.objects.create(
epg=self.epg, tvg_id="test.channel.1",
start_time=start1, end_time=start1 + timedelta(hours=1),
title="Show A", sub_title="Episode 1",
)
ProgramData.objects.create(
epg=epg2, tvg_id="test.channel.2",
start_time=start2, end_time=start2 + timedelta(hours=1),
title="Show B", sub_title="Episode 1",
)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 2,
"No duplicates across multiple series rules after EPG refresh")
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_rapid_epg_refreshes_simulate_user_report(
self, mock_release, mock_lock, mock_schedule, mock_artwork
):
"""Reproduce the user-reported scenario: series rule + multiple EPG refreshes
causing count to balloon from 6 to 25 and 5 simultaneous recordings.
Simulates 6 episodes with 5 EPG refreshes (each assigning new ProgramData IDs).
"""
from apps.channels.tasks import evaluate_series_rules_impl
# Create 6 episodes (the user had "next of 6")
episodes = []
for i in range(6):
start = self.now + timedelta(hours=2 + i * 2)
episodes.append({
"tvg_id": "test.channel.1",
"start_time": start,
"end_time": start + timedelta(hours=1),
"title": "Test Show",
"sub_title": f"Episode {i + 1}",
})
# Create initial ProgramData
for ep in episodes:
ProgramData.objects.create(epg=self.epg, **ep)
# First evaluation: should create exactly 6 recordings
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 6)
# Simulate 5 EPG refreshes (the user saw count balloon to 25)
for refresh_num in range(5):
self._simulate_epg_refresh(episodes)
result = evaluate_series_rules_impl()
self.assertEqual(
Recording.objects.count(), 6,
f"After EPG refresh #{refresh_num + 1}, expected 6 recordings "
f"but got {Recording.objects.count()}"
)
self.assertEqual(result["scheduled"], 0)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_recording_survives_program_removal_and_readd(
self, mock_release, mock_lock, mock_schedule, mock_artwork
):
"""Program temporarily disappears from EPG then reappears — no duplicate."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2, sub_title="Episode 1")
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
# EPG refresh removes the program entirely
self._simulate_epg_refresh([])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"Existing recording preserved when program disappears from EPG")
# EPG refresh adds the program back (new ID)
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"No duplicate when program reappears with new ID")
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_celery_task_wrapper_calls_impl(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""The @shared_task evaluate_series_rules delegates to _impl correctly."""
from apps.channels.tasks import evaluate_series_rules
self._create_program(hours_from_now=2)
result = evaluate_series_rules()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
# Call again (simulating a second EPG refresh trigger)
result2 = evaluate_series_rules()
self.assertEqual(result2["scheduled"], 0)
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_tvg_id_scoped_evaluation(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Scoped evaluation (tvg_id parameter) still prevents duplicates."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
result1 = evaluate_series_rules_impl(tvg_id="test.channel.1")
self.assertEqual(result1["scheduled"], 1)
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result2 = evaluate_series_rules_impl(tvg_id="test.channel.1")
self.assertEqual(result2["scheduled"], 0)
self.assertEqual(Recording.objects.count(), 1)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_full_flow_offset_change_between_refreshes(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Changing DVR offsets between EPG refreshes doesn't create duplicates.
Even though Recording.start_time/end_time change when offsets change,
the dedup key uses the original program times from custom_properties.
"""
from apps.channels.tasks import evaluate_series_rules_impl
_set_dvr_offsets(pre_min=5, post_min=5)
prog = self._create_program(hours_from_now=2)
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
rec = Recording.objects.first()
original_start = rec.start_time
original_end = rec.end_time
# Change offsets
_set_dvr_offsets(pre_min=10, post_min=15)
# EPG refresh
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
result = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"Changing offsets between refreshes should not create duplicates")
self.assertEqual(result["scheduled"], 0)
# ---------------------------------------------------------------------------
# Edge case tests: Redis unavailability, non-series recordings, robustness
# ---------------------------------------------------------------------------
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class RedisUnavailabilityTests(SeriesRuleDedupBaseTestCase):
"""Verify evaluation works when Redis is unavailable (lock cannot be acquired)."""
def test_proceeds_when_redis_down(self, mock_schedule, mock_artwork):
"""Evaluation succeeds (with dedup guards) when Redis raises on lock acquire."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")):
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 1)
def test_dedup_still_works_without_lock(self, mock_schedule, mock_artwork):
"""Dedup guards prevent duplicates even when the lock is unavailable."""
from apps.channels.tasks import evaluate_series_rules_impl
prog = self._create_program(hours_from_now=2)
# First call: Redis down, proceeds without lock
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")):
evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1)
# EPG refresh
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
# Second call: Redis still down
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")):
result = evaluate_series_rules_impl()
self.assertEqual(Recording.objects.count(), 1,
"Dedup guards prevent duplicates even without lock")
self.assertEqual(result["scheduled"], 0)
def test_lock_not_released_when_not_acquired(self, mock_schedule, mock_artwork):
"""release_task_lock is not called if acquire raised an exception."""
from apps.channels.tasks import evaluate_series_rules_impl
self._create_program(hours_from_now=2)
with patch("apps.channels.tasks.acquire_task_lock",
side_effect=ConnectionError("Redis unavailable")), \
patch("apps.channels.tasks.release_task_lock") as mock_release:
evaluate_series_rules_impl()
mock_release.assert_not_called()
@patch("apps.channels.tasks.prefetch_recording_artwork")
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
class NonSeriesRecordingTests(SeriesRuleDedupBaseTestCase):
"""Verify non-series recordings don't interfere with series rule dedup."""
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_manual_recording_without_program_data_ignored(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Recordings without custom_properties.program are skipped by dedup key builder."""
from apps.channels.tasks import evaluate_series_rules_impl
# Manual recording with no program metadata
Recording.objects.create(
channel=self.channel,
start_time=self.now + timedelta(hours=2),
end_time=self.now + timedelta(hours=3),
custom_properties={},
)
prog = self._create_program(hours_from_now=2)
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 2)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_recurring_rule_recording_does_not_interfere(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Recordings from recurring rules (custom_properties.rule) don't block series rules."""
from apps.channels.tasks import evaluate_series_rules_impl
Recording.objects.create(
channel=self.channel,
start_time=self.now + timedelta(hours=2),
end_time=self.now + timedelta(hours=3),
custom_properties={"rule": {"id": 1, "name": "Daily News"}},
)
prog = self._create_program(hours_from_now=2)
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
self.assertEqual(Recording.objects.count(), 2)
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
@patch("apps.channels.tasks.release_task_lock")
def test_recording_with_null_custom_properties_ignored(self, mock_release, mock_lock,
mock_schedule, mock_artwork):
"""Recordings with None custom_properties don't crash the dedup key builder."""
from apps.channels.tasks import evaluate_series_rules_impl
Recording.objects.create(
channel=self.channel,
start_time=self.now + timedelta(hours=2),
end_time=self.now + timedelta(hours=3),
custom_properties=None,
)
prog = self._create_program(hours_from_now=2)
result = evaluate_series_rules_impl()
self.assertEqual(result["scheduled"], 1)
@@ -0,0 +1,385 @@
"""Tests for ghost client detection and cleanup.
Covers:
- ClientManager.remove_ghost_clients() pipelined EXISTS logic
- channel_status detailed stats path removes ghost clients from Redis SET
- channel_status basic stats path removes ghost clients and corrects count
- _check_orphaned_metadata() validates client SET entries and cleans up
channels where all clients are ghosts
"""
from unittest.mock import MagicMock, patch, PropertyMock
from django.test import TestCase
from apps.proxy.ts_proxy.client_manager import ClientManager
from apps.proxy.ts_proxy.constants import ChannelMetadataField, ChannelState
from apps.proxy.ts_proxy.redis_keys import RedisKeys
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
CHANNEL_ID = "00000000-0000-0000-0000-000000000001"
def _make_proxy_server(redis_client=None):
"""Create a minimal mock ProxyServer with a redis_client."""
server = MagicMock()
server.redis_client = redis_client or MagicMock()
server.stream_managers = {}
server.client_managers = {}
server.worker_id = "test-worker-1"
return server
def _metadata_for_channel(state="active"):
"""Return a plausible channel metadata dict (bytes keys/values)."""
return {
ChannelMetadataField.STATE.encode(): state.encode(),
ChannelMetadataField.URL.encode(): b"http://example.com/stream",
ChannelMetadataField.STREAM_PROFILE.encode(): b"default",
ChannelMetadataField.OWNER.encode(): b"test-worker-1",
ChannelMetadataField.INIT_TIME.encode(): b"1773500000.0",
}
# ---------------------------------------------------------------------------
# Unit tests for ClientManager.remove_ghost_clients()
# ---------------------------------------------------------------------------
class RemoveGhostClientsTests(TestCase):
"""Directly exercises the static method that all callers rely on."""
def test_ghost_removed_and_returned(self):
"""Client ID in SET with no metadata hash should be SREM'd."""
redis = MagicMock()
redis.smembers.return_value = {b"ghost_001"}
pipe = MagicMock()
redis.pipeline.return_value = pipe
pipe.execute.return_value = [False] # EXISTS → False
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(result, [b"ghost_001"])
redis.srem.assert_called_once()
def test_live_client_preserved(self):
"""Client with valid metadata hash should NOT be removed."""
redis = MagicMock()
redis.smembers.return_value = {b"live_001"}
pipe = MagicMock()
redis.pipeline.return_value = pipe
pipe.execute.return_value = [True] # EXISTS → True
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(result, [])
redis.srem.assert_not_called()
def test_mixed_ghost_and_live(self):
"""Only ghost clients should be removed; live ones preserved."""
redis = MagicMock()
redis.smembers.return_value = {b"ghost_001", b"live_001"}
pipe = MagicMock()
redis.pipeline.return_value = pipe
# Order matches list(smembers), which is non-deterministic —
# map both IDs so the test is stable regardless of iteration order.
client_id_list = list(redis.smembers.return_value)
def exists_results():
return [
b"ghost_001" not in cid.decode() == False
for cid in client_id_list
]
# Simpler: mock based on key content
def pipe_exists(key):
pass # just enqueued; results come from execute()
pipe.exists.side_effect = pipe_exists
pipe.execute.return_value = [
"live" in cid.decode() for cid in client_id_list
]
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(len(result), 1)
self.assertTrue(any(b"ghost" in cid for cid in result))
redis.srem.assert_called_once()
def test_empty_set_returns_empty(self):
"""No clients means nothing to clean."""
redis = MagicMock()
redis.smembers.return_value = set()
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
self.assertEqual(result, [])
redis.pipeline.assert_not_called()
def test_pre_fetched_client_ids_skips_smembers(self):
"""When client_ids is passed, SMEMBERS should not be called."""
redis = MagicMock()
pipe = MagicMock()
redis.pipeline.return_value = pipe
pipe.execute.return_value = [False]
pre_fetched = {b"ghost_001"}
result = ClientManager.remove_ghost_clients(
redis, CHANNEL_ID, client_ids=pre_fetched
)
redis.smembers.assert_not_called()
self.assertEqual(len(result), 1)
# ---------------------------------------------------------------------------
# Detailed stats path: exercises get_detailed_channel_info()
# ---------------------------------------------------------------------------
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
class DetailedStatsGhostClientTests(TestCase):
"""get_detailed_channel_info() should remove ghost clients whose metadata
hash has expired from the Redis client SET."""
def _setup_redis(self, mock_proxy_cls, client_ids, hgetall_side_effect):
"""Wire up a mock ProxyServer with controlled Redis responses."""
redis = MagicMock()
server = _make_proxy_server(redis)
mock_proxy_cls.get_instance.return_value = server
redis.hgetall.side_effect = hgetall_side_effect
redis.smembers.return_value = client_ids
# buffer_index, ttl, exists all need safe defaults
redis.get.return_value = b"10"
redis.ttl.return_value = 300
redis.exists.return_value = True
return redis
def test_ghost_client_removed_from_set(self, mock_proxy_cls):
"""Ghost client should be SREM'd and excluded from result."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
def hgetall_side_effect(key):
if "clients:" in key:
return {} # ghost — metadata expired
return _metadata_for_channel()
redis = self._setup_redis(
mock_proxy_cls, {b"ghost_001"}, hgetall_side_effect
)
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
self.assertEqual(result['client_count'], 0)
self.assertEqual(len(result['clients']), 0)
redis.srem.assert_called_once()
def test_live_client_preserved(self, mock_proxy_cls):
"""Client with valid metadata should appear in results."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
def hgetall_side_effect(key):
if "clients:" in key:
return {
b'user_agent': b'VLC/3.0',
b'worker_id': b'test-worker-1',
b'connected_at': b'1773500000.0',
}
return _metadata_for_channel()
redis = self._setup_redis(
mock_proxy_cls, {b"live_001"}, hgetall_side_effect
)
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
self.assertEqual(result['client_count'], 1)
self.assertEqual(len(result['clients']), 1)
redis.srem.assert_not_called()
def test_mixed_ghost_and_live(self, mock_proxy_cls):
"""Only ghost clients should be removed; live ones preserved."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
def hgetall_side_effect(key):
if "clients:" in key:
if "ghost" in key:
return {}
return {
b'user_agent': b'VLC/3.0',
b'worker_id': b'test-worker-1',
}
return _metadata_for_channel()
redis = self._setup_redis(
mock_proxy_cls, {b"ghost_001", b"live_001"}, hgetall_side_effect
)
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
self.assertEqual(result['client_count'], 1)
self.assertEqual(len(result['clients']), 1)
redis.srem.assert_called_once()
# ---------------------------------------------------------------------------
# Basic stats path: exercises get_basic_channel_info()
# ---------------------------------------------------------------------------
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
class BasicStatsGhostClientTests(TestCase):
"""get_basic_channel_info() should call remove_ghost_clients(), skip
ghosts from display, and correct client_count."""
def _setup_redis(self, mock_proxy_cls, client_ids, ghost_ids):
"""Wire up mock ProxyServer. ghost_ids controls which EXISTS return False."""
redis = MagicMock()
server = _make_proxy_server(redis)
mock_proxy_cls.get_instance.return_value = server
redis.hgetall.return_value = _metadata_for_channel()
redis.get.return_value = b"10" # buffer_index
redis.scard.return_value = len(client_ids)
redis.smembers.return_value = client_ids
redis.hget.return_value = None # individual field lookups
# Pipeline for remove_ghost_clients
pipe = MagicMock()
redis.pipeline.return_value = pipe
client_id_list = list(client_ids)
pipe.execute.return_value = [
cid not in ghost_ids for cid in client_id_list
]
return redis
def test_ghost_removed_and_count_corrected(self, mock_proxy_cls):
"""Ghost client should be cleaned and client_count decremented."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
redis = self._setup_redis(
mock_proxy_cls,
client_ids={b"ghost_001"},
ghost_ids={b"ghost_001"},
)
result = ChannelStatus.get_basic_channel_info(CHANNEL_ID)
self.assertIsNotNone(result)
self.assertEqual(result['client_count'], 0)
redis.srem.assert_called_once()
def test_live_client_count_preserved(self, mock_proxy_cls):
"""Live clients should be counted correctly."""
from apps.proxy.ts_proxy.channel_status import ChannelStatus
redis = self._setup_redis(
mock_proxy_cls,
client_ids={b"live_001"},
ghost_ids=set(),
)
result = ChannelStatus.get_basic_channel_info(CHANNEL_ID)
self.assertIsNotNone(result)
self.assertEqual(result['client_count'], 1)
redis.srem.assert_not_called()
# ---------------------------------------------------------------------------
# Orphaned channel cleanup: exercises _check_orphaned_metadata()
# ---------------------------------------------------------------------------
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
class OrphanedChannelGhostValidationTests(TestCase):
"""_check_orphaned_metadata() should validate client SET entries when
owner is dead and client_count > 0. If all clients are ghosts, it
should clean up the channel."""
def _make_server_for_orphan_check(self, mock_proxy_cls, channel_id,
client_ids, ghost_ids, owner="dead-worker"):
"""Build a mock ProxyServer whose Redis state simulates an orphaned channel."""
redis = MagicMock()
server = _make_proxy_server(redis)
mock_proxy_cls.get_instance.return_value = server
metadata_key = RedisKeys.channel_metadata(channel_id)
metadata = _metadata_for_channel()
metadata[ChannelMetadataField.OWNER.encode()] = owner.encode()
# scan returns the one channel metadata key
redis.scan.return_value = (0, [metadata_key.encode()])
redis.hgetall.return_value = metadata
redis.scard.return_value = len(client_ids)
redis.smembers.return_value = client_ids
# Owner heartbeat is dead
redis.exists.side_effect = lambda key: (
False if "heartbeat" in key else True
)
# Pipeline for remove_ghost_clients
pipe = MagicMock()
redis.pipeline.return_value = pipe
client_id_list = list(client_ids)
pipe.execute.return_value = [
cid not in ghost_ids for cid in client_id_list
]
return server, redis
def test_all_ghosts_triggers_cleanup(self, mock_proxy_cls):
"""When all clients are ghosts, channel should be cleaned up."""
from apps.proxy.ts_proxy.server import ProxyServer
channel_id = "00000000-0000-0000-0000-000000000005"
server, redis = self._make_server_for_orphan_check(
mock_proxy_cls, channel_id,
client_ids={b"ghost_001", b"ghost_002"},
ghost_ids={b"ghost_001", b"ghost_002"},
)
# Call the real method on a real-ish ProxyServer
# The method lives on the server instance, so invoke it directly.
# We need to call _check_orphaned_metadata on the actual server mock,
# but it's a MagicMock. Instead, test via remove_ghost_clients directly
# and verify the cleanup decision logic.
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
real_count = max(0, len({b"ghost_001", b"ghost_002"}) - len(stale_ids))
self.assertEqual(len(stale_ids), 2)
self.assertEqual(real_count, 0)
redis.srem.assert_called_once()
def test_mixed_preserves_live_clients(self, mock_proxy_cls):
"""When some clients are live, real_count should be > 0."""
channel_id = "00000000-0000-0000-0000-000000000006"
server, redis = self._make_server_for_orphan_check(
mock_proxy_cls, channel_id,
client_ids={b"ghost_001", b"live_001"},
ghost_ids={b"ghost_001"},
)
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
real_count = max(0, 2 - len(stale_ids))
self.assertEqual(len(stale_ids), 1)
self.assertEqual(real_count, 1)
def test_no_ghosts_no_cleanup(self, mock_proxy_cls):
"""When all clients are live, no SREM should be called."""
channel_id = "00000000-0000-0000-0000-000000000007"
server, redis = self._make_server_for_orphan_check(
mock_proxy_cls, channel_id,
client_ids={b"live_001"},
ghost_ids=set(),
)
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
self.assertEqual(len(stale_ids), 0)
redis.srem.assert_not_called()
@@ -0,0 +1,231 @@
"""Tests for stuck INITIALIZING state fix.
Covers:
- stream_manager.run() finally block: ownership check + state guard fallback
- ChannelState.PRE_ACTIVE contains the correct states
- INITIALIZING is included in the cleanup task grace period check
"""
import time
import threading
from unittest.mock import MagicMock, patch
from django.test import TestCase
from apps.proxy.ts_proxy.constants import ChannelMetadataField, ChannelState
from apps.proxy.ts_proxy.redis_keys import RedisKeys
from apps.proxy.ts_proxy.stream_manager import StreamManager
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
CHANNEL_ID = "00000000-0000-0000-0000-000000000001"
def _make_stream_manager(tried_stream_ids=None, max_retries=3):
"""Build a StreamManager via __new__ (bypasses __init__) with the
minimum attributes required by the run() finally block."""
sm = StreamManager.__new__(StreamManager)
sm.channel_id = CHANNEL_ID
sm.worker_id = "worker-1"
sm.max_retries = max_retries
sm.tried_stream_ids = tried_stream_ids if tried_stream_ids is not None else set()
sm.running = False # while-loop exits immediately
sm.connected = False
sm.transcode_process_active = False
sm._buffer_check_timers = []
sm.url = "http://example.com/stream"
sm.url_switching = False
sm.url_switch_start_time = 0
sm.url_switch_timeout = 30
sm.stop_requested = False
sm.stopping = False
sm.socket = None
sm.transcode_process = None
sm.current_response = None
sm.current_session = None
sm.current_stream_id = None
buffer = MagicMock()
buffer.redis_client = MagicMock()
buffer.channel_id = CHANNEL_ID
sm.buffer = buffer
return sm
def _run_finally_block(sm, owner_value, current_state):
"""Invoke StreamManager.run() so its finally block executes against real code.
Patches threading.Thread and ConfigHelper so the try-block is inert
(self.running=False makes the while-loop exit immediately).
Returns True if the finally block wrote ERROR to Redis.
"""
redis = sm.buffer.redis_client
# Mock the owner key GET — the finally block calls redis.get(owner_key)
def get_side_effect(key):
if "owner" in key:
return owner_value
return None
redis.get.side_effect = get_side_effect
# Mock hget for state field lookup in the PRE_ACTIVE guard
if current_state is not None:
redis.hget.return_value = current_state.encode('utf-8')
else:
redis.hget.return_value = None
# Reset hset so we can detect whether ERROR was written
redis.hset.reset_mock()
redis.setex.reset_mock()
with patch.object(threading, 'Thread', return_value=MagicMock()):
with patch('apps.proxy.ts_proxy.stream_manager.ConfigHelper') as mock_cfg:
mock_cfg.max_stream_switches.return_value = 0
mock_cfg.max_retries.return_value = sm.max_retries
sm.run()
# Check if hset was called with ERROR state
if redis.hset.called:
mapping = redis.hset.call_args[1].get('mapping', {})
return mapping.get(ChannelMetadataField.STATE) == ChannelState.ERROR
return False
# ---------------------------------------------------------------------------
# stream_manager.run() finally block: ownership + state guard behavior
# ---------------------------------------------------------------------------
class StreamManagerFinallyBlockTests(TestCase):
"""The run() finally block writes ERROR if the worker is still the owner
(normal case) OR if ownership expired and the channel is still in a
pre-active state (no new owner has taken over)."""
# --- Owner still valid: always write ERROR ---
def test_owner_writes_error_regardless_of_state(self):
"""When we're still the owner, always write ERROR."""
sm = _make_stream_manager()
owner = sm.worker_id.encode('utf-8')
self.assertTrue(_run_finally_block(sm, owner, ChannelState.ACTIVE))
def test_owner_writes_error_on_initializing(self):
"""Owner + INITIALIZING = write ERROR."""
sm = _make_stream_manager()
owner = sm.worker_id.encode('utf-8')
self.assertTrue(_run_finally_block(sm, owner, ChannelState.INITIALIZING))
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
self.assertEqual(mapping[ChannelMetadataField.STATE], ChannelState.ERROR)
# --- Ownership expired, no new owner: use state guard ---
def test_no_owner_initializing_writes_error(self):
"""Ownership expired + INITIALIZING = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.INITIALIZING))
def test_no_owner_connecting_writes_error(self):
"""Ownership expired + CONNECTING = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.CONNECTING))
def test_no_owner_buffering_writes_error(self):
"""Ownership expired + BUFFERING = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.BUFFERING))
def test_no_owner_waiting_for_clients_writes_error(self):
"""Ownership expired + WAITING_FOR_CLIENTS = write ERROR."""
sm = _make_stream_manager()
self.assertTrue(_run_finally_block(sm, None, ChannelState.WAITING_FOR_CLIENTS))
def test_no_owner_active_does_not_write(self):
"""Ownership expired + ACTIVE = do NOT write ERROR."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, None, ChannelState.ACTIVE))
def test_no_owner_error_does_not_write(self):
"""Ownership expired + already ERROR = do NOT write again."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, None, ChannelState.ERROR))
def test_no_owner_no_state_does_not_write(self):
"""Ownership expired + no state metadata = do NOT write."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, None, None))
# --- New owner took over: never clobber ---
def test_new_owner_initializing_does_not_write(self):
"""Another worker owns the channel — do NOT clobber."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, b"other-worker", ChannelState.INITIALIZING))
def test_new_owner_active_does_not_write(self):
"""Another worker owns the channel and is ACTIVE — do NOT write."""
sm = _make_stream_manager()
self.assertFalse(_run_finally_block(sm, b"other-worker", ChannelState.ACTIVE))
# --- Stopping key and error messages ---
def test_stopping_key_set_on_error_update(self):
"""When ERROR is written, stopping key must also be set."""
sm = _make_stream_manager()
_run_finally_block(sm, None, ChannelState.INITIALIZING)
sm.buffer.redis_client.setex.assert_called_once()
args = sm.buffer.redis_client.setex.call_args[0]
self.assertIn("stopping", args[0])
self.assertEqual(args[1], 60)
def test_error_message_includes_stream_count(self):
"""When multiple streams were tried, error message reflects that."""
sm = _make_stream_manager(tried_stream_ids={1, 2, 3})
_run_finally_block(sm, None, ChannelState.INITIALIZING)
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
error_msg = mapping[ChannelMetadataField.ERROR_MESSAGE]
self.assertIn("3 stream options failed", error_msg)
def test_error_message_with_no_streams_tried(self):
"""When no alternate streams were tried, shows retry count."""
sm = _make_stream_manager(tried_stream_ids=set(), max_retries=5)
_run_finally_block(sm, None, ChannelState.INITIALIZING)
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
error_msg = mapping[ChannelMetadataField.ERROR_MESSAGE]
self.assertIn("5", error_msg)
# ---------------------------------------------------------------------------
# ChannelState.PRE_ACTIVE: verify contents and immutability
# ---------------------------------------------------------------------------
class PreActiveStateTests(TestCase):
"""Verify PRE_ACTIVE contains the correct states and is immutable."""
def test_initializing_in_pre_active(self):
self.assertIn(ChannelState.INITIALIZING, ChannelState.PRE_ACTIVE)
def test_connecting_in_pre_active(self):
self.assertIn(ChannelState.CONNECTING, ChannelState.PRE_ACTIVE)
def test_buffering_in_pre_active(self):
self.assertIn(ChannelState.BUFFERING, ChannelState.PRE_ACTIVE)
def test_waiting_for_clients_in_pre_active(self):
self.assertIn(ChannelState.WAITING_FOR_CLIENTS, ChannelState.PRE_ACTIVE)
def test_active_not_in_pre_active(self):
self.assertNotIn(ChannelState.ACTIVE, ChannelState.PRE_ACTIVE)
def test_error_not_in_pre_active(self):
self.assertNotIn(ChannelState.ERROR, ChannelState.PRE_ACTIVE)
def test_pre_active_is_frozenset(self):
self.assertIsInstance(ChannelState.PRE_ACTIVE, frozenset)
@@ -0,0 +1,331 @@
"""Tests for ts_proxy keepalive and stats-update behavior.
Covers:
- stream_generator._should_send_keepalive() owner vs non-owner worker paths
- stream_generator._should_send_keepalive() Redis last_data health check
- client_manager._do_stats_update() error handling and WebSocket dispatch
- client_manager.remove_client() non-blocking stats update
- Keepalive/DVR-timeout timing invariants
"""
import threading
import time
from unittest.mock import MagicMock, patch
from django.test import TestCase
# ---------------------------------------------------------------------------
# _should_send_keepalive: owner worker path
# ---------------------------------------------------------------------------
class OwnerWorkerKeepaliveTests(TestCase):
"""Owner worker has a stream_manager; keepalive logic uses it directly."""
def _make_generator(self, healthy, at_buffer_head, consecutive_empty):
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
gen = StreamGenerator.__new__(StreamGenerator)
gen.channel_id = "00000000-0000-0000-0000-000000000001"
gen.client_id = "test-client"
buffer = MagicMock()
buffer.index = 10 if at_buffer_head else 100
gen.local_index = 10
gen.buffer = buffer
stream_manager = MagicMock()
stream_manager.healthy = healthy
gen.stream_manager = stream_manager
gen.consecutive_empty = consecutive_empty
return gen
def test_owner_healthy_returns_false(self):
"""Owner worker, healthy stream -> no keepalive."""
gen = self._make_generator(healthy=True, at_buffer_head=True, consecutive_empty=10)
self.assertFalse(gen._should_send_keepalive(gen.local_index))
def test_owner_unhealthy_at_head_returns_true(self):
"""Owner worker, unhealthy stream, at buffer head -> send keepalive."""
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=10)
self.assertTrue(gen._should_send_keepalive(gen.local_index))
def test_owner_unhealthy_not_at_head_returns_false(self):
"""Owner worker, unhealthy stream, but NOT at buffer head -> no keepalive."""
gen = self._make_generator(healthy=False, at_buffer_head=False, consecutive_empty=10)
self.assertFalse(gen._should_send_keepalive(gen.local_index))
def test_owner_insufficient_consecutive_empty_returns_false(self):
"""Owner worker, unhealthy, at head but consecutive_empty < 5 -> no keepalive."""
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=3)
self.assertFalse(gen._should_send_keepalive(gen.local_index))
def test_owner_exactly_5_consecutive_empty_returns_true(self):
"""consecutive_empty == 5 is the minimum threshold."""
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=5)
self.assertTrue(gen._should_send_keepalive(gen.local_index))
# ---------------------------------------------------------------------------
# _should_send_keepalive: non-owner worker path
# ---------------------------------------------------------------------------
class NonOwnerWorkerKeepaliveTests(TestCase):
"""Non-owner worker has stream_manager=None; health determined from Redis."""
def _make_generator(self, consecutive_empty=10):
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
gen = StreamGenerator.__new__(StreamGenerator)
gen.channel_id = "00000000-0000-0000-0000-000000000002"
gen.client_id = "test-client-nonowner"
buffer = MagicMock()
buffer.index = 10
gen.local_index = 10
gen.buffer = buffer
gen.stream_manager = None # non-owner worker
gen.consecutive_empty = consecutive_empty
# Attributes added by health-check throttling (set in __init__)
gen._last_health_check_time = 0.0
gen._last_health_check_result = False
gen._health_check_interval = 2.0
gen.proxy_server = None
return gen
def _mock_proxy_server(self, last_data_value):
"""Return a mock ProxyServer with a redis_client pre-configured."""
server = MagicMock()
redis_client = MagicMock()
server.redis_client = redis_client
redis_client.get.return_value = last_data_value
return server
def test_non_owner_fresh_data_returns_false(self):
"""Non-owner, last_data < 10s ago -> stream healthy -> no keepalive."""
gen = self._make_generator()
fresh_ts = str(time.time() - 2.0).encode()
server = self._mock_proxy_server(fresh_ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result, "Fresh data should NOT trigger keepalive")
def test_non_owner_stale_data_returns_true(self):
"""Non-owner, last_data >= 10s ago -> stream unhealthy -> send keepalive."""
gen = self._make_generator()
stale_ts = str(time.time() - 12.0).encode()
server = self._mock_proxy_server(stale_ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertTrue(result, "Stale data (12s) should trigger keepalive")
def test_non_owner_exactly_at_timeout_returns_true(self):
"""Data age exactly equal to CONNECTION_TIMEOUT (10s) -> send keepalive."""
gen = self._make_generator()
ts = str(time.time() - 10.0).encode()
server = self._mock_proxy_server(ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertTrue(result, "Data at exactly timeout threshold should trigger keepalive")
def test_non_owner_no_redis_key_returns_true(self):
"""Non-owner, last_data key missing from Redis -> assume unhealthy."""
gen = self._make_generator()
server = self._mock_proxy_server(None)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertTrue(result, "Missing last_data key should trigger keepalive")
def test_non_owner_redis_client_none_returns_false(self):
"""Non-owner, redis_client is None (disconnected) -> conservative, no keepalive."""
gen = self._make_generator()
server = MagicMock()
server.redis_client = None
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result, "No redis_client -> conservative, no keepalive")
def test_non_owner_redis_exception_returns_false(self):
"""Non-owner, Redis raises an exception -> conservative, no keepalive."""
gen = self._make_generator()
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.side_effect = Exception("Redis error")
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result, "Redis error -> conservative, no keepalive")
def test_non_owner_not_at_buffer_head_returns_false(self):
"""Non-owner, NOT at buffer head -> no keepalive regardless of Redis."""
gen = self._make_generator()
gen.buffer.index = 100 # far ahead of local_index=10
server = self._mock_proxy_server(None)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result)
def test_non_owner_insufficient_consecutive_empty_returns_false(self):
"""Non-owner, at head, but consecutive_empty < 5 -> no keepalive."""
gen = self._make_generator(consecutive_empty=2)
stale_ts = str(time.time() - 30.0).encode()
server = self._mock_proxy_server(stale_ts)
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
MockPS.get_instance.return_value = server
result = gen._should_send_keepalive(gen.local_index)
self.assertFalse(result)
# ---------------------------------------------------------------------------
# _do_stats_update: error handling and WebSocket dispatch
# ---------------------------------------------------------------------------
class DoStatsUpdateTests(TestCase):
"""_do_stats_update runs the actual Redis scan + WebSocket call."""
def _make_client_manager(self):
from apps.proxy.ts_proxy.client_manager import ClientManager
cm = ClientManager.__new__(ClientManager)
cm.channel_id = "00000000-0000-0000-0000-000000000004"
cm._heartbeat_running = False
return cm
def test_do_stats_update_calls_send_websocket_update(self):
"""_do_stats_update must call send_websocket_update with channel_stats."""
cm = self._make_client_manager()
mock_redis = MagicMock()
mock_redis.scan.return_value = (0, [])
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update") as mock_ws, \
patch("redis.Redis.from_url", return_value=mock_redis):
cm._do_stats_update()
mock_ws.assert_called_once()
event_type = mock_ws.call_args[0][1]
self.assertEqual(event_type, "update")
payload = mock_ws.call_args[0][2]
self.assertEqual(payload["type"], "channel_stats")
def test_do_stats_update_does_not_raise_on_redis_error(self):
"""Redis failure must be swallowed (logged), not propagated."""
cm = self._make_client_manager()
with patch("redis.Redis.from_url", side_effect=Exception("Redis down")):
try:
cm._do_stats_update()
except Exception as e:
self.fail(f"_do_stats_update raised an exception: {e}")
def test_do_stats_update_scans_channel_client_keys(self):
"""Must scan for ts_proxy:channel:*:clients pattern."""
cm = self._make_client_manager()
mock_redis = MagicMock()
mock_redis.scan.return_value = (0, [])
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update"), \
patch("redis.Redis.from_url", return_value=mock_redis):
cm._do_stats_update()
scan_call = mock_redis.scan.call_args
self.assertIn("ts_proxy:channel:*:clients", str(scan_call))
# ---------------------------------------------------------------------------
# Integration: remove_client must not block on WebSocket
# ---------------------------------------------------------------------------
class ClientRemoveIntegrationTests(TestCase):
"""When remove_client() fires, _trigger_stats_update must not block."""
def test_remove_client_does_not_block_on_websocket(self):
"""remove_client() must return quickly even if WebSocket is slow."""
from apps.proxy.ts_proxy.client_manager import ClientManager
cm = ClientManager.__new__(ClientManager)
cm.channel_id = "00000000-0000-0000-0000-000000000005"
cm._heartbeat_running = False
cm.clients = {"test-client-1"}
cm.last_heartbeat_time = {"test-client-1": time.time()}
cm.last_active_time = time.time()
cm.client_set_key = f"ts_proxy:channel:{cm.channel_id}:clients"
cm.client_ttl = 60
cm.worker_id = "worker-1"
cm.proxy_server = MagicMock()
cm.proxy_server.am_i_owner.return_value = False
cm.lock = threading.Lock()
mock_redis = MagicMock()
mock_redis.hgetall.return_value = {b"ip_address": b"127.0.0.1"}
mock_redis.scard.return_value = 1
cm.redis_client = mock_redis
slow_ws_called = threading.Event()
def slow_websocket(*args, **kwargs):
time.sleep(2.0)
slow_ws_called.set()
start = time.time()
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update", side_effect=slow_websocket):
cm.remove_client("test-client-1")
elapsed = time.time() - start
self.assertLess(elapsed, 1.0,
f"remove_client() blocked for {elapsed:.2f}s waiting for WebSocket "
f"(should dispatch to background thread and return immediately)")
# ---------------------------------------------------------------------------
# DVR timeout threshold vs keepalive timing
# ---------------------------------------------------------------------------
class KeepaliveTimingTests(TestCase):
"""Verify that keepalive threshold gives sufficient margin before DVR timeout."""
def test_keepalive_threshold_less_than_dvr_timeout(self):
"""CONNECTION_TIMEOUT (keepalive trigger) must be < DVR read timeout (15s)."""
from apps.proxy.config import TSConfig as Config
connection_timeout = getattr(Config, "CONNECTION_TIMEOUT", 10)
dvr_read_timeout = 15 # hard-coded in run_recording: timeout=(10, 15)
self.assertLess(
connection_timeout,
dvr_read_timeout,
f"CONNECTION_TIMEOUT ({connection_timeout}s) must be < DVR timeout ({dvr_read_timeout}s) "
f"so keepalives fire before DVR times out",
)
def test_keepalive_interval_is_short(self):
"""KEEPALIVE_INTERVAL must be short enough to send multiple keepalives in the gap."""
from apps.proxy.config import TSConfig as Config
interval = getattr(Config, "KEEPALIVE_INTERVAL", 0.5)
connection_timeout = getattr(Config, "CONNECTION_TIMEOUT", 10)
remaining_window = 15 - connection_timeout
self.assertGreater(
remaining_window / interval,
3,
f"KEEPALIVE_INTERVAL ({interval}s) is too long: only "
f"{remaining_window/interval:.1f} keepalives would fit in the "
f"{remaining_window}s window before DVR timeout",
)
@@ -0,0 +1,195 @@
"""
Unit tests for the keepalive duration cap in StreamGenerator._stream_data_generator.
Verifies that a client held in keepalive mode is disconnected after
MAX_KEEPALIVE_DURATION seconds, and that the timer resets when real data resumes.
"""
import time
from unittest.mock import MagicMock, patch, call
from django.test import TestCase
def _make_generator(consecutive_empty=10, local_index=10, buffer_index=10):
"""Minimal StreamGenerator stub for testing _stream_data_generator logic."""
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
gen = StreamGenerator.__new__(StreamGenerator)
gen.channel_id = "00000000-0000-0000-0000-000000000099"
gen.client_id = "test-client-duration"
gen.consecutive_empty = consecutive_empty
gen.empty_reads = 0
gen.local_index = local_index
gen.bytes_sent = 0
gen.chunks_sent = 0
gen.last_yield_time = time.time()
gen.stream_start_time = time.time()
gen.last_stats_time = time.time()
gen.last_stats_bytes = 0
gen.current_rate = 0.0
gen.last_ttl_refresh = time.time()
gen.ttl_refresh_interval = 3
gen.is_owner_worker = False
gen.stream_manager = None
gen._last_health_check_time = 0.0
gen._last_health_check_result = False
gen._health_check_interval = 2.0
gen.proxy_server = None
buffer = MagicMock()
buffer.index = buffer_index
buffer.get_optimized_client_data.return_value = ([], local_index)
buffer.find_oldest_available_chunk.return_value = None
gen.buffer = buffer
return gen
class KeepaliveDurationCapTests(TestCase):
"""MAX_KEEPALIVE_DURATION cap disconnects clients stuck in keepalive mode."""
def _run_generator_to_break(self, gen, max_iterations=20):
"""Drive _stream_data_generator until it breaks or hits iteration limit."""
iterations = 0
for _ in gen._stream_data_generator():
iterations += 1
if iterations >= max_iterations:
break
return iterations
def test_cap_fires_after_max_duration_exceeded(self):
"""Generator exits when keepalive has run longer than MAX_KEEPALIVE_DURATION."""
gen = _make_generator()
with patch.object(gen, '_check_resources', return_value=True), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent') as mock_gevent, \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
MockConfig.MAX_KEEPALIVE_DURATION = 30
# First call: keepalive_start_time not yet set (returns current)
# Second call: inside the cap check — simulate time elapsed > 30s
mock_time.time.side_effect = [
1000.0, # keepalive_start_time assignment
1031.0, # cap check: 31s elapsed > 30s limit
]
packets = list(gen._stream_data_generator())
# No packets should be yielded — cap fires before yield
self.assertEqual(len(packets), 0)
def test_cap_does_not_fire_before_max_duration(self):
"""Generator yields keepalive packets while within MAX_KEEPALIVE_DURATION."""
gen = _make_generator()
call_count = 0
def time_side_effect():
nonlocal call_count
call_count += 1
# keepalive_start_time set at t=1000; cap checks always see <30s elapsed
if call_count == 1:
return 1000.0 # keepalive_start_time
return 1010.0 # always 10s elapsed — under the 30s cap
with patch.object(gen, '_check_resources', side_effect=[True, True, False]), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
MockConfig.MAX_KEEPALIVE_DURATION = 30
mock_time.time.side_effect = time_side_effect
packets = list(gen._stream_data_generator())
# Two iterations with _check_resources=True should yield two keepalive packets
self.assertGreater(len(packets), 0)
def test_timer_resets_when_real_data_resumes(self):
"""keepalive_start_time is cleared to None when real chunks are received."""
gen = _make_generator()
chunk = b'\x47' * 188
real_chunks = ([chunk], gen.local_index + 1)
no_chunks = ([], gen.local_index)
# Sequence: no data (keepalive), then real data, then stop
gen.buffer.get_optimized_client_data.side_effect = [
no_chunks, # iteration 1: keepalive
real_chunks, # iteration 2: real data — should reset timer
no_chunks, # iteration 3: keepalive again — timer restarts fresh
]
captured_start_times = []
original_gen = gen
with patch.object(gen, '_check_resources', side_effect=[True, True, True, False]), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch.object(gen, '_process_chunks', return_value=iter([chunk])), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
MockConfig.MAX_KEEPALIVE_DURATION = 300
mock_time.time.return_value = 1000.0
list(gen._stream_data_generator())
# Test passes if no exception and generator completes normally —
# if the timer were NOT reset, the second keepalive block would
# carry over the old start time rather than starting fresh.
def test_cap_uses_config_value(self):
"""Cap threshold reads MAX_KEEPALIVE_DURATION from Config, not a hardcoded value."""
gen = _make_generator()
with patch.object(gen, '_check_resources', return_value=True), \
patch.object(gen, '_should_send_keepalive', return_value=True), \
patch.object(gen, '_is_ghost_client', return_value=False), \
patch.object(gen, '_is_timeout', return_value=False), \
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
MockPS.get_instance.return_value = None
MockConfig.KEEPALIVE_INTERVAL = 0
# Set a custom cap of 60s
MockConfig.MAX_KEEPALIVE_DURATION = 60
mock_time.time.side_effect = [
1000.0, # keepalive_start_time
1050.0, # cap check: 50s elapsed — under 60s, should NOT fire
1000.0, # last_yield_time update
1070.0, # cap check on next iteration: 70s elapsed — fires
]
packets = list(gen._stream_data_generator())
# First iteration: 50s < 60s cap — one keepalive yielded
# Second iteration: 70s > 60s cap — generator exits
self.assertEqual(len(packets), 1)
+169
View File
@@ -0,0 +1,169 @@
"""Tests for the _validate_url() helper in tasks.py.
Covers:
- Rejection of None, empty, and non-string inputs
- Non-HTTP URLs pass through without network requests
- HTTP(S) URLs validated via HEAD request (2xx/3xx pass, 4xx/5xx fail)
- Network errors (timeout, connection) treated as failures
- Per-worker result cache: hits, expiry, eviction
"""
from unittest.mock import patch, MagicMock
from django.test import TestCase
from apps.channels.tasks import _validate_url, _url_validation_cache, _URL_CACHE_TTL
class ValidateUrlInputTests(TestCase):
"""Input validation — no network requests should be made."""
def setUp(self):
_url_validation_cache.clear()
def test_none_returns_false(self):
self.assertFalse(_validate_url(None))
def test_empty_string_returns_false(self):
self.assertFalse(_validate_url(""))
def test_non_string_returns_false(self):
self.assertFalse(_validate_url(123))
self.assertFalse(_validate_url(["http://example.com"]))
@patch("apps.channels.tasks.requests.head")
def test_non_http_url_returns_true_without_request(self, mock_head):
"""file:// and other non-HTTP schemes skip validation."""
self.assertTrue(_validate_url("file:///local/path.jpg"))
self.assertTrue(_validate_url("/data/images/poster.jpg"))
mock_head.assert_not_called()
class ValidateUrlNetworkTests(TestCase):
"""HTTP(S) URL validation via HEAD request."""
def setUp(self):
_url_validation_cache.clear()
@patch("apps.channels.tasks.requests.head")
def test_200_returns_true(self, mock_head):
mock_head.return_value = MagicMock(status_code=200)
self.assertTrue(_validate_url("https://example.com/poster.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_302_redirect_returns_true(self, mock_head):
mock_head.return_value = MagicMock(status_code=302)
self.assertTrue(_validate_url("https://example.com/redirect"))
@patch("apps.channels.tasks.requests.head")
def test_404_returns_false(self, mock_head):
mock_head.return_value = MagicMock(status_code=404)
self.assertFalse(_validate_url("https://dead-cdn.com/missing.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_500_returns_false(self, mock_head):
mock_head.return_value = MagicMock(status_code=500)
self.assertFalse(_validate_url("https://broken.com/error"))
@patch("apps.channels.tasks.requests.head")
def test_timeout_returns_false(self, mock_head):
import requests
mock_head.side_effect = requests.Timeout("timed out")
self.assertFalse(_validate_url("https://slow-cdn.com/poster.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_connection_error_returns_false(self, mock_head):
import requests
mock_head.side_effect = requests.ConnectionError("refused")
self.assertFalse(_validate_url("https://unreachable.com/poster.jpg"))
@patch("apps.channels.tasks.requests.head")
def test_custom_timeout_passed_to_head(self, mock_head):
mock_head.return_value = MagicMock(status_code=200)
_validate_url("https://example.com/img.jpg", timeout=10)
mock_head.assert_called_once_with(
"https://example.com/img.jpg", timeout=10, allow_redirects=True
)
@patch("apps.channels.tasks.requests.get")
@patch("apps.channels.tasks.requests.head")
def test_405_falls_back_to_get(self, mock_head, mock_get):
"""When HEAD returns 405, fall back to a ranged GET request."""
mock_head.return_value = MagicMock(status_code=405)
mock_resp = MagicMock(status_code=200)
mock_get.return_value = mock_resp
self.assertTrue(_validate_url("https://no-head.com/poster.jpg"))
mock_get.assert_called_once()
mock_resp.close.assert_called_once()
@patch("apps.channels.tasks.requests.get")
@patch("apps.channels.tasks.requests.head")
def test_405_fallback_get_also_fails(self, mock_head, mock_get):
"""When HEAD returns 405 and GET also fails, return False."""
mock_head.return_value = MagicMock(status_code=405)
mock_get.return_value = MagicMock(status_code=403)
self.assertFalse(_validate_url("https://blocked.com/poster.jpg"))
class ValidateUrlCacheTests(TestCase):
"""Per-worker result caching."""
def setUp(self):
_url_validation_cache.clear()
@patch("apps.channels.tasks.requests.head")
def test_cache_hit_avoids_second_request(self, mock_head):
mock_head.return_value = MagicMock(status_code=200)
url = "https://cached.com/poster.jpg"
self.assertTrue(_validate_url(url))
self.assertTrue(_validate_url(url))
mock_head.assert_called_once()
@patch("apps.channels.tasks.requests.head")
def test_cache_hit_returns_false_for_failed_url(self, mock_head):
mock_head.return_value = MagicMock(status_code=404)
url = "https://dead.com/missing.jpg"
self.assertFalse(_validate_url(url))
self.assertFalse(_validate_url(url))
mock_head.assert_called_once()
@patch("apps.channels.tasks.time.monotonic")
@patch("apps.channels.tasks.requests.head")
def test_cache_expiry_triggers_new_request(self, mock_head, mock_time):
"""After TTL expires, a new HEAD request is made."""
mock_head.return_value = MagicMock(status_code=200)
url = "https://expiring.com/poster.jpg"
mock_time.return_value = 1000.0
self.assertTrue(_validate_url(url))
self.assertEqual(mock_head.call_count, 1)
# Within TTL — cache hit
mock_time.return_value = 1000.0 + _URL_CACHE_TTL - 1
self.assertTrue(_validate_url(url))
self.assertEqual(mock_head.call_count, 1)
# Past TTL — new request
mock_time.return_value = 1000.0 + _URL_CACHE_TTL + 1
self.assertTrue(_validate_url(url))
self.assertEqual(mock_head.call_count, 2)
@patch("apps.channels.tasks.time.monotonic")
@patch("apps.channels.tasks.requests.head")
def test_eviction_when_cache_exceeds_limit(self, mock_head, mock_time):
"""Expired entries are evicted when cache grows past 512 entries."""
mock_head.return_value = MagicMock(status_code=200)
# Fill cache with 513 entries at time 0
mock_time.return_value = 0.0
for i in range(513):
_url_validation_cache[f"https://fill-{i}.com/img.jpg"] = (True, 0.0)
# Advance past TTL and add one more — triggers eviction
mock_time.return_value = _URL_CACHE_TTL + 1
_validate_url("https://trigger-eviction.com/img.jpg")
# All 513 old entries expired and should be evicted
remaining = [k for k in _url_validation_cache if k.startswith("https://fill-")]
self.assertEqual(len(remaining), 0)
# The new entry should remain
self.assertIn("https://trigger-eviction.com/img.jpg", _url_validation_cache)
+12
View File
@@ -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'),
]
+25
View File
@@ -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'])
+41
View File
@@ -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')
View File
+17
View File
@@ -0,0 +1,17 @@
from django.urls import path
from rest_framework.routers import DefaultRouter
from .api_views import (
IntegrationViewSet,
EventSubscriptionViewSet,
DeliveryLogViewSet,
)
app_name = 'connect'
router = DefaultRouter()
router.register(r'integrations', IntegrationViewSet, basename='integration')
router.register(r'subscriptions', EventSubscriptionViewSet, basename='subscription')
router.register(r'logs', DeliveryLogViewSet, basename='delivery-log')
urlpatterns = []
urlpatterns += router.urls
+226
View File
@@ -0,0 +1,226 @@
from rest_framework import viewsets, status, serializers
from rest_framework.pagination import PageNumberPagination
from django_filters.rest_framework import DjangoFilterBackend
from rest_framework.response import Response
from rest_framework.decorators import action
from django.utils import timezone
from drf_spectacular.utils import extend_schema, inline_serializer
from .models import Integration, EventSubscription, DeliveryLog
from .serializers import (
IntegrationSerializer,
EventSubscriptionSerializer,
DeliveryLogSerializer,
)
from apps.accounts.permissions import (
Authenticated,
permission_classes_by_action,
IsAdmin,
)
from .handlers.webhook import WebhookHandler
from .handlers.script import ScriptHandler
class IntegrationViewSet(viewsets.ModelViewSet):
queryset = Integration.objects.all()
serializer_class = IntegrationSerializer
def get_permissions(self):
try:
perms = permission_classes_by_action[self.action]
except KeyError:
# Respect view/action-specific permission_classes if provided; fallback to Authenticated
perms = getattr(self, "permission_classes", [Authenticated])
return [perm() for perm in perms]
@action(detail=True, methods=["get"], url_path="subscriptions")
def list_subscriptions(self, request, pk=None):
qs = EventSubscription.objects.filter(integration_id=pk)
serializer = EventSubscriptionSerializer(qs, many=True)
return Response(serializer.data)
@extend_schema(
methods=["PUT"],
description=(
"Replace the integration's event subscriptions with the provided list. "
"Accepts a JSON array of subscription objects. "
"Existing subscriptions not in the list will be deleted. "
"The 'payload_template' field is only relevant for webhook integrations."
),
request=inline_serializer(
name="SetSubscriptionsRequest",
fields={
"event": serializers.CharField(help_text="Event name (e.g. 'channel_start')."),
"enabled": serializers.BooleanField(required=False, default=True),
"payload_template": serializers.CharField(required=False, allow_blank=True, allow_null=True, help_text="Custom payload template (webhook integrations only)."),
},
many=True,
),
responses={200: inline_serializer(
name="SetSubscriptionsResponse",
fields={
"event": serializers.CharField(),
"enabled": serializers.BooleanField(),
"payload_template": serializers.CharField(allow_null=True),
},
many=True,
)},
)
@action(detail=True, methods=["put"], url_path=r"subscriptions/set")
def set_subscriptions(self, request, pk=None):
"""
Replace the integration's subscriptions with the provided list.
Body format: [{"event": "channel_start", "enabled": true, "payload_template": "..."}, ...]
Any existing subscriptions not in the list will be deleted; missing ones will be created/updated.
"""
try:
integration = Integration.objects.get(pk=pk)
except Integration.DoesNotExist:
return Response(
{"detail": "Integration not found"}, status=status.HTTP_404_NOT_FOUND
)
data = request.data
if not isinstance(data, list):
return Response(
{"detail": "Expected a list of subscriptions"},
status=status.HTTP_400_BAD_REQUEST,
)
# Validate incoming items using serializer (without integration field)
# We'll attach the integration explicitly
valid_events = set(evt for evt, _ in EventSubscription.EVENT_CHOICES)
incoming = []
for item in data:
if not isinstance(item, dict):
return Response(
{"detail": "Each subscription must be an object"},
status=status.HTTP_400_BAD_REQUEST,
)
event = item.get("event")
if event not in valid_events:
return Response(
{"detail": f"Invalid event: {event}"},
status=status.HTTP_400_BAD_REQUEST,
)
# Only accept payload_template when the integration is a webhook
payload_template = item.get("payload_template") if integration.type == "webhook" else None
incoming.append(
{
"event": event,
"enabled": bool(item.get("enabled", True)),
"payload_template": payload_template,
}
)
incoming_events = {s["event"] for s in incoming}
# Delete subscriptions that are no longer present
EventSubscription.objects.filter(integration=integration).exclude(
event__in=incoming_events
).delete()
# Upsert incoming subscriptions
updated = []
for sub in incoming:
obj, _created = EventSubscription.objects.update_or_create(
integration=integration,
event=sub["event"],
defaults={
"enabled": sub["enabled"],
"payload_template": sub.get("payload_template"),
},
)
updated.append(obj)
serializer = EventSubscriptionSerializer(updated, many=True)
return Response(serializer.data, status=status.HTTP_200_OK)
@action(detail=True, methods=["post"], url_path="test", permission_classes=[IsAdmin])
def test(self, request, pk=None):
"""
Execute a saved integration (connect) with a dummy payload to verify configuration.
"""
try:
integration = Integration.objects.get(pk=pk)
except Integration.DoesNotExist:
return Response({"detail": "Integration not found"}, status=status.HTTP_404_NOT_FOUND)
# Build a dummy payload similar to system events
now = timezone.now().isoformat()
dummy_payload = {
"event": "test",
"timestamp": now,
"channel_name": "Test Channel",
"stream_name": "Test Stream",
"stream_url": "http://example.com/stream.m3u8",
"channel_url": "http://example.com/stream.m3u8",
"provider_name": "Test Provider",
"profile_used": "Default",
"test": True,
}
# Choose handler based on saved type
if integration.type == "webhook":
handler = WebhookHandler(integration, None, dummy_payload)
elif integration.type == "script":
handler = ScriptHandler(integration, None, dummy_payload)
else:
return Response(
{"success": False, "error": f"Unsupported integration type: {integration.type}"},
status=status.HTTP_400_BAD_REQUEST,
)
try:
result = handler.execute()
return Response(
{
"success": bool(result.get("success")),
"type": integration.type,
"request_payload": dummy_payload,
"result": result,
},
status=status.HTTP_200_OK,
)
except Exception as e:
return Response(
{
"success": False,
"type": integration.type,
"request_payload": dummy_payload,
"error": str(e),
},
status=status.HTTP_502_BAD_GATEWAY,
)
class EventSubscriptionViewSet(viewsets.ModelViewSet):
queryset = EventSubscription.objects.all()
serializer_class = EventSubscriptionSerializer
class DeliveryLogViewSet(viewsets.ReadOnlyModelViewSet):
queryset = DeliveryLog.objects.all().order_by("-created_at")
serializer_class = DeliveryLogSerializer
filter_backends = [DjangoFilterBackend]
# Support server-side pagination with page_size query param
class ConnectLogsPagination(PageNumberPagination):
page_size = 50
page_size_query_param = "page_size"
max_page_size = 250
pagination_class = ConnectLogsPagination
def get_queryset(self):
qs = super().get_queryset()
# Optional filters: integration id and type
integration_id = self.request.query_params.get("integration")
if integration_id:
qs = qs.filter(subscription__integration_id=integration_id)
integration_type = self.request.query_params.get("type")
if integration_type:
qs = qs.filter(subscription__integration__type=integration_type)
return qs

Some files were not shown because too many files have changed in this diff Show More