Proyecto LCX Dispatcharr multicuenta
Base Image Build / prepare (push) Has been cancelled
Build and Push Multi-Arch Docker Image / build-and-push (push) Has been cancelled
Frontend Tests / test (push) Has been cancelled
Base Image Build / docker (amd64, ubuntu-24.04) (push) Has been cancelled
Base Image Build / docker (arm64, ubuntu-24.04-arm) (push) Has been cancelled
Base Image Build / create-manifest (push) Has been cancelled
Base Image Build / prepare (push) Has been cancelled
Build and Push Multi-Arch Docker Image / build-and-push (push) Has been cancelled
Frontend Tests / test (push) Has been cancelled
Base Image Build / docker (amd64, ubuntu-24.04) (push) Has been cancelled
Base Image Build / docker (arm64, ubuntu-24.04-arm) (push) Has been cancelled
Base Image Build / create-manifest (push) Has been cancelled
This commit is contained in:
@@ -0,0 +1,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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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']
|
||||
@@ -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()),
|
||||
],
|
||||
),
|
||||
]
|
||||
+43
@@ -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),
|
||||
),
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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],
|
||||
}
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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())
|
||||
@@ -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
|
||||
]
|
||||
@@ -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'),
|
||||
]
|
||||
@@ -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"),
|
||||
]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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}")
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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}")
|
||||
@@ -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
@@ -0,0 +1,38 @@
|
||||
from django.contrib import admin
|
||||
from .models import Stream, Channel, ChannelGroup
|
||||
|
||||
@admin.register(Stream)
|
||||
class StreamAdmin(admin.ModelAdmin):
|
||||
list_display = (
|
||||
'id', # Primary Key
|
||||
'name',
|
||||
'channel_group',
|
||||
'url',
|
||||
'current_viewers',
|
||||
'updated_at',
|
||||
)
|
||||
|
||||
list_filter = ('channel_group',) # Filter by 'channel_group' (foreign key)
|
||||
|
||||
search_fields = ('id', 'name', 'url', 'channel_group__name') # Search by 'ChannelGroup' name
|
||||
|
||||
ordering = ('-updated_at',)
|
||||
|
||||
@admin.register(Channel)
|
||||
class ChannelAdmin(admin.ModelAdmin):
|
||||
list_display = (
|
||||
'id', # Primary Key
|
||||
'channel_number',
|
||||
'uuid',
|
||||
'name',
|
||||
'channel_group',
|
||||
'epg_data'
|
||||
)
|
||||
list_filter = ('channel_group',)
|
||||
search_fields = ('id', 'name', 'channel_group__name', 'epg_data') # Added 'id'
|
||||
ordering = ('channel_number',)
|
||||
|
||||
@admin.register(ChannelGroup)
|
||||
class ChannelGroupAdmin(admin.ModelAdmin):
|
||||
list_display = ('id', 'name') # Added 'id'
|
||||
search_fields = ('id', 'name') # Added 'id'
|
||||
@@ -0,0 +1,55 @@
|
||||
from django.urls import path, include
|
||||
from rest_framework.routers import DefaultRouter
|
||||
from .api_views import (
|
||||
StreamViewSet,
|
||||
ChannelViewSet,
|
||||
ChannelGroupViewSet,
|
||||
BulkDeleteStreamsAPIView,
|
||||
BulkDeleteChannelsAPIView,
|
||||
BulkDeleteLogosAPIView,
|
||||
CleanupUnusedLogosAPIView,
|
||||
LogoViewSet,
|
||||
ChannelProfileViewSet,
|
||||
UpdateChannelMembershipAPIView,
|
||||
BulkUpdateChannelMembershipAPIView,
|
||||
RecordingViewSet,
|
||||
RecurringRecordingRuleViewSet,
|
||||
GetChannelStreamsAPIView,
|
||||
SeriesRulesAPIView,
|
||||
DeleteSeriesRuleAPIView,
|
||||
EvaluateSeriesRulesAPIView,
|
||||
BulkRemoveSeriesRecordingsAPIView,
|
||||
BulkDeleteUpcomingRecordingsAPIView,
|
||||
ComskipConfigAPIView,
|
||||
)
|
||||
|
||||
app_name = 'channels' # for DRF routing
|
||||
|
||||
router = DefaultRouter()
|
||||
router.register(r'streams', StreamViewSet, basename='stream')
|
||||
router.register(r'groups', ChannelGroupViewSet, basename='channel-group')
|
||||
router.register(r'channels', ChannelViewSet, basename='channel')
|
||||
router.register(r'logos', LogoViewSet, basename='logo')
|
||||
router.register(r'profiles', ChannelProfileViewSet, basename='profile')
|
||||
router.register(r'recordings', RecordingViewSet, basename='recording')
|
||||
router.register(r'recurring-rules', RecurringRecordingRuleViewSet, basename='recurring-rule')
|
||||
|
||||
urlpatterns = [
|
||||
# Bulk delete is a single APIView, not a ViewSet
|
||||
path('streams/bulk-delete/', BulkDeleteStreamsAPIView.as_view(), name='bulk_delete_streams'),
|
||||
path('channels/bulk-delete/', BulkDeleteChannelsAPIView.as_view(), name='bulk_delete_channels'),
|
||||
path('logos/bulk-delete/', BulkDeleteLogosAPIView.as_view(), name='bulk_delete_logos'),
|
||||
path('logos/cleanup/', CleanupUnusedLogosAPIView.as_view(), name='cleanup_unused_logos'),
|
||||
path('channels/<int:channel_id>/streams/', GetChannelStreamsAPIView.as_view(), name='get_channel_streams'),
|
||||
path('profiles/<int:profile_id>/channels/<int:channel_id>/', UpdateChannelMembershipAPIView.as_view(), name='update_channel_membership'),
|
||||
path('profiles/<int:profile_id>/channels/bulk-update/', BulkUpdateChannelMembershipAPIView.as_view(), name='bulk_update_channel_membership'),
|
||||
# DVR series rules (order matters: specific routes before catch-all slug)
|
||||
path('series-rules/', SeriesRulesAPIView.as_view(), name='series_rules'),
|
||||
path('series-rules/evaluate/', EvaluateSeriesRulesAPIView.as_view(), name='evaluate_series_rules'),
|
||||
path('series-rules/bulk-remove/', BulkRemoveSeriesRecordingsAPIView.as_view(), name='bulk_remove_series_recordings'),
|
||||
path('series-rules/<path:tvg_id>/', DeleteSeriesRuleAPIView.as_view(), name='delete_series_rule'),
|
||||
path('recordings/bulk-delete-upcoming/', BulkDeleteUpcomingRecordingsAPIView.as_view(), name='bulk_delete_upcoming_recordings'),
|
||||
path('dvr/comskip-config/', ComskipConfigAPIView.as_view(), name='comskip_config'),
|
||||
]
|
||||
|
||||
urlpatterns += router.urls
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,11 @@
|
||||
from django.apps import AppConfig
|
||||
|
||||
class ChannelsConfig(AppConfig):
|
||||
default_auto_field = 'django.db.models.BigAutoField'
|
||||
name = 'apps.channels'
|
||||
verbose_name = "Channel & Stream Management"
|
||||
label = 'dispatcharr_channels'
|
||||
|
||||
def ready(self):
|
||||
# Import signals so they get registered.
|
||||
import apps.channels.signals
|
||||
@@ -0,0 +1,53 @@
|
||||
from django import forms
|
||||
from .models import Stream, Channel, ChannelGroup
|
||||
|
||||
#
|
||||
# ChannelGroup Form
|
||||
#
|
||||
class ChannelGroupForm(forms.ModelForm):
|
||||
class Meta:
|
||||
model = ChannelGroup
|
||||
fields = ['name']
|
||||
|
||||
|
||||
#
|
||||
# Channel Form
|
||||
#
|
||||
class ChannelForm(forms.ModelForm):
|
||||
# Explicitly define channel_number as FloatField to ensure decimal values work
|
||||
channel_number = forms.FloatField(
|
||||
required=False,
|
||||
widget=forms.NumberInput(attrs={'step': '0.1'}), # Allow decimal steps
|
||||
help_text="Channel number can include decimals (e.g., 1.1, 2.5)"
|
||||
)
|
||||
|
||||
channel_group = forms.ModelChoiceField(
|
||||
queryset=ChannelGroup.objects.all(),
|
||||
required=False,
|
||||
label="Channel Group",
|
||||
empty_label="--- No group ---"
|
||||
)
|
||||
|
||||
class Meta:
|
||||
model = Channel
|
||||
fields = [
|
||||
'channel_number',
|
||||
'name',
|
||||
'channel_group',
|
||||
]
|
||||
|
||||
|
||||
#
|
||||
# Example: Stream Form (optional if you want a ModelForm for Streams)
|
||||
#
|
||||
class StreamForm(forms.ModelForm):
|
||||
class Meta:
|
||||
model = Stream
|
||||
fields = [
|
||||
'name',
|
||||
'url',
|
||||
'logo_url',
|
||||
'epg_data',
|
||||
'local_file',
|
||||
'channel_group',
|
||||
]
|
||||
@@ -0,0 +1,77 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-05 22:07
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
initial = True
|
||||
|
||||
dependencies = [
|
||||
('core', '0001_initial'),
|
||||
('m3u', '0001_initial'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='ChannelGroup',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(max_length=100, unique=True)),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='Channel',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('channel_number', models.IntegerField()),
|
||||
('channel_name', models.CharField(max_length=255)),
|
||||
('logo_url', models.URLField(blank=True, max_length=2000, null=True)),
|
||||
('logo_file', models.ImageField(blank=True, null=True, upload_to='logos/')),
|
||||
('tvg_id', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('tvg_name', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('stream_profile', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='core.streamprofile')),
|
||||
('channel_group', models.ForeignKey(blank=True, help_text='Channel group this channel belongs to.', null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='dispatcharr_channels.channelgroup')),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='Stream',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(default='Default Stream', max_length=255)),
|
||||
('url', models.URLField()),
|
||||
('custom_url', models.URLField(blank=True, max_length=2000, null=True)),
|
||||
('logo_url', models.URLField(blank=True, max_length=2000, null=True)),
|
||||
('tvg_id', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('local_file', models.FileField(blank=True, null=True, upload_to='uploads/')),
|
||||
('current_viewers', models.PositiveIntegerField(default=0)),
|
||||
('updated_at', models.DateTimeField(auto_now=True)),
|
||||
('group_name', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('m3u_account', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.CASCADE, related_name='streams', to='m3u.m3uaccount')),
|
||||
('stream_profile', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='streams', to='core.streamprofile')),
|
||||
],
|
||||
options={
|
||||
'verbose_name': 'Stream',
|
||||
'verbose_name_plural': 'Streams',
|
||||
'ordering': ['-updated_at'],
|
||||
},
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='ChannelStream',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('order', models.PositiveIntegerField(default=0)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.channel')),
|
||||
('stream', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.stream')),
|
||||
],
|
||||
options={
|
||||
'ordering': ['order'],
|
||||
},
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='streams',
|
||||
field=models.ManyToManyField(blank=True, related_name='channels', through='dispatcharr_channels.ChannelStream', to='dispatcharr_channels.stream'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,27 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-16 12:21
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0001_initial'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RenameField(
|
||||
model_name='channel',
|
||||
old_name='channel_name',
|
||||
new_name='name',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='stream',
|
||||
name='url',
|
||||
),
|
||||
migrations.RenameField(
|
||||
model_name='stream',
|
||||
old_name='custom_url',
|
||||
new_name='url',
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,38 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-16 13:25
|
||||
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
# Generated by Django 5.1.6 on 2025-03-16 13:25
|
||||
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
def generate_uuids(apps, schema_editor):
|
||||
Channel = apps.get_model('dispatcharr_channels', 'Channel')
|
||||
for channel in Channel.objects.all():
|
||||
if not channel.uuid:
|
||||
channel.uuid = uuid.uuid4()
|
||||
channel.save()
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0002_rename_channel_name_channel_name_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='uuid',
|
||||
field=models.UUIDField(default=uuid.uuid4, editable=False),
|
||||
),
|
||||
migrations.RunPython(generate_uuids),
|
||||
migrations.AlterField(
|
||||
model_name='channel',
|
||||
name='uuid',
|
||||
field=models.UUIDField(default=uuid.uuid4, editable=False, unique=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-17 21:16
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0003_channel_uuid'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='is_custom',
|
||||
field=models.BooleanField(default=False, help_text='Whether this is a user-created stream or from an M3U account'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,44 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-19 16:33
|
||||
|
||||
import django.db.models.deletion
|
||||
import django.utils.timezone
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0004_stream_is_custom'),
|
||||
('m3u', '0003_create_custom_account'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='channel_group',
|
||||
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='streams', to='dispatcharr_channels.channelgroup'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='last_seen',
|
||||
field=models.DateTimeField(db_index=True, default=django.utils.timezone.now),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='channel',
|
||||
name='uuid',
|
||||
field=models.UUIDField(db_index=True, default=uuid.uuid4, editable=False, unique=True),
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='ChannelGroupM3UAccount',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('enabled', models.BooleanField(default=True)),
|
||||
('channel_group', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='m3u_account', to='dispatcharr_channels.channelgroup')),
|
||||
('m3u_account', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='channel_group', to='m3u.m3uaccount')),
|
||||
],
|
||||
options={
|
||||
'unique_together': {('channel_group', 'm3u_account')},
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,51 @@
|
||||
# In your app's migrations folder, create a new migration file
|
||||
# e.g., migrations/000X_migrate_channel_group_to_foreign_key.py
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
def migrate_channel_group(apps, schema_editor):
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
ChannelGroup = apps.get_model('dispatcharr_channels', 'ChannelGroup')
|
||||
ChannelGroupM3UAccount = apps.get_model('dispatcharr_channels', 'ChannelGroup')
|
||||
M3UAccount = apps.get_model('m3u', 'M3UAccount')
|
||||
|
||||
streams_to_update = []
|
||||
for stream in Stream.objects.all():
|
||||
# If the stream has a 'channel_group' string, try to find or create the ChannelGroup
|
||||
if stream.group_name: # group_name holds the channel group string
|
||||
channel_group_name = stream.group_name.strip()
|
||||
|
||||
# Try to find the ChannelGroup by name
|
||||
channel_group, created = ChannelGroup.objects.get_or_create(name=channel_group_name)
|
||||
|
||||
# Set the foreign key to the found or newly created ChannelGroup
|
||||
stream.channel_group = channel_group
|
||||
|
||||
streams_to_update.append(stream)
|
||||
|
||||
# If the stream has an M3U account, ensure the M3U account is linked
|
||||
if stream.m3u_account:
|
||||
ChannelGroupM3UAccount.objects.get_or_create(
|
||||
channel_group=channel_group,
|
||||
m3u_account=stream.m3u_account,
|
||||
enabled=True # Or set it to whatever the default logic is
|
||||
)
|
||||
|
||||
Stream.objects.bulk_update(streams_to_update, ['channel_group'])
|
||||
|
||||
def reverse_migration(apps, schema_editor):
|
||||
# This reverse migration would undo the changes, setting `channel_group` to `None` and clearing any relationships.
|
||||
Stream = apps.get_model('yourapp', 'Stream')
|
||||
for stream in Stream.objects.all():
|
||||
stream.channel_group = None
|
||||
stream.save()
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0005_stream_channel_group_stream_last_seen_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(migrate_channel_group, reverse_code=reverse_migration),
|
||||
]
|
||||
@@ -0,0 +1,17 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-19 16:43
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0006_migrate_stream_groups'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='stream',
|
||||
name='group_name',
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-19 18:21
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0007_remove_stream_group_name'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_hash',
|
||||
field=models.CharField(db_index=True, help_text='Unique hash for this stream from the M3U account', max_length=255, null=True, unique=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='logo_url',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,24 @@
|
||||
# Generated by Django 5.1.6 on 2025-03-26 12:59
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0008_stream_stream_hash'),
|
||||
('epg', '0004_epgdata_epg_source_alter_epgdata_tvg_id'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='channel',
|
||||
name='tvg_name',
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='epg_data',
|
||||
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='epg.epgdata'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-01 17:36
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0009_remove_channel_tvg_name_channel_epg_data'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='custom_properties',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,35 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-01 22:14
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0010_stream_custom_properties'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='Logo',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(max_length=255)),
|
||||
('url', models.URLField(unique=True)),
|
||||
],
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='channel',
|
||||
name='logo_file',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='channel',
|
||||
name='logo_url',
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='logo',
|
||||
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='channels', to='dispatcharr_channels.logo'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,33 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-02 23:27
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0011_logo_remove_channel_logo_file_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='ChannelProfile',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.CharField(max_length=100, unique=True)),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='ChannelProfileMembership',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('enabled', models.BooleanField(default=True)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.channel')),
|
||||
('channel_profile', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='dispatcharr_channels.channelprofile')),
|
||||
],
|
||||
options={
|
||||
'unique_together': {('channel_profile', 'channel')},
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-04 15:04
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0012_channelprofile_channelprofilemembership'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='logo',
|
||||
name='url',
|
||||
field=models.TextField(unique=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,24 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-05 22:25
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0013_alter_logo_url'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='Recording',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('start_time', models.DateTimeField()),
|
||||
('end_time', models.DateTimeField()),
|
||||
('task_id', models.CharField(blank=True, max_length=255, null=True)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='recordings', to='dispatcharr_channels.channel')),
|
||||
],
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-07 16:47
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0014_recording'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='recording',
|
||||
name='custom_properties',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,38 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-18 16:21
|
||||
|
||||
from django.db import migrations, models
|
||||
from django.db.models import Count
|
||||
|
||||
def remove_duplicate_channel_streams(apps, schema_editor):
|
||||
ChannelStream = apps.get_model('dispatcharr_channels', 'ChannelStream')
|
||||
# Find duplicates by (channel, stream)
|
||||
duplicates = (
|
||||
ChannelStream.objects
|
||||
.values('channel', 'stream')
|
||||
.annotate(count=Count('id'))
|
||||
.filter(count__gt=1)
|
||||
)
|
||||
|
||||
for dupe in duplicates:
|
||||
# Get all duplicates for this pair
|
||||
dups = ChannelStream.objects.filter(
|
||||
channel=dupe['channel'],
|
||||
stream=dupe['stream']
|
||||
).order_by('id')
|
||||
|
||||
# Keep the first one, delete the rest
|
||||
dups.exclude(id=dups.first().id).delete()
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0015_recording_custom_properties'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(remove_duplicate_channel_streams),
|
||||
migrations.AddConstraint(
|
||||
model_name='channelstream',
|
||||
constraint=models.UniqueConstraint(fields=('channel', 'stream'), name='unique_channel_stream'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-21 20:47
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0016_channelstream_unique_channel_stream'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channelgroup',
|
||||
name='name',
|
||||
field=models.TextField(db_index=True, unique=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-04-27 14:12
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0017_alter_channelgroup_name'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='custom_properties',
|
||||
field=models.TextField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-05-04 00:02
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0018_channelgroupm3uaccount_custom_properties_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='tvc_guide_stationid',
|
||||
field=models.CharField(blank=True, max_length=255, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-05-15 19:37
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0019_channel_tvc_guide_stationid'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channel',
|
||||
name='channel_number',
|
||||
field=models.FloatField(db_index=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.1.6 on 2025-05-18 14:31
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0020_alter_channel_channel_number'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='user_level',
|
||||
field=models.IntegerField(default=0),
|
||||
),
|
||||
]
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
# Generated by Django 5.1.6 on 2025-07-13 23:08
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0021_channel_user_level'),
|
||||
('m3u', '0012_alter_m3uaccount_refresh_interval'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='auto_created',
|
||||
field=models.BooleanField(default=False, help_text='Whether this channel was automatically created via M3U auto channel sync'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='auto_created_by',
|
||||
field=models.ForeignKey(blank=True, help_text='The M3U account that auto-created this channel', null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='auto_created_channels', to='m3u.m3uaccount'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='auto_channel_sync',
|
||||
field=models.BooleanField(default=False, help_text='Automatically create/delete channels to match streams in this group'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='auto_sync_channel_start',
|
||||
field=models.FloatField(blank=True, help_text='Starting channel number for auto-created channels in this group', null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.1.6 on 2025-07-29 02:39
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0022_channel_auto_created_channel_auto_created_by_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_stats',
|
||||
field=models.JSONField(blank=True, help_text='JSON object containing stream statistics like video codec, resolution, etc.', null=True),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_stats_updated_at',
|
||||
field=models.DateTimeField(blank=True, db_index=True, help_text='When stream statistics were last updated', null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,19 @@
|
||||
# Generated by Django 5.2.4 on 2025-08-22 20:14
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0023_stream_stream_stats_stream_stream_stats_updated_at'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='channel_group',
|
||||
field=models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='m3u_accounts', to='dispatcharr_channels.channelgroup'),
|
||||
),
|
||||
]
|
||||
+28
@@ -0,0 +1,28 @@
|
||||
# Generated by Django 5.2.4 on 2025-09-02 14:30
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0024_alter_channelgroupm3uaccount_channel_group'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='custom_properties',
|
||||
field=models.JSONField(blank=True, default=dict, null=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='recording',
|
||||
name='custom_properties',
|
||||
field=models.JSONField(blank=True, default=dict, null=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='custom_properties',
|
||||
field=models.JSONField(blank=True, default=dict, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# Generated by Django 5.0.14 on 2025-09-18 14:56
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0025_alter_channelgroupm3uaccount_custom_properties_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='RecurringRecordingRule',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('days_of_week', models.JSONField(default=list)),
|
||||
('start_time', models.TimeField()),
|
||||
('end_time', models.TimeField()),
|
||||
('enabled', models.BooleanField(default=True)),
|
||||
('name', models.CharField(blank=True, max_length=255)),
|
||||
('created_at', models.DateTimeField(auto_now_add=True)),
|
||||
('updated_at', models.DateTimeField(auto_now=True)),
|
||||
('channel', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='recurring_rules', to='dispatcharr_channels.channel')),
|
||||
],
|
||||
options={
|
||||
'ordering': ['channel', 'start_time'],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.2.4 on 2025-10-05 20:50
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0026_recurringrecordingrule'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='recurringrecordingrule',
|
||||
name='end_date',
|
||||
field=models.DateField(blank=True, null=True),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='recurringrecordingrule',
|
||||
name='start_date',
|
||||
field=models.DateField(blank=True, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,25 @@
|
||||
# Generated by Django 5.2.4 on 2025-10-06 22:55
|
||||
|
||||
import django.utils.timezone
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0027_recurringrecordingrule_end_date_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='created_at',
|
||||
field=models.DateTimeField(auto_now_add=True, default=django.utils.timezone.now, help_text='Timestamp when this channel was created'),
|
||||
preserve_default=False,
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='updated_at',
|
||||
field=models.DateTimeField(auto_now=True, help_text='Timestamp when this channel was last updated'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,54 @@
|
||||
# Generated migration to backfill stream_hash for existing custom streams
|
||||
|
||||
from django.db import migrations
|
||||
import hashlib
|
||||
|
||||
|
||||
def backfill_custom_stream_hashes(apps, schema_editor):
|
||||
"""
|
||||
Generate stream_hash for all custom streams that don't have one.
|
||||
Uses stream ID to create a stable hash that won't change when name/url is edited.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
|
||||
custom_streams_without_hash = Stream.objects.filter(
|
||||
is_custom=True,
|
||||
stream_hash__isnull=True
|
||||
)
|
||||
|
||||
updated_count = 0
|
||||
for stream in custom_streams_without_hash:
|
||||
# Generate a stable hash using the stream's ID
|
||||
# This ensures the hash never changes even if name/url is edited
|
||||
unique_string = f"custom_stream_{stream.id}"
|
||||
stream.stream_hash = hashlib.sha256(unique_string.encode()).hexdigest()
|
||||
stream.save(update_fields=['stream_hash'])
|
||||
updated_count += 1
|
||||
|
||||
if updated_count > 0:
|
||||
print(f"Backfilled stream_hash for {updated_count} custom streams")
|
||||
else:
|
||||
print("No custom streams needed stream_hash backfill")
|
||||
|
||||
|
||||
def reverse_backfill(apps, schema_editor):
|
||||
"""
|
||||
Reverse migration - clear stream_hash for custom streams.
|
||||
Note: This will break preview functionality for custom streams.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
|
||||
custom_streams = Stream.objects.filter(is_custom=True)
|
||||
count = custom_streams.update(stream_hash=None)
|
||||
print(f"Cleared stream_hash for {count} custom streams")
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0028_channel_created_at_channel_updated_at'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(backfill_custom_stream_hashes, reverse_backfill),
|
||||
]
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.2.4 on 2025-10-28 20:00
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0029_backfill_custom_stream_hashes'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='url',
|
||||
field=models.URLField(blank=True, max_length=4096, null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,29 @@
|
||||
# Generated by Django 5.2.9 on 2026-01-09 18:19
|
||||
|
||||
import django.utils.timezone
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0030_alter_stream_url'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='is_stale',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this group relationship is stale (not seen in recent refresh, pending deletion)'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='channelgroupm3uaccount',
|
||||
name='last_seen',
|
||||
field=models.DateTimeField(db_index=True, default=django.utils.timezone.now, help_text='Last time this group was seen in the M3U source during a refresh'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='is_stale',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this stream is stale (not seen in recent refresh, pending deletion)'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,23 @@
|
||||
# Generated by Django 5.2.9 on 2026-01-17 16:56
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0031_channelgroupm3uaccount_is_stale_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='channel',
|
||||
name='is_adult',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this channel contains adult content'),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='is_adult',
|
||||
field=models.BooleanField(db_index=True, default=False, help_text='Whether this stream contains adult content'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,205 @@
|
||||
# Generated by Django - Add stream_id and channel_number fields with data migration
|
||||
|
||||
from django.db import migrations, models
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def populate_fields_and_rehash(apps, schema_editor):
|
||||
"""
|
||||
Populate stream_id and stream_chno from custom_properties for XC account streams,
|
||||
populate stream_chno from tvg-chno for standard M3U accounts,
|
||||
then rehash XC streams using stable hash keys.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
M3UAccount = apps.get_model('m3u', 'M3UAccount')
|
||||
CoreSettings = apps.get_model('core', 'CoreSettings')
|
||||
|
||||
# Get hash keys from settings
|
||||
try:
|
||||
stream_settings = CoreSettings.objects.get(key='stream_settings')
|
||||
hash_key_str = stream_settings.value.get('m3u_hash_key', '') if stream_settings.value else ''
|
||||
keys = [k.strip() for k in hash_key_str.split(',') if k.strip()] if hash_key_str else []
|
||||
except CoreSettings.DoesNotExist:
|
||||
keys = []
|
||||
|
||||
logger.info(f"Using hash keys: {keys}")
|
||||
|
||||
# Get XC account IDs
|
||||
xc_account_ids = set(
|
||||
M3UAccount.objects.filter(account_type='XC').values_list('id', flat=True)
|
||||
)
|
||||
|
||||
logger.info(f"Found {len(xc_account_ids)} XC accounts")
|
||||
|
||||
# Track hash collisions for XC streams
|
||||
hash_map = {} # new_hash -> stream_id
|
||||
duplicates_to_delete = []
|
||||
|
||||
# Process all streams in batches
|
||||
batch_size = 1000
|
||||
processed = 0
|
||||
updated = 0
|
||||
|
||||
total_count = Stream.objects.count()
|
||||
logger.info(f"Processing {total_count} total streams")
|
||||
|
||||
streams_to_update = []
|
||||
|
||||
for stream in Stream.objects.select_related('channel_group', 'm3u_account').iterator(chunk_size=batch_size):
|
||||
processed += 1
|
||||
needs_update = False
|
||||
|
||||
custom_props = stream.custom_properties or {}
|
||||
is_xc = stream.m3u_account_id in xc_account_ids if stream.m3u_account_id else False
|
||||
|
||||
# Extract stream_id (XC accounts only)
|
||||
if is_xc and isinstance(custom_props, dict):
|
||||
provider_stream_id = custom_props.get('stream_id')
|
||||
if provider_stream_id:
|
||||
try:
|
||||
stream.stream_id = int(provider_stream_id)
|
||||
needs_update = True
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# Extract stream_chno
|
||||
channel_num = None
|
||||
if isinstance(custom_props, dict):
|
||||
if is_xc:
|
||||
# XC accounts use 'num'
|
||||
channel_num = custom_props.get('num')
|
||||
else:
|
||||
# Standard M3U accounts use 'tvg-chno' or 'channel-number' (case insensitive check)
|
||||
for key in ['tvg-chno', 'TVG-CHNO', 'tvg-Chno', 'Tvg-Chno', 'channel-number', 'Channel-Number', 'CHANNEL-NUMBER']:
|
||||
if key in custom_props:
|
||||
channel_num = custom_props.get(key)
|
||||
break
|
||||
|
||||
if channel_num is not None:
|
||||
try:
|
||||
stream.stream_chno = float(channel_num)
|
||||
needs_update = True
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# Rehash XC streams only when 'url' is in hash keys (otherwise hash wouldn't change)
|
||||
if is_xc and stream.stream_id and keys and 'url' in keys:
|
||||
# For XC accounts, use stream_id instead of url when 'url' is in the hash keys
|
||||
# This ensures credential/URL changes don't break stream identity
|
||||
effective_url = stream.stream_id
|
||||
|
||||
# Get group name
|
||||
group_name = stream.channel_group.name if stream.channel_group else None
|
||||
|
||||
# Build hash parts
|
||||
stream_parts = {
|
||||
"name": stream.name,
|
||||
"url": effective_url,
|
||||
"tvg_id": stream.tvg_id,
|
||||
"m3u_id": stream.m3u_account_id,
|
||||
"group": group_name
|
||||
}
|
||||
hash_parts = {key: stream_parts[key] for key in keys if key in stream_parts}
|
||||
|
||||
# When using stream_id instead of URL, we MUST include m3u_id to prevent
|
||||
# collisions across different XC accounts (stream_id is only unique per account)
|
||||
if 'm3u_id' not in hash_parts:
|
||||
hash_parts['m3u_id'] = stream.m3u_account_id
|
||||
|
||||
# Generate hash
|
||||
serialized_obj = json.dumps(hash_parts, sort_keys=True)
|
||||
new_hash = hashlib.sha256(serialized_obj.encode()).hexdigest()
|
||||
|
||||
# Check for collisions
|
||||
if new_hash in hash_map:
|
||||
# Duplicate - mark for deletion (keep the first one)
|
||||
duplicates_to_delete.append(stream.id)
|
||||
continue
|
||||
|
||||
hash_map[new_hash] = stream.id
|
||||
stream.stream_hash = new_hash
|
||||
needs_update = True
|
||||
|
||||
if needs_update:
|
||||
streams_to_update.append(stream)
|
||||
updated += 1
|
||||
|
||||
# Bulk update in batches
|
||||
if len(streams_to_update) >= batch_size:
|
||||
Stream.objects.bulk_update(
|
||||
streams_to_update,
|
||||
['stream_id', 'stream_chno', 'stream_hash'],
|
||||
batch_size=500
|
||||
)
|
||||
logger.info(f"Updated batch: {processed}/{total_count} streams processed")
|
||||
streams_to_update = []
|
||||
|
||||
# Final batch
|
||||
if streams_to_update:
|
||||
Stream.objects.bulk_update(
|
||||
streams_to_update,
|
||||
['stream_id', 'stream_chno', 'stream_hash'],
|
||||
batch_size=500
|
||||
)
|
||||
|
||||
# Delete duplicates if any
|
||||
if duplicates_to_delete:
|
||||
logger.warning(f"Deleting {len(duplicates_to_delete)} duplicate streams due to hash collisions")
|
||||
Stream.objects.filter(id__in=duplicates_to_delete).delete()
|
||||
|
||||
logger.info(f"Migration complete: {updated} streams updated, {len(duplicates_to_delete)} duplicates removed")
|
||||
|
||||
|
||||
def reverse_migration(apps, schema_editor):
|
||||
"""
|
||||
Reverse migration - clear fields but don't attempt to reverse hash changes.
|
||||
"""
|
||||
Stream = apps.get_model('dispatcharr_channels', 'Stream')
|
||||
Stream.objects.all().update(stream_id=None, stream_chno=None)
|
||||
logger.info("Cleared stream_id and stream_chno fields. Note: stream hashes were not reverted.")
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0032_channel_is_adult_stream_is_adult'),
|
||||
('m3u', '0018_add_profile_custom_properties'),
|
||||
('core', '0020_change_coresettings_value_to_jsonfield'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
# Schema changes - add fields WITHOUT indexes first
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_id',
|
||||
field=models.IntegerField(
|
||||
blank=True,
|
||||
help_text='Provider stream ID (e.g., XC stream_id) for stable identity across credential changes',
|
||||
null=True,
|
||||
),
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='stream',
|
||||
name='stream_chno',
|
||||
field=models.FloatField(
|
||||
blank=True,
|
||||
help_text='Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1',
|
||||
null=True,
|
||||
),
|
||||
),
|
||||
# Data migration (may delete duplicates, which would conflict with pending index creation)
|
||||
migrations.RunPython(populate_fields_and_rehash, reverse_migration),
|
||||
# Add indexes AFTER data migration completes
|
||||
migrations.AddIndex(
|
||||
model_name='stream',
|
||||
index=models.Index(fields=['stream_id'], name='dispatcharr_stream_id_idx'),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='stream',
|
||||
index=models.Index(fields=['stream_chno'], name='dispatcharr_stream_chno_idx'),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# Generated by Django 5.2.9 on 2026-02-01 03:21
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('dispatcharr_channels', '0033_stream_id_stream_chno'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveIndex(
|
||||
model_name='stream',
|
||||
name='dispatcharr_stream_id_idx',
|
||||
),
|
||||
migrations.RemoveIndex(
|
||||
model_name='stream',
|
||||
name='dispatcharr_stream_chno_idx',
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='stream_chno',
|
||||
field=models.FloatField(blank=True, db_index=True, help_text='Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1', null=True),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='stream',
|
||||
name='stream_id',
|
||||
field=models.IntegerField(blank=True, db_index=True, help_text='Provider stream ID (e.g., XC stream_id) for stable identity across credential changes', null=True),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,988 @@
|
||||
from django.db import models
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.conf import settings
|
||||
from core.models import StreamProfile, CoreSettings
|
||||
from core.utils import RedisClient
|
||||
from apps.proxy.ts_proxy.redis_keys import RedisKeys
|
||||
from apps.proxy.ts_proxy.constants import ChannelMetadataField
|
||||
import logging
|
||||
import uuid
|
||||
from django.utils import timezone
|
||||
import hashlib
|
||||
import json
|
||||
from apps.epg.models import EPGData
|
||||
from apps.accounts.models import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# If you have an M3UAccount model in apps.m3u, you can still import it:
|
||||
from apps.m3u.models import M3UAccount
|
||||
|
||||
|
||||
# Add fallback functions if Redis isn't available
|
||||
def get_total_viewers(channel_id):
|
||||
"""Get viewer count from Redis or return 0 if Redis isn't available"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
try:
|
||||
return int(redis_client.get(f"channel:{channel_id}:viewers") or 0)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
class ChannelGroup(models.Model):
|
||||
name = models.TextField(unique=True, db_index=True)
|
||||
|
||||
def related_channels(self):
|
||||
# local import if needed to avoid cyc. Usually fine in a single file though
|
||||
return Channel.objects.filter(channel_group=self)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def bulk_create_and_fetch(cls, objects):
|
||||
# Perform the bulk create operation
|
||||
cls.objects.bulk_create(objects)
|
||||
|
||||
# Use a unique field to fetch the created objects (assuming 'name' is unique)
|
||||
created_objects = cls.objects.filter(name__in=[obj.name for obj in objects])
|
||||
|
||||
return created_objects
|
||||
|
||||
|
||||
class Stream(models.Model):
|
||||
"""
|
||||
Represents a single stream (e.g. from an M3U source or custom URL).
|
||||
"""
|
||||
|
||||
name = models.CharField(max_length=255, default="Default Stream")
|
||||
url = models.URLField(max_length=4096, blank=True, null=True)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount,
|
||||
on_delete=models.CASCADE,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
logo_url = models.TextField(blank=True, null=True)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
local_file = models.FileField(upload_to="uploads/", blank=True, null=True)
|
||||
current_viewers = models.PositiveIntegerField(default=0)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
null=True,
|
||||
blank=True,
|
||||
on_delete=models.SET_NULL,
|
||||
related_name="streams",
|
||||
)
|
||||
is_custom = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this is a user-created stream or from an M3U account",
|
||||
)
|
||||
stream_hash = models.CharField(
|
||||
max_length=255,
|
||||
null=True,
|
||||
unique=True,
|
||||
help_text="Unique hash for this stream from the M3U account",
|
||||
db_index=True,
|
||||
)
|
||||
last_seen = models.DateTimeField(db_index=True, default=timezone.now)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream is stale (not seen in recent refresh, pending deletion)"
|
||||
)
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream contains adult content"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
stream_id = models.IntegerField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider stream ID (e.g., XC stream_id) for stable identity across credential changes"
|
||||
)
|
||||
stream_chno = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1"
|
||||
)
|
||||
|
||||
# Stream statistics fields
|
||||
stream_stats = models.JSONField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="JSON object containing stream statistics like video codec, resolution, etc."
|
||||
)
|
||||
stream_stats_updated_at = models.DateTimeField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="When stream statistics were last updated",
|
||||
db_index=True
|
||||
)
|
||||
|
||||
class Meta:
|
||||
# If you use m3u_account, you might do unique_together = ('name','url','m3u_account')
|
||||
verbose_name = "Stream"
|
||||
verbose_name_plural = "Streams"
|
||||
ordering = ["-updated_at"]
|
||||
|
||||
def __str__(self):
|
||||
return self.name or self.url or f"Stream ID {self.id}"
|
||||
|
||||
@classmethod
|
||||
def generate_hash_key(cls, name, url, tvg_id, keys=None, m3u_id=None, group=None,
|
||||
account_type=None, stream_id=None):
|
||||
if keys is None:
|
||||
keys = CoreSettings.get_m3u_hash_key().split(",")
|
||||
|
||||
# For XC accounts, use stream_id instead of url when 'url' is in the hash keys
|
||||
# This ensures credential/URL changes don't break stream identity
|
||||
effective_url = url
|
||||
use_stream_id = account_type == 'XC' and stream_id and 'url' in keys
|
||||
if use_stream_id:
|
||||
effective_url = stream_id
|
||||
|
||||
stream_parts = {"name": name, "url": effective_url, "tvg_id": tvg_id, "m3u_id": m3u_id, "group": group}
|
||||
|
||||
hash_parts = {key: stream_parts[key] for key in keys if key in stream_parts}
|
||||
|
||||
# When using stream_id instead of URL, we MUST include m3u_id to prevent
|
||||
# collisions across different XC accounts (stream_id is only unique per account)
|
||||
if use_stream_id and 'm3u_id' not in hash_parts:
|
||||
hash_parts['m3u_id'] = m3u_id
|
||||
|
||||
# Serialize and hash the dictionary
|
||||
serialized_obj = json.dumps(
|
||||
hash_parts, sort_keys=True
|
||||
) # sort_keys ensures consistent ordering
|
||||
hash_object = hashlib.sha256(serialized_obj.encode())
|
||||
return hash_object.hexdigest()
|
||||
|
||||
@classmethod
|
||||
def update_or_create_by_hash(cls, hash_value, **fields_to_update):
|
||||
try:
|
||||
# Try to find the Stream object with the given hash
|
||||
stream = cls.objects.get(stream_hash=hash_value)
|
||||
# If it exists, update the fields
|
||||
for field, value in fields_to_update.items():
|
||||
setattr(stream, field, value)
|
||||
stream.save() # Save the updated object
|
||||
return stream, False # False means it was updated, not created
|
||||
except cls.DoesNotExist:
|
||||
# If it doesn't exist, create a new object with the given hash
|
||||
fields_to_update["stream_hash"] = (
|
||||
hash_value # Make sure the hash field is set
|
||||
)
|
||||
stream = cls.objects.create(**fields_to_update)
|
||||
return stream, True # True means it was created
|
||||
|
||||
def get_stream_profile(self):
|
||||
"""
|
||||
Get the stream profile for this stream.
|
||||
Uses the stream's own profile if set, otherwise returns the default.
|
||||
"""
|
||||
if self.stream_profile:
|
||||
return self.stream_profile
|
||||
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available profile for this stream and reserves a connection slot.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
profile_id = redis_client.get(f"stream_profile:{self.id}")
|
||||
if profile_id:
|
||||
profile_id = int(profile_id)
|
||||
return self.id, profile_id, None
|
||||
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = self.m3u_account
|
||||
m3u_profiles = m3u_account.profiles.all()
|
||||
default_profile = next((obj for obj in m3u_profiles if obj.is_default), None)
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
logger.info(profile)
|
||||
# Skip inactive profiles
|
||||
if profile.is_active == False:
|
||||
continue
|
||||
|
||||
# Atomic slot reservation: INCR first, check, rollback if over capacity
|
||||
if profile.max_streams == 0:
|
||||
reserved = True
|
||||
else:
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
if new_count <= profile.max_streams:
|
||||
reserved = True
|
||||
else:
|
||||
redis_client.decr(profile_connections_key)
|
||||
reserved = False
|
||||
|
||||
if reserved:
|
||||
redis_client.set(f"channel_stream:{self.id}", self.id)
|
||||
redis_client.set(f"stream_profile:{self.id}", profile.id)
|
||||
return self.id, profile.id, None
|
||||
|
||||
return None, None, "All active M3U profiles have reached maximum connection limits"
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = self.id
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not profile_id:
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: no profile found in "
|
||||
f"stream_profile:{stream_id}"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
|
||||
profile_id = int(profile_id)
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: found profile_id={profile_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
class ChannelManager(models.Manager):
|
||||
def active(self):
|
||||
return self.all()
|
||||
|
||||
|
||||
class Channel(models.Model):
|
||||
channel_number = models.FloatField(db_index=True)
|
||||
name = models.CharField(max_length=255)
|
||||
logo = models.ForeignKey(
|
||||
"Logo",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
# M2M to Stream now in the same file
|
||||
streams = models.ManyToManyField(
|
||||
Stream, blank=True, through="ChannelStream", related_name="channels"
|
||||
)
|
||||
|
||||
channel_group = models.ForeignKey(
|
||||
"ChannelGroup",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
help_text="Channel group this channel belongs to.",
|
||||
)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
tvc_guide_stationid = models.CharField(max_length=255, blank=True, null=True)
|
||||
|
||||
epg_data = models.ForeignKey(
|
||||
EPGData,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
uuid = models.UUIDField(
|
||||
default=uuid.uuid4, editable=False, unique=True, db_index=True
|
||||
)
|
||||
|
||||
user_level = models.IntegerField(default=0)
|
||||
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this channel contains adult content"
|
||||
)
|
||||
|
||||
auto_created = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this channel was automatically created via M3U auto channel sync"
|
||||
)
|
||||
auto_created_by = models.ForeignKey(
|
||||
"m3u.M3UAccount",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="auto_created_channels",
|
||||
help_text="The M3U account that auto-created this channel"
|
||||
)
|
||||
|
||||
created_at = models.DateTimeField(
|
||||
auto_now_add=True,
|
||||
help_text="Timestamp when this channel was created"
|
||||
)
|
||||
updated_at = models.DateTimeField(
|
||||
auto_now=True,
|
||||
help_text="Timestamp when this channel was last updated"
|
||||
)
|
||||
|
||||
def clean(self):
|
||||
# Enforce unique channel_number within a given group
|
||||
existing = Channel.objects.filter(
|
||||
channel_number=self.channel_number, channel_group=self.channel_group
|
||||
).exclude(id=self.id)
|
||||
if existing.exists():
|
||||
raise ValidationError(
|
||||
f"Channel number {self.channel_number} already exists in group {self.channel_group}."
|
||||
)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_number} - {self.name}"
|
||||
|
||||
@classmethod
|
||||
def get_next_available_channel_number(cls, starting_from=1):
|
||||
used_numbers = set(cls.objects.all().values_list("channel_number", flat=True))
|
||||
n = starting_from
|
||||
while n in used_numbers:
|
||||
n += 1
|
||||
return n
|
||||
|
||||
# @TODO: honor stream's stream profile
|
||||
def get_stream_profile(self):
|
||||
stream_profile = self.stream_profile
|
||||
if not stream_profile:
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def _account_active_connections(self, redis_client, m3u_account, cache: dict[int, int]) -> int:
|
||||
"""
|
||||
Return active connection count for an M3U account across all its active profiles.
|
||||
Uses Redis `profile_connections:*`, which is shared by live and VOD proxy paths.
|
||||
"""
|
||||
if not m3u_account:
|
||||
return 0
|
||||
account_id = int(m3u_account.id)
|
||||
if account_id in cache:
|
||||
return cache[account_id]
|
||||
|
||||
total = 0
|
||||
try:
|
||||
for profile in m3u_account.profiles.filter(is_active=True):
|
||||
total += int(redis_client.get(f"profile_connections:{profile.id}") or 0)
|
||||
except Exception:
|
||||
total = 0
|
||||
cache[account_id] = total
|
||||
return total
|
||||
|
||||
def _pick_channel_to_preempt(
|
||||
self,
|
||||
profile_id,
|
||||
requester_level,
|
||||
redis_client,
|
||||
exclude_channel_ids=None,
|
||||
cooldown_seconds=30,
|
||||
):
|
||||
"""
|
||||
Pick the lowest-impact channel to terminate on the given profile.
|
||||
Returns: Optional[int] channel_id to preempt
|
||||
"""
|
||||
exclude_channel_ids = set(exclude_channel_ids or [])
|
||||
candidates = []
|
||||
|
||||
# 1) Try to get active channel IDs for this profile from an index set if available
|
||||
ch_set_key = f"ts_proxy:profile:{profile_id}:channels"
|
||||
try:
|
||||
ch_ids = { (int(x) if not isinstance(x, int) else x) for x in (redis_client.smembers(ch_set_key) or set()) }
|
||||
except Exception:
|
||||
ch_ids = set()
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
# 2) Fallback: scan metadata keys and filter by m3u_profile == profile_id
|
||||
if not ch_ids:
|
||||
cursor = 0
|
||||
pattern = "ts_proxy:channel:*:metadata"
|
||||
while True:
|
||||
cursor, keys = redis_client.scan(cursor=cursor, match=pattern, count=500)
|
||||
if keys:
|
||||
# Prefer HGET m3u_profile if metadata is a hash
|
||||
pipe = redis_client.pipeline()
|
||||
for k in keys:
|
||||
pipe.hget(k, "m3u_profile")
|
||||
prof_vals = pipe.execute()
|
||||
for k, prof_val in zip(keys, prof_vals):
|
||||
try:
|
||||
pid = int(prof_val) if prof_val is not None else None
|
||||
except Exception:
|
||||
pid = None
|
||||
|
||||
if pid == profile_id:
|
||||
parts = k.split(":") # ts_proxy:channel:{id}:metadata
|
||||
if len(parts) >= 4:
|
||||
try:
|
||||
ch_ids.add(int(parts[2]))
|
||||
except Exception:
|
||||
pass
|
||||
if cursor == 0:
|
||||
break
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
if not ch_ids:
|
||||
return None
|
||||
|
||||
# 3) Score candidates
|
||||
for ch_id in ch_ids:
|
||||
if ch_id in exclude_channel_ids:
|
||||
continue
|
||||
|
||||
# Skip if recently preempted
|
||||
last_preempt_key = f"ts_proxy:channel:{ch_id}:last_preempt"
|
||||
try:
|
||||
last_preempt = float(redis_client.get(last_preempt_key) or 0.0)
|
||||
except Exception:
|
||||
last_preempt = 0.0
|
||||
if last_preempt and (time.time() - last_preempt) < cooldown_seconds:
|
||||
continue
|
||||
|
||||
# Clients and their levels
|
||||
clients_key = f"ts_proxy:channel:{ch_id}:clients"
|
||||
member_ids = list(redis_client.smembers(clients_key) or [])
|
||||
viewer_count = len(member_ids)
|
||||
max_viewer_level = 0
|
||||
if viewer_count:
|
||||
pipe = redis_client.pipeline()
|
||||
for cid in member_ids:
|
||||
pipe.hget(f"ts_proxy:channel:{ch_id}:clients:{cid}", "user_level")
|
||||
levels_raw = pipe.execute()
|
||||
levels = []
|
||||
for lv in levels_raw:
|
||||
try:
|
||||
levels.append(int(lv or 0))
|
||||
except Exception:
|
||||
levels.append(0)
|
||||
max_viewer_level = max(levels or [0])
|
||||
|
||||
# Only preempt if requester strictly outranks this channel's viewers
|
||||
if requester_level <= max_viewer_level:
|
||||
continue
|
||||
|
||||
# Metadata (protected/recording/started_at_ts)
|
||||
meta_key = f"ts_proxy:channel:{ch_id}:metadata"
|
||||
try:
|
||||
protected, recording, started_at_ts = redis_client.hmget(
|
||||
meta_key, "protected", "recording", "started_at_ts"
|
||||
)
|
||||
except Exception:
|
||||
protected = recording = started_at_ts = None
|
||||
|
||||
protected = str(protected or "0") in ("1", "true", "True")
|
||||
recording = str(recording or "0") in ("1", "true", "True")
|
||||
if protected or recording:
|
||||
continue
|
||||
|
||||
try:
|
||||
started_at_ts = float(started_at_ts) if started_at_ts is not None else None
|
||||
except Exception:
|
||||
started_at_ts = None
|
||||
if started_at_ts is None:
|
||||
started_at_ts = time.time() # treat unknown as newest
|
||||
|
||||
# Score: lower is safer to terminate
|
||||
has_viewers = 1 if viewer_count > 0 else 0
|
||||
score = (has_viewers, max_viewer_level, viewer_count, started_at_ts)
|
||||
candidates.append((score, ch_id))
|
||||
|
||||
logger.debug("Candidate channels after scoring:")
|
||||
logger.debug(candidates)
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
victim_id = candidates[0][1]
|
||||
|
||||
# Mark preempt timestamp to avoid thrashing
|
||||
try:
|
||||
redis_client.set(f"ts_proxy:channel:{victim_id}:last_preempt", str(time.time()), ex=3600)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return victim_id
|
||||
|
||||
def _check_and_reserve_profile_slot(self, profile, redis_client):
|
||||
"""
|
||||
Atomically check and reserve a connection slot for the given profile.
|
||||
|
||||
Uses an INCR-first-then-check pattern to eliminate the TOCTOU race
|
||||
condition where separate GET + check + INCR operations could allow
|
||||
concurrent requests to both pass the capacity check.
|
||||
|
||||
For profiles with max_streams=0 (unlimited), no reservation is needed.
|
||||
|
||||
Args:
|
||||
profile: M3UAccountProfile instance
|
||||
redis_client: Redis client instance
|
||||
|
||||
Returns:
|
||||
tuple: (reserved: bool, current_count: int)
|
||||
"""
|
||||
if profile.max_streams == 0:
|
||||
return (True, 0)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
|
||||
# Atomically increment first — this is a single Redis command
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
|
||||
if new_count <= profile.max_streams:
|
||||
return (True, new_count)
|
||||
|
||||
# Over capacity — roll back the increment
|
||||
redis_client.decr(profile_connections_key)
|
||||
return (False, new_count - 1)
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available stream for the requested channel and returns the selected stream and profile.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
error_reason = None
|
||||
|
||||
# Check if this channel has any streams
|
||||
if not self.streams.exists():
|
||||
error_reason = "No streams assigned to channel"
|
||||
return None, None, error_reason
|
||||
|
||||
# Check if a stream is already active for this channel
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if stream_id_bytes:
|
||||
try:
|
||||
stream_id = int(stream_id_bytes)
|
||||
profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id_bytes:
|
||||
try:
|
||||
profile_id = int(profile_id_bytes)
|
||||
return stream_id, profile_id, None
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid profile ID retrieved from Redis: {profile_id_bytes}"
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid stream ID retrieved from Redis: {stream_id_bytes}"
|
||||
)
|
||||
|
||||
# No existing active stream, attempt to assign a new one
|
||||
has_streams_but_maxed_out = False
|
||||
has_active_profiles = False
|
||||
account_load_cache: dict[int, int] = {}
|
||||
|
||||
ordered_streams = list(self.streams.all().order_by("channelstream__order"))
|
||||
original_order = {stream.id: idx for idx, stream in enumerate(ordered_streams)}
|
||||
ordered_streams.sort(
|
||||
key=lambda s: (
|
||||
self._account_active_connections(redis_client, s.m3u_account, account_load_cache),
|
||||
original_order.get(s.id, 0),
|
||||
)
|
||||
)
|
||||
|
||||
# Iterate through channel streams and their profiles
|
||||
for stream in ordered_streams:
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = stream.m3u_account
|
||||
if not m3u_account:
|
||||
logger.debug(f"Stream {stream.id} has no M3U account")
|
||||
continue
|
||||
if m3u_account.is_active == False:
|
||||
logger.debug(f"M3U account {m3u_account.id} is inactive, skipping.")
|
||||
continue
|
||||
|
||||
m3u_profiles = m3u_account.profiles.filter(is_active=True)
|
||||
default_profile = next(
|
||||
(obj for obj in m3u_profiles if obj.is_default), None
|
||||
)
|
||||
|
||||
if not default_profile:
|
||||
logger.debug(f"M3U account {m3u_account.id} has no active default profile")
|
||||
continue
|
||||
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
has_active_profiles = True
|
||||
|
||||
# Atomically check and reserve a slot (INCR-first pattern)
|
||||
reserved, current_count = self._check_and_reserve_profile_slot(
|
||||
profile, redis_client
|
||||
)
|
||||
|
||||
if reserved:
|
||||
# Slot reserved — assign stream to this channel
|
||||
redis_client.set(f"channel_stream:{self.id}", stream.id)
|
||||
redis_client.set(f"stream_profile:{stream.id}", profile.id)
|
||||
|
||||
return (
|
||||
stream.id,
|
||||
profile.id,
|
||||
None,
|
||||
) # Return newly assigned stream and matched profile
|
||||
else:
|
||||
# At capacity: try to preempt a lower-impact channel on this profile
|
||||
victim_channel_id = self._pick_channel_to_preempt(
|
||||
profile_id=profile.id,
|
||||
requester_level=requester.user_level if requester else 100,
|
||||
redis_client=redis_client,
|
||||
exclude_channel_ids=None,
|
||||
)
|
||||
if victim_channel_id:
|
||||
logger.info(f"Preempting channel {victim_channel_id} for new stream on profile {profile.id}")
|
||||
# return self.id, profile.id, victim_channel_id
|
||||
|
||||
# This profile is at max connections
|
||||
has_streams_but_maxed_out = True
|
||||
logger.debug(
|
||||
f"Profile {profile.id} at max connections: "
|
||||
f"{current_count}/{profile.max_streams}"
|
||||
)
|
||||
|
||||
# No available streams - determine specific reason
|
||||
if has_streams_but_maxed_out:
|
||||
error_reason = "All active M3U profiles have reached maximum connection limits"
|
||||
elif has_active_profiles:
|
||||
error_reason = "No compatible active profile found for any assigned stream"
|
||||
else:
|
||||
error_reason = "No active profiles found for any assigned stream"
|
||||
|
||||
return None, None, error_reason
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no stream/profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id:
|
||||
# Primary key missing — try metadata hash fallback.
|
||||
# The proxy may have already cleaned up channel_stream/stream_profile
|
||||
# keys, but the metadata hash can still have the stream_id and profile.
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_stream_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.STREAM_ID
|
||||
)
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
|
||||
if meta_stream_id and meta_profile_id:
|
||||
stream_id = int(meta_stream_id)
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered stream_id={stream_id}, "
|
||||
f"profile_id={profile_id} from metadata fallback"
|
||||
)
|
||||
# Clean up any remaining keys
|
||||
redis_client.delete(f"channel_stream:{self.id}")
|
||||
redis_client.delete(f"stream_profile:{stream_id}")
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# won't find them and DECR again
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
current_count = int(
|
||||
redis_client.get(profile_connections_key) or 0
|
||||
)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
return True
|
||||
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: no stream info found in primary keys "
|
||||
f"or metadata fallback"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"channel_stream:{self.id}") # Remove active stream
|
||||
|
||||
stream_id = int(stream_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found stream_id={stream_id} for "
|
||||
f"channel_stream:{self.id}"
|
||||
)
|
||||
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id:
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
profile_id = int(profile_id)
|
||||
else:
|
||||
# stream_profile key missing — try metadata hash fallback
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
if meta_profile_id:
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered profile_id={profile_id} "
|
||||
f"from metadata fallback (stream_profile:{stream_id} was missing)"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"Channel {self.uuid}: no profile found for "
|
||||
f"stream_profile:{stream_id} or in metadata fallback"
|
||||
)
|
||||
return False
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found profile_id={profile_id} for "
|
||||
f"stream {stream_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# (e.g. from _clean_redis_keys or ChannelService.stop_channel)
|
||||
# won't find them via fallback and DECR again
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
def update_stream_profile(self, new_profile_id):
|
||||
"""
|
||||
Updates the profile for the current stream and adjusts connection counts.
|
||||
|
||||
Args:
|
||||
new_profile_id: The ID of the new stream profile to use
|
||||
|
||||
Returns:
|
||||
bool: True if successful, False otherwise
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
# Get current stream ID
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id_bytes:
|
||||
logger.debug("No active stream found for channel")
|
||||
return False
|
||||
|
||||
stream_id = int(stream_id_bytes)
|
||||
|
||||
# Get current profile ID
|
||||
current_profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not current_profile_id_bytes:
|
||||
logger.debug("No profile found for current stream")
|
||||
return False
|
||||
|
||||
current_profile_id = int(current_profile_id_bytes)
|
||||
|
||||
# Don't do anything if the profile is already set to the requested one
|
||||
if current_profile_id == new_profile_id:
|
||||
return True
|
||||
|
||||
# Use pipeline for atomic profile switch to prevent counter drift
|
||||
# if an exception occurs between DECR and INCR
|
||||
old_profile_connections_key = f"profile_connections:{current_profile_id}"
|
||||
new_profile_connections_key = f"profile_connections:{new_profile_id}"
|
||||
old_count = int(redis_client.get(old_profile_connections_key) or 0)
|
||||
|
||||
pipe = redis_client.pipeline()
|
||||
if old_count > 0:
|
||||
pipe.decr(old_profile_connections_key)
|
||||
pipe.set(f"stream_profile:{stream_id}", new_profile_id)
|
||||
pipe.incr(new_profile_connections_key)
|
||||
pipe.execute()
|
||||
logger.info(
|
||||
f"Updated stream {stream_id} profile from {current_profile_id} to {new_profile_id}"
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
class ChannelProfile(models.Model):
|
||||
name = models.CharField(max_length=100, unique=True)
|
||||
|
||||
|
||||
class ChannelProfileMembership(models.Model):
|
||||
channel_profile = models.ForeignKey(ChannelProfile, on_delete=models.CASCADE)
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
enabled = models.BooleanField(
|
||||
default=True
|
||||
) # Track if the channel is enabled for this group
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_profile", "channel")
|
||||
|
||||
|
||||
class ChannelStream(models.Model):
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
stream = models.ForeignKey(Stream, on_delete=models.CASCADE)
|
||||
order = models.PositiveIntegerField(default=0) # Ordering field
|
||||
|
||||
class Meta:
|
||||
ordering = ["order"] # Ensure streams are retrieved in order
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=["channel", "stream"], name="unique_channel_stream"
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class ChannelGroupM3UAccount(models.Model):
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup, on_delete=models.CASCADE, related_name="m3u_accounts"
|
||||
)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount, on_delete=models.CASCADE, related_name="channel_group"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
enabled = models.BooleanField(default=True)
|
||||
auto_channel_sync = models.BooleanField(
|
||||
default=False,
|
||||
help_text='Automatically create/delete channels to match streams in this group'
|
||||
)
|
||||
auto_sync_channel_start = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text='Starting channel number for auto-created channels in this group'
|
||||
)
|
||||
last_seen = models.DateTimeField(
|
||||
default=timezone.now,
|
||||
db_index=True,
|
||||
help_text='Last time this group was seen in the M3U source during a refresh'
|
||||
)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text='Whether this group relationship is stale (not seen in recent refresh, pending deletion)'
|
||||
)
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_group", "m3u_account")
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_group.name} - {self.m3u_account.name} (Enabled: {self.enabled})"
|
||||
|
||||
|
||||
class Logo(models.Model):
|
||||
name = models.CharField(max_length=255)
|
||||
url = models.TextField(unique=True)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
|
||||
class Recording(models.Model):
|
||||
channel = models.ForeignKey(
|
||||
"Channel", on_delete=models.CASCADE, related_name="recordings"
|
||||
)
|
||||
start_time = models.DateTimeField()
|
||||
end_time = models.DateTimeField()
|
||||
task_id = models.CharField(max_length=255, null=True, blank=True)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel.name} - {self.start_time} to {self.end_time}"
|
||||
|
||||
|
||||
class RecurringRecordingRule(models.Model):
|
||||
"""Rule describing a recurring manual DVR schedule."""
|
||||
|
||||
channel = models.ForeignKey(
|
||||
"Channel",
|
||||
on_delete=models.CASCADE,
|
||||
related_name="recurring_rules",
|
||||
)
|
||||
days_of_week = models.JSONField(default=list)
|
||||
start_time = models.TimeField()
|
||||
end_time = models.TimeField()
|
||||
enabled = models.BooleanField(default=True)
|
||||
name = models.CharField(max_length=255, blank=True)
|
||||
start_date = models.DateField(null=True, blank=True)
|
||||
end_date = models.DateField(null=True, blank=True)
|
||||
created_at = models.DateTimeField(auto_now_add=True)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
|
||||
class Meta:
|
||||
ordering = ["channel", "start_time"]
|
||||
|
||||
def __str__(self):
|
||||
channel_name = getattr(self.channel, "name", str(self.channel_id))
|
||||
return f"Recurring rule for {channel_name}"
|
||||
|
||||
def cleaned_days(self):
|
||||
try:
|
||||
return sorted({int(d) for d in (self.days_of_week or []) if 0 <= int(d) <= 6})
|
||||
except Exception:
|
||||
return []
|
||||
@@ -0,0 +1,958 @@
|
||||
from django.db import models
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.conf import settings
|
||||
from core.models import StreamProfile, CoreSettings
|
||||
from core.utils import RedisClient
|
||||
from apps.proxy.ts_proxy.redis_keys import RedisKeys
|
||||
from apps.proxy.ts_proxy.constants import ChannelMetadataField
|
||||
import logging
|
||||
import uuid
|
||||
from django.utils import timezone
|
||||
import hashlib
|
||||
import json
|
||||
from apps.epg.models import EPGData
|
||||
from apps.accounts.models import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# If you have an M3UAccount model in apps.m3u, you can still import it:
|
||||
from apps.m3u.models import M3UAccount
|
||||
|
||||
|
||||
# Add fallback functions if Redis isn't available
|
||||
def get_total_viewers(channel_id):
|
||||
"""Get viewer count from Redis or return 0 if Redis isn't available"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
try:
|
||||
return int(redis_client.get(f"channel:{channel_id}:viewers") or 0)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
class ChannelGroup(models.Model):
|
||||
name = models.TextField(unique=True, db_index=True)
|
||||
|
||||
def related_channels(self):
|
||||
# local import if needed to avoid cyc. Usually fine in a single file though
|
||||
return Channel.objects.filter(channel_group=self)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def bulk_create_and_fetch(cls, objects):
|
||||
# Perform the bulk create operation
|
||||
cls.objects.bulk_create(objects)
|
||||
|
||||
# Use a unique field to fetch the created objects (assuming 'name' is unique)
|
||||
created_objects = cls.objects.filter(name__in=[obj.name for obj in objects])
|
||||
|
||||
return created_objects
|
||||
|
||||
|
||||
class Stream(models.Model):
|
||||
"""
|
||||
Represents a single stream (e.g. from an M3U source or custom URL).
|
||||
"""
|
||||
|
||||
name = models.CharField(max_length=255, default="Default Stream")
|
||||
url = models.URLField(max_length=4096, blank=True, null=True)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount,
|
||||
on_delete=models.CASCADE,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
logo_url = models.TextField(blank=True, null=True)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
local_file = models.FileField(upload_to="uploads/", blank=True, null=True)
|
||||
current_viewers = models.PositiveIntegerField(default=0)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="streams",
|
||||
)
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
null=True,
|
||||
blank=True,
|
||||
on_delete=models.SET_NULL,
|
||||
related_name="streams",
|
||||
)
|
||||
is_custom = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this is a user-created stream or from an M3U account",
|
||||
)
|
||||
stream_hash = models.CharField(
|
||||
max_length=255,
|
||||
null=True,
|
||||
unique=True,
|
||||
help_text="Unique hash for this stream from the M3U account",
|
||||
db_index=True,
|
||||
)
|
||||
last_seen = models.DateTimeField(db_index=True, default=timezone.now)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream is stale (not seen in recent refresh, pending deletion)"
|
||||
)
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this stream contains adult content"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
stream_id = models.IntegerField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider stream ID (e.g., XC stream_id) for stable identity across credential changes"
|
||||
)
|
||||
stream_chno = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
db_index=True,
|
||||
help_text="Provider channel number (XC num or M3U tvg-chno) for ordering - supports decimals like 2.1"
|
||||
)
|
||||
|
||||
# Stream statistics fields
|
||||
stream_stats = models.JSONField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="JSON object containing stream statistics like video codec, resolution, etc."
|
||||
)
|
||||
stream_stats_updated_at = models.DateTimeField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text="When stream statistics were last updated",
|
||||
db_index=True
|
||||
)
|
||||
|
||||
class Meta:
|
||||
# If you use m3u_account, you might do unique_together = ('name','url','m3u_account')
|
||||
verbose_name = "Stream"
|
||||
verbose_name_plural = "Streams"
|
||||
ordering = ["-updated_at"]
|
||||
|
||||
def __str__(self):
|
||||
return self.name or self.url or f"Stream ID {self.id}"
|
||||
|
||||
@classmethod
|
||||
def generate_hash_key(cls, name, url, tvg_id, keys=None, m3u_id=None, group=None,
|
||||
account_type=None, stream_id=None):
|
||||
if keys is None:
|
||||
keys = CoreSettings.get_m3u_hash_key().split(",")
|
||||
|
||||
# For XC accounts, use stream_id instead of url when 'url' is in the hash keys
|
||||
# This ensures credential/URL changes don't break stream identity
|
||||
effective_url = url
|
||||
use_stream_id = account_type == 'XC' and stream_id and 'url' in keys
|
||||
if use_stream_id:
|
||||
effective_url = stream_id
|
||||
|
||||
stream_parts = {"name": name, "url": effective_url, "tvg_id": tvg_id, "m3u_id": m3u_id, "group": group}
|
||||
|
||||
hash_parts = {key: stream_parts[key] for key in keys if key in stream_parts}
|
||||
|
||||
# When using stream_id instead of URL, we MUST include m3u_id to prevent
|
||||
# collisions across different XC accounts (stream_id is only unique per account)
|
||||
if use_stream_id and 'm3u_id' not in hash_parts:
|
||||
hash_parts['m3u_id'] = m3u_id
|
||||
|
||||
# Serialize and hash the dictionary
|
||||
serialized_obj = json.dumps(
|
||||
hash_parts, sort_keys=True
|
||||
) # sort_keys ensures consistent ordering
|
||||
hash_object = hashlib.sha256(serialized_obj.encode())
|
||||
return hash_object.hexdigest()
|
||||
|
||||
@classmethod
|
||||
def update_or_create_by_hash(cls, hash_value, **fields_to_update):
|
||||
try:
|
||||
# Try to find the Stream object with the given hash
|
||||
stream = cls.objects.get(stream_hash=hash_value)
|
||||
# If it exists, update the fields
|
||||
for field, value in fields_to_update.items():
|
||||
setattr(stream, field, value)
|
||||
stream.save() # Save the updated object
|
||||
return stream, False # False means it was updated, not created
|
||||
except cls.DoesNotExist:
|
||||
# If it doesn't exist, create a new object with the given hash
|
||||
fields_to_update["stream_hash"] = (
|
||||
hash_value # Make sure the hash field is set
|
||||
)
|
||||
stream = cls.objects.create(**fields_to_update)
|
||||
return stream, True # True means it was created
|
||||
|
||||
def get_stream_profile(self):
|
||||
"""
|
||||
Get the stream profile for this stream.
|
||||
Uses the stream's own profile if set, otherwise returns the default.
|
||||
"""
|
||||
if self.stream_profile:
|
||||
return self.stream_profile
|
||||
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available profile for this stream and reserves a connection slot.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
profile_id = redis_client.get(f"stream_profile:{self.id}")
|
||||
if profile_id:
|
||||
profile_id = int(profile_id)
|
||||
return self.id, profile_id, None
|
||||
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = self.m3u_account
|
||||
m3u_profiles = m3u_account.profiles.all()
|
||||
default_profile = next((obj for obj in m3u_profiles if obj.is_default), None)
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
logger.info(profile)
|
||||
# Skip inactive profiles
|
||||
if profile.is_active == False:
|
||||
continue
|
||||
|
||||
# Atomic slot reservation: INCR first, check, rollback if over capacity
|
||||
if profile.max_streams == 0:
|
||||
reserved = True
|
||||
else:
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
if new_count <= profile.max_streams:
|
||||
reserved = True
|
||||
else:
|
||||
redis_client.decr(profile_connections_key)
|
||||
reserved = False
|
||||
|
||||
if reserved:
|
||||
redis_client.set(f"channel_stream:{self.id}", self.id)
|
||||
redis_client.set(f"stream_profile:{self.id}", profile.id)
|
||||
return self.id, profile.id, None
|
||||
|
||||
return None, None, "All active M3U profiles have reached maximum connection limits"
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = self.id
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not profile_id:
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: no profile found in "
|
||||
f"stream_profile:{stream_id}"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
|
||||
profile_id = int(profile_id)
|
||||
logger.debug(
|
||||
f"Stream {stream_id}: found profile_id={profile_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
class ChannelManager(models.Manager):
|
||||
def active(self):
|
||||
return self.all()
|
||||
|
||||
|
||||
class Channel(models.Model):
|
||||
channel_number = models.FloatField(db_index=True)
|
||||
name = models.CharField(max_length=255)
|
||||
logo = models.ForeignKey(
|
||||
"Logo",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
# M2M to Stream now in the same file
|
||||
streams = models.ManyToManyField(
|
||||
Stream, blank=True, through="ChannelStream", related_name="channels"
|
||||
)
|
||||
|
||||
channel_group = models.ForeignKey(
|
||||
"ChannelGroup",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
help_text="Channel group this channel belongs to.",
|
||||
)
|
||||
tvg_id = models.CharField(max_length=255, blank=True, null=True)
|
||||
tvc_guide_stationid = models.CharField(max_length=255, blank=True, null=True)
|
||||
|
||||
epg_data = models.ForeignKey(
|
||||
EPGData,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
stream_profile = models.ForeignKey(
|
||||
StreamProfile,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="channels",
|
||||
)
|
||||
|
||||
uuid = models.UUIDField(
|
||||
default=uuid.uuid4, editable=False, unique=True, db_index=True
|
||||
)
|
||||
|
||||
user_level = models.IntegerField(default=0)
|
||||
|
||||
is_adult = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text="Whether this channel contains adult content"
|
||||
)
|
||||
|
||||
auto_created = models.BooleanField(
|
||||
default=False,
|
||||
help_text="Whether this channel was automatically created via M3U auto channel sync"
|
||||
)
|
||||
auto_created_by = models.ForeignKey(
|
||||
"m3u.M3UAccount",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="auto_created_channels",
|
||||
help_text="The M3U account that auto-created this channel"
|
||||
)
|
||||
|
||||
created_at = models.DateTimeField(
|
||||
auto_now_add=True,
|
||||
help_text="Timestamp when this channel was created"
|
||||
)
|
||||
updated_at = models.DateTimeField(
|
||||
auto_now=True,
|
||||
help_text="Timestamp when this channel was last updated"
|
||||
)
|
||||
|
||||
def clean(self):
|
||||
# Enforce unique channel_number within a given group
|
||||
existing = Channel.objects.filter(
|
||||
channel_number=self.channel_number, channel_group=self.channel_group
|
||||
).exclude(id=self.id)
|
||||
if existing.exists():
|
||||
raise ValidationError(
|
||||
f"Channel number {self.channel_number} already exists in group {self.channel_group}."
|
||||
)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_number} - {self.name}"
|
||||
|
||||
@classmethod
|
||||
def get_next_available_channel_number(cls, starting_from=1):
|
||||
used_numbers = set(cls.objects.all().values_list("channel_number", flat=True))
|
||||
n = starting_from
|
||||
while n in used_numbers:
|
||||
n += 1
|
||||
return n
|
||||
|
||||
# @TODO: honor stream's stream profile
|
||||
def get_stream_profile(self):
|
||||
stream_profile = self.stream_profile
|
||||
if not stream_profile:
|
||||
stream_profile = StreamProfile.objects.get(
|
||||
id=CoreSettings.get_default_stream_profile_id()
|
||||
)
|
||||
|
||||
return stream_profile
|
||||
|
||||
def _pick_channel_to_preempt(
|
||||
self,
|
||||
profile_id,
|
||||
requester_level,
|
||||
redis_client,
|
||||
exclude_channel_ids=None,
|
||||
cooldown_seconds=30,
|
||||
):
|
||||
"""
|
||||
Pick the lowest-impact channel to terminate on the given profile.
|
||||
Returns: Optional[int] channel_id to preempt
|
||||
"""
|
||||
exclude_channel_ids = set(exclude_channel_ids or [])
|
||||
candidates = []
|
||||
|
||||
# 1) Try to get active channel IDs for this profile from an index set if available
|
||||
ch_set_key = f"ts_proxy:profile:{profile_id}:channels"
|
||||
try:
|
||||
ch_ids = { (int(x) if not isinstance(x, int) else x) for x in (redis_client.smembers(ch_set_key) or set()) }
|
||||
except Exception:
|
||||
ch_ids = set()
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
# 2) Fallback: scan metadata keys and filter by m3u_profile == profile_id
|
||||
if not ch_ids:
|
||||
cursor = 0
|
||||
pattern = "ts_proxy:channel:*:metadata"
|
||||
while True:
|
||||
cursor, keys = redis_client.scan(cursor=cursor, match=pattern, count=500)
|
||||
if keys:
|
||||
# Prefer HGET m3u_profile if metadata is a hash
|
||||
pipe = redis_client.pipeline()
|
||||
for k in keys:
|
||||
pipe.hget(k, "m3u_profile")
|
||||
prof_vals = pipe.execute()
|
||||
for k, prof_val in zip(keys, prof_vals):
|
||||
try:
|
||||
pid = int(prof_val) if prof_val is not None else None
|
||||
except Exception:
|
||||
pid = None
|
||||
|
||||
if pid == profile_id:
|
||||
parts = k.split(":") # ts_proxy:channel:{id}:metadata
|
||||
if len(parts) >= 4:
|
||||
try:
|
||||
ch_ids.add(int(parts[2]))
|
||||
except Exception:
|
||||
pass
|
||||
if cursor == 0:
|
||||
break
|
||||
|
||||
logger.debug("Candidate channels for preemption:")
|
||||
logger.debug(ch_ids)
|
||||
|
||||
if not ch_ids:
|
||||
return None
|
||||
|
||||
# 3) Score candidates
|
||||
for ch_id in ch_ids:
|
||||
if ch_id in exclude_channel_ids:
|
||||
continue
|
||||
|
||||
# Skip if recently preempted
|
||||
last_preempt_key = f"ts_proxy:channel:{ch_id}:last_preempt"
|
||||
try:
|
||||
last_preempt = float(redis_client.get(last_preempt_key) or 0.0)
|
||||
except Exception:
|
||||
last_preempt = 0.0
|
||||
if last_preempt and (time.time() - last_preempt) < cooldown_seconds:
|
||||
continue
|
||||
|
||||
# Clients and their levels
|
||||
clients_key = f"ts_proxy:channel:{ch_id}:clients"
|
||||
member_ids = list(redis_client.smembers(clients_key) or [])
|
||||
viewer_count = len(member_ids)
|
||||
max_viewer_level = 0
|
||||
if viewer_count:
|
||||
pipe = redis_client.pipeline()
|
||||
for cid in member_ids:
|
||||
pipe.hget(f"ts_proxy:channel:{ch_id}:clients:{cid}", "user_level")
|
||||
levels_raw = pipe.execute()
|
||||
levels = []
|
||||
for lv in levels_raw:
|
||||
try:
|
||||
levels.append(int(lv or 0))
|
||||
except Exception:
|
||||
levels.append(0)
|
||||
max_viewer_level = max(levels or [0])
|
||||
|
||||
# Only preempt if requester strictly outranks this channel's viewers
|
||||
if requester_level <= max_viewer_level:
|
||||
continue
|
||||
|
||||
# Metadata (protected/recording/started_at_ts)
|
||||
meta_key = f"ts_proxy:channel:{ch_id}:metadata"
|
||||
try:
|
||||
protected, recording, started_at_ts = redis_client.hmget(
|
||||
meta_key, "protected", "recording", "started_at_ts"
|
||||
)
|
||||
except Exception:
|
||||
protected = recording = started_at_ts = None
|
||||
|
||||
protected = str(protected or "0") in ("1", "true", "True")
|
||||
recording = str(recording or "0") in ("1", "true", "True")
|
||||
if protected or recording:
|
||||
continue
|
||||
|
||||
try:
|
||||
started_at_ts = float(started_at_ts) if started_at_ts is not None else None
|
||||
except Exception:
|
||||
started_at_ts = None
|
||||
if started_at_ts is None:
|
||||
started_at_ts = time.time() # treat unknown as newest
|
||||
|
||||
# Score: lower is safer to terminate
|
||||
has_viewers = 1 if viewer_count > 0 else 0
|
||||
score = (has_viewers, max_viewer_level, viewer_count, started_at_ts)
|
||||
candidates.append((score, ch_id))
|
||||
|
||||
logger.debug("Candidate channels after scoring:")
|
||||
logger.debug(candidates)
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
candidates.sort(key=lambda x: x[0])
|
||||
victim_id = candidates[0][1]
|
||||
|
||||
# Mark preempt timestamp to avoid thrashing
|
||||
try:
|
||||
redis_client.set(f"ts_proxy:channel:{victim_id}:last_preempt", str(time.time()), ex=3600)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return victim_id
|
||||
|
||||
def _check_and_reserve_profile_slot(self, profile, redis_client):
|
||||
"""
|
||||
Atomically check and reserve a connection slot for the given profile.
|
||||
|
||||
Uses an INCR-first-then-check pattern to eliminate the TOCTOU race
|
||||
condition where separate GET + check + INCR operations could allow
|
||||
concurrent requests to both pass the capacity check.
|
||||
|
||||
For profiles with max_streams=0 (unlimited), no reservation is needed.
|
||||
|
||||
Args:
|
||||
profile: M3UAccountProfile instance
|
||||
redis_client: Redis client instance
|
||||
|
||||
Returns:
|
||||
tuple: (reserved: bool, current_count: int)
|
||||
"""
|
||||
if profile.max_streams == 0:
|
||||
return (True, 0)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile.id}"
|
||||
|
||||
# Atomically increment first — this is a single Redis command
|
||||
new_count = redis_client.incr(profile_connections_key)
|
||||
|
||||
if new_count <= profile.max_streams:
|
||||
return (True, new_count)
|
||||
|
||||
# Over capacity — roll back the increment
|
||||
redis_client.decr(profile_connections_key)
|
||||
return (False, new_count - 1)
|
||||
|
||||
def get_stream(self, requester=None):
|
||||
"""
|
||||
Finds an available stream for the requested channel and returns the selected stream and profile.
|
||||
|
||||
Returns:
|
||||
Tuple[Optional[int], Optional[int], Optional[str]]: (stream_id, profile_id, error_reason)
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
error_reason = None
|
||||
|
||||
# Check if this channel has any streams
|
||||
if not self.streams.exists():
|
||||
error_reason = "No streams assigned to channel"
|
||||
return None, None, error_reason
|
||||
|
||||
# Check if a stream is already active for this channel
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if stream_id_bytes:
|
||||
try:
|
||||
stream_id = int(stream_id_bytes)
|
||||
profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id_bytes:
|
||||
try:
|
||||
profile_id = int(profile_id_bytes)
|
||||
return stream_id, profile_id, None
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid profile ID retrieved from Redis: {profile_id_bytes}"
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
f"Invalid stream ID retrieved from Redis: {stream_id_bytes}"
|
||||
)
|
||||
|
||||
# No existing active stream, attempt to assign a new one
|
||||
has_streams_but_maxed_out = False
|
||||
has_active_profiles = False
|
||||
|
||||
# Iterate through channel streams and their profiles
|
||||
for stream in self.streams.all().order_by("channelstream__order"):
|
||||
# Retrieve the M3U account associated with the stream.
|
||||
m3u_account = stream.m3u_account
|
||||
if not m3u_account:
|
||||
logger.debug(f"Stream {stream.id} has no M3U account")
|
||||
continue
|
||||
if m3u_account.is_active == False:
|
||||
logger.debug(f"M3U account {m3u_account.id} is inactive, skipping.")
|
||||
continue
|
||||
|
||||
m3u_profiles = m3u_account.profiles.filter(is_active=True)
|
||||
default_profile = next(
|
||||
(obj for obj in m3u_profiles if obj.is_default), None
|
||||
)
|
||||
|
||||
if not default_profile:
|
||||
logger.debug(f"M3U account {m3u_account.id} has no active default profile")
|
||||
continue
|
||||
|
||||
profiles = [default_profile] + [
|
||||
obj for obj in m3u_profiles if not obj.is_default
|
||||
]
|
||||
|
||||
for profile in profiles:
|
||||
has_active_profiles = True
|
||||
|
||||
# Atomically check and reserve a slot (INCR-first pattern)
|
||||
reserved, current_count = self._check_and_reserve_profile_slot(
|
||||
profile, redis_client
|
||||
)
|
||||
|
||||
if reserved:
|
||||
# Slot reserved — assign stream to this channel
|
||||
redis_client.set(f"channel_stream:{self.id}", stream.id)
|
||||
redis_client.set(f"stream_profile:{stream.id}", profile.id)
|
||||
|
||||
return (
|
||||
stream.id,
|
||||
profile.id,
|
||||
None,
|
||||
) # Return newly assigned stream and matched profile
|
||||
else:
|
||||
# At capacity: try to preempt a lower-impact channel on this profile
|
||||
victim_channel_id = self._pick_channel_to_preempt(
|
||||
profile_id=profile.id,
|
||||
requester_level=requester.user_level if requester else 100,
|
||||
redis_client=redis_client,
|
||||
exclude_channel_ids=None,
|
||||
)
|
||||
if victim_channel_id:
|
||||
logger.info(f"Preempting channel {victim_channel_id} for new stream on profile {profile.id}")
|
||||
# return self.id, profile.id, victim_channel_id
|
||||
|
||||
# This profile is at max connections
|
||||
has_streams_but_maxed_out = True
|
||||
logger.debug(
|
||||
f"Profile {profile.id} at max connections: "
|
||||
f"{current_count}/{profile.max_streams}"
|
||||
)
|
||||
|
||||
# No available streams - determine specific reason
|
||||
if has_streams_but_maxed_out:
|
||||
error_reason = "All active M3U profiles have reached maximum connection limits"
|
||||
elif has_active_profiles:
|
||||
error_reason = "No compatible active profile found for any assigned stream"
|
||||
else:
|
||||
error_reason = "No active profiles found for any assigned stream"
|
||||
|
||||
return None, None, error_reason
|
||||
|
||||
def release_stream(self):
|
||||
"""
|
||||
Called when a stream is finished to release the lock.
|
||||
|
||||
Returns:
|
||||
bool: True if stream was successfully released, False if
|
||||
no stream/profile info could be found for cleanup.
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
stream_id = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id:
|
||||
# Primary key missing — try metadata hash fallback.
|
||||
# The proxy may have already cleaned up channel_stream/stream_profile
|
||||
# keys, but the metadata hash can still have the stream_id and profile.
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_stream_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.STREAM_ID
|
||||
)
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
|
||||
if meta_stream_id and meta_profile_id:
|
||||
stream_id = int(meta_stream_id)
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered stream_id={stream_id}, "
|
||||
f"profile_id={profile_id} from metadata fallback"
|
||||
)
|
||||
# Clean up any remaining keys
|
||||
redis_client.delete(f"channel_stream:{self.id}")
|
||||
redis_client.delete(f"stream_profile:{stream_id}")
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# won't find them and DECR again
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
current_count = int(
|
||||
redis_client.get(profile_connections_key) or 0
|
||||
)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
return True
|
||||
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: no stream info found in primary keys "
|
||||
f"or metadata fallback"
|
||||
)
|
||||
return False
|
||||
|
||||
redis_client.delete(f"channel_stream:{self.id}") # Remove active stream
|
||||
|
||||
stream_id = int(stream_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found stream_id={stream_id} for "
|
||||
f"channel_stream:{self.id}"
|
||||
)
|
||||
|
||||
# Get the matched profile for cleanup
|
||||
profile_id = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if profile_id:
|
||||
redis_client.delete(f"stream_profile:{stream_id}") # Remove profile association
|
||||
profile_id = int(profile_id)
|
||||
else:
|
||||
# stream_profile key missing — try metadata hash fallback
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
meta_profile_id = redis_client.hget(
|
||||
metadata_key, ChannelMetadataField.M3U_PROFILE
|
||||
)
|
||||
if meta_profile_id:
|
||||
profile_id = int(meta_profile_id)
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: recovered profile_id={profile_id} "
|
||||
f"from metadata fallback (stream_profile:{stream_id} was missing)"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"Channel {self.uuid}: no profile found for "
|
||||
f"stream_profile:{stream_id} or in metadata fallback"
|
||||
)
|
||||
return False
|
||||
logger.debug(
|
||||
f"Channel {self.uuid}: found profile_id={profile_id} for "
|
||||
f"stream {stream_id}"
|
||||
)
|
||||
|
||||
profile_connections_key = f"profile_connections:{profile_id}"
|
||||
|
||||
# Only decrement if the profile had a max_connections limit
|
||||
current_count = int(redis_client.get(profile_connections_key) or 0)
|
||||
if current_count > 0:
|
||||
redis_client.decr(profile_connections_key)
|
||||
|
||||
# Clear metadata fields so duplicate release_stream() calls
|
||||
# (e.g. from _clean_redis_keys or ChannelService.stop_channel)
|
||||
# won't find them via fallback and DECR again
|
||||
metadata_key = RedisKeys.channel_metadata(str(self.uuid))
|
||||
redis_client.hdel(
|
||||
metadata_key,
|
||||
ChannelMetadataField.STREAM_ID,
|
||||
ChannelMetadataField.M3U_PROFILE,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
def update_stream_profile(self, new_profile_id):
|
||||
"""
|
||||
Updates the profile for the current stream and adjusts connection counts.
|
||||
|
||||
Args:
|
||||
new_profile_id: The ID of the new stream profile to use
|
||||
|
||||
Returns:
|
||||
bool: True if successful, False otherwise
|
||||
"""
|
||||
redis_client = RedisClient.get_client()
|
||||
|
||||
# Get current stream ID
|
||||
stream_id_bytes = redis_client.get(f"channel_stream:{self.id}")
|
||||
if not stream_id_bytes:
|
||||
logger.debug("No active stream found for channel")
|
||||
return False
|
||||
|
||||
stream_id = int(stream_id_bytes)
|
||||
|
||||
# Get current profile ID
|
||||
current_profile_id_bytes = redis_client.get(f"stream_profile:{stream_id}")
|
||||
if not current_profile_id_bytes:
|
||||
logger.debug("No profile found for current stream")
|
||||
return False
|
||||
|
||||
current_profile_id = int(current_profile_id_bytes)
|
||||
|
||||
# Don't do anything if the profile is already set to the requested one
|
||||
if current_profile_id == new_profile_id:
|
||||
return True
|
||||
|
||||
# Use pipeline for atomic profile switch to prevent counter drift
|
||||
# if an exception occurs between DECR and INCR
|
||||
old_profile_connections_key = f"profile_connections:{current_profile_id}"
|
||||
new_profile_connections_key = f"profile_connections:{new_profile_id}"
|
||||
old_count = int(redis_client.get(old_profile_connections_key) or 0)
|
||||
|
||||
pipe = redis_client.pipeline()
|
||||
if old_count > 0:
|
||||
pipe.decr(old_profile_connections_key)
|
||||
pipe.set(f"stream_profile:{stream_id}", new_profile_id)
|
||||
pipe.incr(new_profile_connections_key)
|
||||
pipe.execute()
|
||||
logger.info(
|
||||
f"Updated stream {stream_id} profile from {current_profile_id} to {new_profile_id}"
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
class ChannelProfile(models.Model):
|
||||
name = models.CharField(max_length=100, unique=True)
|
||||
|
||||
|
||||
class ChannelProfileMembership(models.Model):
|
||||
channel_profile = models.ForeignKey(ChannelProfile, on_delete=models.CASCADE)
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
enabled = models.BooleanField(
|
||||
default=True
|
||||
) # Track if the channel is enabled for this group
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_profile", "channel")
|
||||
|
||||
|
||||
class ChannelStream(models.Model):
|
||||
channel = models.ForeignKey(Channel, on_delete=models.CASCADE)
|
||||
stream = models.ForeignKey(Stream, on_delete=models.CASCADE)
|
||||
order = models.PositiveIntegerField(default=0) # Ordering field
|
||||
|
||||
class Meta:
|
||||
ordering = ["order"] # Ensure streams are retrieved in order
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=["channel", "stream"], name="unique_channel_stream"
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class ChannelGroupM3UAccount(models.Model):
|
||||
channel_group = models.ForeignKey(
|
||||
ChannelGroup, on_delete=models.CASCADE, related_name="m3u_accounts"
|
||||
)
|
||||
m3u_account = models.ForeignKey(
|
||||
M3UAccount, on_delete=models.CASCADE, related_name="channel_group"
|
||||
)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
enabled = models.BooleanField(default=True)
|
||||
auto_channel_sync = models.BooleanField(
|
||||
default=False,
|
||||
help_text='Automatically create/delete channels to match streams in this group'
|
||||
)
|
||||
auto_sync_channel_start = models.FloatField(
|
||||
null=True,
|
||||
blank=True,
|
||||
help_text='Starting channel number for auto-created channels in this group'
|
||||
)
|
||||
last_seen = models.DateTimeField(
|
||||
default=timezone.now,
|
||||
db_index=True,
|
||||
help_text='Last time this group was seen in the M3U source during a refresh'
|
||||
)
|
||||
is_stale = models.BooleanField(
|
||||
default=False,
|
||||
db_index=True,
|
||||
help_text='Whether this group relationship is stale (not seen in recent refresh, pending deletion)'
|
||||
)
|
||||
|
||||
class Meta:
|
||||
unique_together = ("channel_group", "m3u_account")
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel_group.name} - {self.m3u_account.name} (Enabled: {self.enabled})"
|
||||
|
||||
|
||||
class Logo(models.Model):
|
||||
name = models.CharField(max_length=255)
|
||||
url = models.TextField(unique=True)
|
||||
|
||||
def __str__(self):
|
||||
return self.name
|
||||
|
||||
|
||||
class Recording(models.Model):
|
||||
channel = models.ForeignKey(
|
||||
"Channel", on_delete=models.CASCADE, related_name="recordings"
|
||||
)
|
||||
start_time = models.DateTimeField()
|
||||
end_time = models.DateTimeField()
|
||||
task_id = models.CharField(max_length=255, null=True, blank=True)
|
||||
custom_properties = models.JSONField(default=dict, blank=True, null=True)
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.channel.name} - {self.start_time} to {self.end_time}"
|
||||
|
||||
|
||||
class RecurringRecordingRule(models.Model):
|
||||
"""Rule describing a recurring manual DVR schedule."""
|
||||
|
||||
channel = models.ForeignKey(
|
||||
"Channel",
|
||||
on_delete=models.CASCADE,
|
||||
related_name="recurring_rules",
|
||||
)
|
||||
days_of_week = models.JSONField(default=list)
|
||||
start_time = models.TimeField()
|
||||
end_time = models.TimeField()
|
||||
enabled = models.BooleanField(default=True)
|
||||
name = models.CharField(max_length=255, blank=True)
|
||||
start_date = models.DateField(null=True, blank=True)
|
||||
end_date = models.DateField(null=True, blank=True)
|
||||
created_at = models.DateTimeField(auto_now_add=True)
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
|
||||
class Meta:
|
||||
ordering = ["channel", "start_time"]
|
||||
|
||||
def __str__(self):
|
||||
channel_name = getattr(self.channel, "name", str(self.channel_id))
|
||||
return f"Recurring rule for {channel_name}"
|
||||
|
||||
def cleaned_days(self):
|
||||
try:
|
||||
return sorted({int(d) for d in (self.days_of_week or []) if 0 <= int(d) <= 6})
|
||||
except Exception:
|
||||
return []
|
||||
@@ -0,0 +1,537 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from rest_framework import serializers
|
||||
from .models import (
|
||||
Stream,
|
||||
Channel,
|
||||
ChannelGroup,
|
||||
ChannelStream,
|
||||
ChannelGroupM3UAccount,
|
||||
Logo,
|
||||
ChannelProfile,
|
||||
ChannelProfileMembership,
|
||||
Recording,
|
||||
RecurringRecordingRule,
|
||||
)
|
||||
from apps.epg.serializers import EPGDataSerializer
|
||||
from core.models import StreamProfile
|
||||
from apps.epg.models import EPGData
|
||||
from django.urls import reverse
|
||||
from rest_framework import serializers
|
||||
from django.utils import timezone
|
||||
from core.utils import validate_flexible_url
|
||||
|
||||
|
||||
class LogoSerializer(serializers.ModelSerializer):
|
||||
cache_url = serializers.SerializerMethodField()
|
||||
channel_count = serializers.SerializerMethodField()
|
||||
is_used = serializers.SerializerMethodField()
|
||||
channel_names = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Logo
|
||||
fields = ["id", "name", "url", "cache_url", "channel_count", "is_used", "channel_names"]
|
||||
|
||||
def validate_url(self, value):
|
||||
"""Validate that the URL is unique for creation or update"""
|
||||
if self.instance and self.instance.url == value:
|
||||
return value
|
||||
|
||||
if Logo.objects.filter(url=value).exists():
|
||||
raise serializers.ValidationError("A logo with this URL already exists.")
|
||||
|
||||
return value
|
||||
|
||||
def create(self, validated_data):
|
||||
"""Handle logo creation with proper URL validation"""
|
||||
return Logo.objects.create(**validated_data)
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
"""Handle logo updates"""
|
||||
for attr, value in validated_data.items():
|
||||
setattr(instance, attr, value)
|
||||
instance.save()
|
||||
return instance
|
||||
|
||||
def get_cache_url(self, obj):
|
||||
# return f"/api/channels/logos/{obj.id}/cache/"
|
||||
request = self.context.get("request")
|
||||
if request:
|
||||
return request.build_absolute_uri(
|
||||
reverse("api:channels:logo-cache", args=[obj.id])
|
||||
)
|
||||
return reverse("api:channels:logo-cache", args=[obj.id])
|
||||
|
||||
def get_channel_count(self, obj):
|
||||
"""Get the number of channels using this logo"""
|
||||
return obj.channels.count()
|
||||
|
||||
def get_is_used(self, obj):
|
||||
"""Check if this logo is used by any channels"""
|
||||
return obj.channels.exists()
|
||||
|
||||
def get_channel_names(self, obj):
|
||||
"""Get the names of channels using this logo (limited to first 5)"""
|
||||
names = []
|
||||
|
||||
# Get channel names
|
||||
channels = obj.channels.all()[:5]
|
||||
for channel in channels:
|
||||
names.append(f"Channel: {channel.name}")
|
||||
|
||||
# Calculate total count for "more" message
|
||||
total_count = self.get_channel_count(obj)
|
||||
if total_count > 5:
|
||||
names.append(f"...and {total_count - 5} more")
|
||||
|
||||
return names
|
||||
|
||||
|
||||
#
|
||||
# Stream
|
||||
#
|
||||
class StreamSerializer(serializers.ModelSerializer):
|
||||
url = serializers.CharField(
|
||||
required=False,
|
||||
allow_blank=True,
|
||||
allow_null=True,
|
||||
validators=[validate_flexible_url]
|
||||
)
|
||||
stream_profile_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=StreamProfile.objects.all(),
|
||||
source="stream_profile",
|
||||
allow_null=True,
|
||||
required=False,
|
||||
)
|
||||
read_only_fields = ["is_custom", "m3u_account", "stream_hash", "stream_id", "stream_chno"]
|
||||
|
||||
class Meta:
|
||||
model = Stream
|
||||
fields = [
|
||||
"id",
|
||||
"name",
|
||||
"url",
|
||||
"m3u_account", # Uncomment if using M3U fields
|
||||
"logo_url",
|
||||
"tvg_id",
|
||||
"local_file",
|
||||
"current_viewers",
|
||||
"updated_at",
|
||||
"last_seen",
|
||||
"is_stale",
|
||||
"is_adult",
|
||||
"stream_profile_id",
|
||||
"is_custom",
|
||||
"channel_group",
|
||||
"stream_hash",
|
||||
"stream_stats",
|
||||
"stream_stats_updated_at",
|
||||
"stream_id",
|
||||
"stream_chno",
|
||||
]
|
||||
|
||||
def get_fields(self):
|
||||
fields = super().get_fields()
|
||||
|
||||
# Unable to edit specific properties if this stream was created from an M3U account
|
||||
if (
|
||||
self.instance
|
||||
and getattr(self.instance, "m3u_account", None)
|
||||
and not self.instance.is_custom
|
||||
):
|
||||
fields["id"].read_only = True
|
||||
fields["name"].read_only = True
|
||||
fields["url"].read_only = True
|
||||
fields["m3u_account"].read_only = True
|
||||
fields["tvg_id"].read_only = True
|
||||
fields["channel_group"].read_only = True
|
||||
|
||||
return fields
|
||||
|
||||
|
||||
class ChannelGroupM3UAccountSerializer(serializers.ModelSerializer):
|
||||
m3u_accounts = serializers.IntegerField(source="m3u_accounts.id", read_only=True)
|
||||
enabled = serializers.BooleanField()
|
||||
auto_channel_sync = serializers.BooleanField(default=False)
|
||||
auto_sync_channel_start = serializers.FloatField(allow_null=True, required=False)
|
||||
custom_properties = serializers.JSONField(required=False)
|
||||
|
||||
class Meta:
|
||||
model = ChannelGroupM3UAccount
|
||||
fields = ["m3u_accounts", "channel_group", "enabled", "auto_channel_sync", "auto_sync_channel_start", "custom_properties", "is_stale", "last_seen"]
|
||||
|
||||
def to_representation(self, instance):
|
||||
data = super().to_representation(instance)
|
||||
|
||||
custom_props = instance.custom_properties or {}
|
||||
|
||||
return data
|
||||
|
||||
def to_internal_value(self, data):
|
||||
# Accept both dict and JSON string for custom_properties (for backward compatibility)
|
||||
val = data.get("custom_properties")
|
||||
if isinstance(val, str):
|
||||
try:
|
||||
data["custom_properties"] = json.loads(val)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return super().to_internal_value(data)
|
||||
|
||||
#
|
||||
# Channel Group
|
||||
#
|
||||
class ChannelGroupSerializer(serializers.ModelSerializer):
|
||||
channel_count = serializers.SerializerMethodField()
|
||||
m3u_account_count = serializers.SerializerMethodField()
|
||||
m3u_accounts = ChannelGroupM3UAccountSerializer(
|
||||
many=True,
|
||||
read_only=True
|
||||
)
|
||||
|
||||
class Meta:
|
||||
model = ChannelGroup
|
||||
fields = ["id", "name", "channel_count", "m3u_account_count", "m3u_accounts"]
|
||||
|
||||
def get_channel_count(self, obj):
|
||||
"""Get count of channels in this group"""
|
||||
return obj.channels.count()
|
||||
|
||||
def get_m3u_account_count(self, obj):
|
||||
"""Get count of M3U accounts associated with this group"""
|
||||
return obj.m3u_accounts.count()
|
||||
|
||||
|
||||
class ChannelProfileSerializer(serializers.ModelSerializer):
|
||||
channels = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = ChannelProfile
|
||||
fields = ["id", "name", "channels"]
|
||||
|
||||
def get_channels(self, obj):
|
||||
memberships = ChannelProfileMembership.objects.filter(
|
||||
channel_profile=obj, enabled=True
|
||||
)
|
||||
return [membership.channel.id for membership in memberships]
|
||||
|
||||
|
||||
class ChannelProfileMembershipSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = ChannelProfileMembership
|
||||
fields = ["channel", "enabled"]
|
||||
|
||||
|
||||
class ChanneProfilelMembershipUpdateSerializer(serializers.Serializer):
|
||||
channel_id = serializers.IntegerField() # Ensure channel_id is an integer
|
||||
enabled = serializers.BooleanField()
|
||||
|
||||
|
||||
class BulkChannelProfileMembershipSerializer(serializers.Serializer):
|
||||
channels = serializers.ListField(
|
||||
child=ChanneProfilelMembershipUpdateSerializer(), # Use the nested serializer
|
||||
allow_empty=False,
|
||||
)
|
||||
|
||||
def validate_channels(self, value):
|
||||
if not value:
|
||||
raise serializers.ValidationError("At least one channel must be provided.")
|
||||
return value
|
||||
|
||||
|
||||
#
|
||||
# Channel
|
||||
#
|
||||
class ChannelSerializer(serializers.ModelSerializer):
|
||||
# Show nested group data, or ID
|
||||
# Ensure channel_number is explicitly typed as FloatField and properly validated
|
||||
channel_number = serializers.FloatField(
|
||||
allow_null=True,
|
||||
required=False,
|
||||
error_messages={"invalid": "Channel number must be a valid decimal number."},
|
||||
)
|
||||
channel_group_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=ChannelGroup.objects.all(), source="channel_group", required=False
|
||||
)
|
||||
epg_data_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=EPGData.objects.all(),
|
||||
source="epg_data",
|
||||
required=False,
|
||||
allow_null=True,
|
||||
)
|
||||
|
||||
stream_profile_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=StreamProfile.objects.all(),
|
||||
source="stream_profile",
|
||||
allow_null=True,
|
||||
required=False,
|
||||
)
|
||||
|
||||
streams = serializers.PrimaryKeyRelatedField(
|
||||
queryset=Stream.objects.all(), many=True, required=False
|
||||
)
|
||||
|
||||
logo_id = serializers.PrimaryKeyRelatedField(
|
||||
queryset=Logo.objects.all(),
|
||||
source="logo",
|
||||
allow_null=True,
|
||||
required=False,
|
||||
)
|
||||
|
||||
auto_created_by_name = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Channel
|
||||
fields = [
|
||||
"id",
|
||||
"channel_number",
|
||||
"name",
|
||||
"channel_group_id",
|
||||
"tvg_id",
|
||||
"tvc_guide_stationid",
|
||||
"epg_data_id",
|
||||
"streams",
|
||||
"stream_profile_id",
|
||||
"uuid",
|
||||
"logo_id",
|
||||
"user_level",
|
||||
"is_adult",
|
||||
"auto_created",
|
||||
"auto_created_by",
|
||||
"auto_created_by_name",
|
||||
]
|
||||
|
||||
def to_representation(self, instance):
|
||||
include_streams = self.context.get("include_streams", False)
|
||||
|
||||
if include_streams:
|
||||
self.fields["streams"] = serializers.SerializerMethodField()
|
||||
return super().to_representation(instance)
|
||||
else:
|
||||
# Fix: For PATCH/PUT responses, ensure streams are ordered
|
||||
representation = super().to_representation(instance)
|
||||
if "streams" in representation:
|
||||
representation["streams"] = list(
|
||||
instance.streams.all()
|
||||
.order_by("channelstream__order")
|
||||
.values_list("id", flat=True)
|
||||
)
|
||||
return representation
|
||||
|
||||
def get_logo(self, obj):
|
||||
return LogoSerializer(obj.logo).data
|
||||
|
||||
def get_streams(self, obj):
|
||||
"""Retrieve ordered stream IDs for GET requests."""
|
||||
return StreamSerializer(
|
||||
obj.streams.all().order_by("channelstream__order"), many=True
|
||||
).data
|
||||
|
||||
def create(self, validated_data):
|
||||
streams = validated_data.pop("streams", [])
|
||||
channel_number = validated_data.pop(
|
||||
"channel_number", Channel.get_next_available_channel_number()
|
||||
)
|
||||
validated_data["channel_number"] = channel_number
|
||||
|
||||
# Auto-assign Default Group if no channel_group is specified
|
||||
if "channel_group" not in validated_data or validated_data.get("channel_group") is None:
|
||||
from apps.channels.models import ChannelGroup
|
||||
default_group, _ = ChannelGroup.objects.get_or_create(name="Default Group")
|
||||
validated_data["channel_group"] = default_group
|
||||
|
||||
channel = Channel.objects.create(**validated_data)
|
||||
|
||||
# Add streams in the specified order
|
||||
for index, stream in enumerate(streams):
|
||||
ChannelStream.objects.create(
|
||||
channel=channel, stream_id=stream.id, order=index
|
||||
)
|
||||
|
||||
return channel
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
streams = validated_data.pop("streams", None)
|
||||
|
||||
# Update standard fields
|
||||
for attr, value in validated_data.items():
|
||||
setattr(instance, attr, value)
|
||||
|
||||
instance.save()
|
||||
|
||||
if streams is not None:
|
||||
# Normalize stream IDs
|
||||
normalized_ids = [
|
||||
stream.id if hasattr(stream, "id") else stream for stream in streams
|
||||
]
|
||||
print(normalized_ids)
|
||||
|
||||
# Get current mapping of stream_id -> ChannelStream
|
||||
current_links = {
|
||||
cs.stream_id: cs for cs in instance.channelstream_set.all()
|
||||
}
|
||||
|
||||
# Track existing stream IDs
|
||||
existing_ids = set(current_links.keys())
|
||||
new_ids = set(normalized_ids)
|
||||
|
||||
# Delete any links not in the new list
|
||||
to_remove = existing_ids - new_ids
|
||||
if to_remove:
|
||||
instance.channelstream_set.filter(stream_id__in=to_remove).delete()
|
||||
|
||||
# Update or create with new order
|
||||
for order, stream_id in enumerate(normalized_ids):
|
||||
if stream_id in current_links:
|
||||
cs = current_links[stream_id]
|
||||
if cs.order != order:
|
||||
cs.order = order
|
||||
cs.save(update_fields=["order"])
|
||||
else:
|
||||
ChannelStream.objects.create(
|
||||
channel=instance, stream_id=stream_id, order=order
|
||||
)
|
||||
|
||||
return instance
|
||||
|
||||
def validate_channel_number(self, value):
|
||||
"""Ensure channel_number is properly processed as a float"""
|
||||
if value is None:
|
||||
return value
|
||||
|
||||
try:
|
||||
# Ensure it's processed as a float
|
||||
return float(value)
|
||||
except (ValueError, TypeError):
|
||||
raise serializers.ValidationError(
|
||||
"Channel number must be a valid decimal number."
|
||||
)
|
||||
|
||||
def validate_stream_profile(self, value):
|
||||
"""Handle special case where empty/0 values mean 'use default' (null)"""
|
||||
if value == "0" or value == 0 or value == "" or value is None:
|
||||
return None
|
||||
return value # PrimaryKeyRelatedField will handle the conversion to object
|
||||
|
||||
def get_auto_created_by_name(self, obj):
|
||||
"""Get the name of the M3U account that auto-created this channel."""
|
||||
if obj.auto_created_by:
|
||||
return obj.auto_created_by.name
|
||||
return None
|
||||
|
||||
|
||||
class RecordingSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = Recording
|
||||
fields = "__all__"
|
||||
read_only_fields = ["task_id"]
|
||||
|
||||
def validate(self, data):
|
||||
from core.models import CoreSettings
|
||||
start_time = data.get("start_time")
|
||||
end_time = data.get("end_time")
|
||||
|
||||
if start_time and timezone.is_naive(start_time):
|
||||
start_time = timezone.make_aware(start_time, timezone.get_current_timezone())
|
||||
data["start_time"] = start_time
|
||||
if end_time and timezone.is_naive(end_time):
|
||||
end_time = timezone.make_aware(end_time, timezone.get_current_timezone())
|
||||
data["end_time"] = end_time
|
||||
|
||||
# If this is an EPG-based recording (program provided), apply global pre/post offsets
|
||||
try:
|
||||
cp = data.get("custom_properties") or {}
|
||||
is_epg_based = isinstance(cp, dict) and isinstance(cp.get("program"), (dict,))
|
||||
except Exception:
|
||||
is_epg_based = False
|
||||
|
||||
if is_epg_based and start_time and end_time:
|
||||
try:
|
||||
pre_min = int(CoreSettings.get_dvr_pre_offset_minutes())
|
||||
except Exception:
|
||||
pre_min = 0
|
||||
try:
|
||||
post_min = int(CoreSettings.get_dvr_post_offset_minutes())
|
||||
except Exception:
|
||||
post_min = 0
|
||||
from datetime import timedelta
|
||||
try:
|
||||
if pre_min and pre_min > 0:
|
||||
start_time = start_time - timedelta(minutes=pre_min)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if post_min and post_min > 0:
|
||||
end_time = end_time + timedelta(minutes=post_min)
|
||||
except Exception:
|
||||
pass
|
||||
# write back adjusted times so scheduling uses them
|
||||
data["start_time"] = start_time
|
||||
data["end_time"] = end_time
|
||||
|
||||
now = timezone.now() # timezone-aware current time
|
||||
|
||||
if end_time < now:
|
||||
raise serializers.ValidationError("End time must be in the future.")
|
||||
|
||||
if start_time < now:
|
||||
# Optional: Adjust start_time if it's in the past but end_time is in the future
|
||||
data["start_time"] = now # or: timezone.now() + timedelta(seconds=1)
|
||||
if end_time <= data["start_time"]:
|
||||
raise serializers.ValidationError("End time must be after start time.")
|
||||
|
||||
return data
|
||||
|
||||
|
||||
class RecurringRecordingRuleSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = RecurringRecordingRule
|
||||
fields = "__all__"
|
||||
read_only_fields = ["created_at", "updated_at"]
|
||||
|
||||
def validate_days_of_week(self, value):
|
||||
if not value:
|
||||
raise serializers.ValidationError("Select at least one day of the week")
|
||||
cleaned = []
|
||||
for entry in value:
|
||||
try:
|
||||
iv = int(entry)
|
||||
except (TypeError, ValueError):
|
||||
raise serializers.ValidationError("Days of week must be integers 0-6")
|
||||
if iv < 0 or iv > 6:
|
||||
raise serializers.ValidationError("Days of week must be between 0 (Monday) and 6 (Sunday)")
|
||||
cleaned.append(iv)
|
||||
return sorted(set(cleaned))
|
||||
|
||||
def validate(self, attrs):
|
||||
start = attrs.get("start_time") or getattr(self.instance, "start_time", None)
|
||||
end = attrs.get("end_time") or getattr(self.instance, "end_time", None)
|
||||
start_date = attrs.get("start_date") if "start_date" in attrs else getattr(self.instance, "start_date", None)
|
||||
end_date = attrs.get("end_date") if "end_date" in attrs else getattr(self.instance, "end_date", None)
|
||||
if start_date is None:
|
||||
existing_start = getattr(self.instance, "start_date", None)
|
||||
if existing_start is None:
|
||||
raise serializers.ValidationError("Start date is required")
|
||||
if start_date and end_date and end_date < start_date:
|
||||
raise serializers.ValidationError("End date must be on or after start date")
|
||||
if end_date is None:
|
||||
existing_end = getattr(self.instance, "end_date", None)
|
||||
if existing_end is None:
|
||||
raise serializers.ValidationError("End date is required")
|
||||
if start and end and start_date and end_date:
|
||||
start_dt = datetime.combine(start_date, start)
|
||||
end_dt = datetime.combine(end_date, end)
|
||||
if end_dt <= start_dt:
|
||||
raise serializers.ValidationError("End datetime must be after start datetime")
|
||||
elif start and end and end == start:
|
||||
raise serializers.ValidationError("End time must be different from start time")
|
||||
# Normalize empty strings to None for dates
|
||||
if attrs.get("end_date") == "":
|
||||
attrs["end_date"] = None
|
||||
if attrs.get("start_date") == "":
|
||||
attrs["start_date"] = None
|
||||
return super().validate(attrs)
|
||||
|
||||
def create(self, validated_data):
|
||||
return super().create(validated_data)
|
||||
@@ -0,0 +1,236 @@
|
||||
# apps/channels/signals.py
|
||||
|
||||
from django.db.models.signals import m2m_changed, pre_save, post_save, post_delete
|
||||
from django.dispatch import receiver
|
||||
from django.utils.timezone import now, is_aware, make_aware
|
||||
from celery.result import AsyncResult
|
||||
from django_celery_beat.models import ClockedSchedule, PeriodicTask
|
||||
from .models import Channel, Stream, ChannelProfile, ChannelProfileMembership, Recording
|
||||
from apps.m3u.models import M3UAccount
|
||||
from apps.epg.tasks import parse_programs_for_tvg_id
|
||||
import json
|
||||
import logging
|
||||
from .tasks import run_recording, prefetch_recording_artwork
|
||||
from datetime import timedelta
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@receiver(m2m_changed, sender=Channel.streams.through)
|
||||
def update_channel_tvg_id_and_logo(sender, instance, action, reverse, model, pk_set, **kwargs):
|
||||
"""
|
||||
Whenever streams are added to a channel:
|
||||
1) If the channel doesn't have a tvg_id, fill it from the first newly-added stream that has one.
|
||||
"""
|
||||
# We only care about post_add, i.e. once the new streams are fully associated
|
||||
if action == "post_add":
|
||||
# --- 1) Populate channel.tvg_id if empty ---
|
||||
if not instance.tvg_id:
|
||||
# Look for newly added streams that have a nonempty tvg_id
|
||||
streams_with_tvg = model.objects.filter(pk__in=pk_set).exclude(tvg_id__exact='')
|
||||
if streams_with_tvg.exists():
|
||||
instance.tvg_id = streams_with_tvg.first().tvg_id
|
||||
instance.save(update_fields=['tvg_id'])
|
||||
|
||||
@receiver(pre_save, sender=Stream)
|
||||
def set_default_m3u_account(sender, instance, **kwargs):
|
||||
"""
|
||||
This function will be triggered before saving a Stream instance.
|
||||
It sets the default m3u_account if not provided.
|
||||
"""
|
||||
if not instance.m3u_account:
|
||||
instance.is_custom = True
|
||||
default_account = M3UAccount.get_custom_account()
|
||||
|
||||
if default_account:
|
||||
instance.m3u_account = default_account
|
||||
else:
|
||||
raise ValueError("No default M3UAccount found.")
|
||||
|
||||
@receiver(post_save, sender=Stream)
|
||||
def generate_custom_stream_hash(sender, instance, created, **kwargs):
|
||||
"""
|
||||
Generate a stable stream_hash for custom streams after creation.
|
||||
Uses the stream's ID to ensure the hash never changes even if name/url is edited.
|
||||
"""
|
||||
if instance.is_custom and not instance.stream_hash and created:
|
||||
import hashlib
|
||||
# Use stream ID for a stable, unique hash that never changes
|
||||
unique_string = f"custom_stream_{instance.id}"
|
||||
instance.stream_hash = hashlib.sha256(unique_string.encode()).hexdigest()
|
||||
# Use update to avoid triggering signals again
|
||||
Stream.objects.filter(id=instance.id).update(stream_hash=instance.stream_hash)
|
||||
|
||||
@receiver(post_save, sender=Channel)
|
||||
def refresh_epg_programs(sender, instance, created, **kwargs):
|
||||
"""
|
||||
When a channel is saved, check if the EPG data has changed.
|
||||
If so, trigger a refresh of the program data for the EPG.
|
||||
"""
|
||||
# Check if this is an update (not a new channel) and the epg_data has changed
|
||||
if not created and kwargs.get('update_fields') and 'epg_data' in kwargs['update_fields']:
|
||||
logger.info(f"Channel {instance.id} ({instance.name}) EPG data updated, refreshing program data")
|
||||
if instance.epg_data:
|
||||
logger.info(f"Triggering EPG program refresh for {instance.epg_data.tvg_id}")
|
||||
parse_programs_for_tvg_id.delay(instance.epg_data.id)
|
||||
# For new channels with EPG data, also refresh
|
||||
elif created and instance.epg_data:
|
||||
logger.info(f"New channel {instance.id} ({instance.name}) created with EPG data, refreshing program data")
|
||||
parse_programs_for_tvg_id.delay(instance.epg_data.id)
|
||||
|
||||
@receiver(post_save, sender=ChannelProfile)
|
||||
def create_profile_memberships(sender, instance, created, **kwargs):
|
||||
if created:
|
||||
channels = Channel.objects.all()
|
||||
ChannelProfileMembership.objects.bulk_create([
|
||||
ChannelProfileMembership(channel_profile=instance, channel=channel)
|
||||
for channel in channels
|
||||
])
|
||||
|
||||
def _dvr_task_name(recording_id):
|
||||
"""Predictable PeriodicTask name for a DVR recording."""
|
||||
return f"dvr-recording-{recording_id}"
|
||||
|
||||
|
||||
def schedule_recording_task(instance, eta=None):
|
||||
"""Schedule a recording task via ClockedSchedule + one-off PeriodicTask.
|
||||
|
||||
The task is stored in the database and dispatched by Celery Beat at the
|
||||
scheduled time with no countdown. This avoids the Redis visibility_timeout
|
||||
redelivery bug that caused duplicate recordings when using apply_async
|
||||
with long countdowns.
|
||||
"""
|
||||
if eta is None:
|
||||
eta = instance.start_time
|
||||
if eta is not None and not is_aware(eta):
|
||||
eta = make_aware(eta)
|
||||
# Clamp to now so Beat dispatches immediately for past/current start times
|
||||
if eta <= now():
|
||||
eta = now()
|
||||
|
||||
task_args = [
|
||||
instance.id,
|
||||
instance.channel_id,
|
||||
str(instance.start_time),
|
||||
str(instance.end_time),
|
||||
]
|
||||
|
||||
clocked, _ = ClockedSchedule.objects.get_or_create(clocked_time=eta)
|
||||
task_name = _dvr_task_name(instance.id)
|
||||
PeriodicTask.objects.update_or_create(
|
||||
name=task_name,
|
||||
defaults={
|
||||
"task": "apps.channels.tasks.run_recording",
|
||||
"clocked": clocked,
|
||||
"args": json.dumps(task_args),
|
||||
"one_off": True,
|
||||
"enabled": True,
|
||||
"interval": None,
|
||||
"crontab": None,
|
||||
"solar": None,
|
||||
},
|
||||
)
|
||||
return task_name
|
||||
|
||||
|
||||
def revoke_task(task_id):
|
||||
"""Cancel a pending recording task.
|
||||
|
||||
task_id is normally a PeriodicTask name (e.g. "dvr-recording-42").
|
||||
For backwards compatibility with legacy Celery async-result UUIDs,
|
||||
falls back to AsyncResult.revoke().
|
||||
"""
|
||||
if not task_id:
|
||||
return
|
||||
# Primary path: delete the PeriodicTask and clean up its ClockedSchedule
|
||||
try:
|
||||
pt = PeriodicTask.objects.get(name=task_id)
|
||||
old_clocked = pt.clocked
|
||||
pt.delete()
|
||||
if old_clocked and not PeriodicTask.objects.filter(clocked=old_clocked).exists():
|
||||
old_clocked.delete()
|
||||
return
|
||||
except PeriodicTask.DoesNotExist:
|
||||
pass
|
||||
# Fallback for legacy Celery task UUIDs
|
||||
try:
|
||||
AsyncResult(task_id).revoke()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@receiver(pre_save, sender=Recording)
|
||||
def revoke_old_task_on_update(sender, instance, **kwargs):
|
||||
if not instance.pk:
|
||||
return # New instance
|
||||
try:
|
||||
old = Recording.objects.get(pk=instance.pk)
|
||||
if old.task_id and (
|
||||
old.start_time != instance.start_time or
|
||||
old.end_time != instance.end_time or
|
||||
old.channel_id != instance.channel_id
|
||||
):
|
||||
# Do NOT revoke while the recording is actively streaming.
|
||||
# run_recording re-reads end_time from the DB every ~2 s and extends
|
||||
# its internal deadline dynamically — revoking here would kill the task.
|
||||
old_status = (old.custom_properties or {}).get("status", "")
|
||||
if old_status == "recording":
|
||||
return
|
||||
revoke_task(old.task_id)
|
||||
instance.task_id = None
|
||||
except Recording.DoesNotExist:
|
||||
pass
|
||||
|
||||
@receiver(post_save, sender=Recording)
|
||||
def schedule_task_on_save(sender, instance, created, **kwargs):
|
||||
try:
|
||||
# Skip processing for internal field-only saves (metadata updates,
|
||||
# task_id assignment, end_time extensions) to prevent re-entrant
|
||||
# artwork dispatch and redundant recording_updated WS events.
|
||||
update_fields = kwargs.get('update_fields')
|
||||
if not created and update_fields is not None and set(update_fields) <= {'custom_properties', 'task_id', 'end_time'}:
|
||||
return
|
||||
|
||||
if not instance.task_id:
|
||||
start_time = instance.start_time
|
||||
end_time = instance.end_time
|
||||
|
||||
# Make datetimes aware (in UTC)
|
||||
if not is_aware(start_time):
|
||||
start_time = make_aware(start_time)
|
||||
if end_time and not is_aware(end_time):
|
||||
end_time = make_aware(end_time)
|
||||
|
||||
current_time = now()
|
||||
|
||||
if start_time > current_time - timedelta(seconds=1):
|
||||
# Future recording — schedule at start_time
|
||||
logger.info(f"Recording {instance.id}: scheduling task at {start_time}")
|
||||
task_id = schedule_recording_task(instance, eta=start_time)
|
||||
instance.task_id = task_id
|
||||
instance.save(update_fields=['task_id'])
|
||||
elif end_time and end_time > current_time:
|
||||
# Currently-playing — start immediately (e.g. series rule for in-progress program)
|
||||
logger.info(f"Recording {instance.id}: start_time in past but end_time still future, scheduling immediately")
|
||||
task_id = schedule_recording_task(instance, eta=current_time)
|
||||
instance.task_id = task_id
|
||||
instance.save(update_fields=['task_id'])
|
||||
else:
|
||||
logger.info(f"Recording {instance.id}: start_time and end_time both in past, not scheduling")
|
||||
# Kick off poster/artwork prefetch to enrich Upcoming cards.
|
||||
# Skip when the recording is already active or finished — run_recording
|
||||
# handles its own poster resolution, and scheduling artwork prefetch
|
||||
# while the task is running causes a race that can overwrite status.
|
||||
cp = instance.custom_properties or {}
|
||||
rec_status = cp.get("status", "")
|
||||
if rec_status not in ("recording", "completed", "stopped", "interrupted"):
|
||||
try:
|
||||
prefetch_recording_artwork.apply_async(args=[instance.id], countdown=1)
|
||||
except Exception as e:
|
||||
print("Error scheduling artwork prefetch:", e)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
print("Error in post_save signal:", e)
|
||||
traceback.print_exc()
|
||||
|
||||
@receiver(post_delete, sender=Recording)
|
||||
def revoke_task_on_delete(sender, instance, **kwargs):
|
||||
revoke_task(instance.task_id)
|
||||
Executable
+3973
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,211 @@
|
||||
from django.test import TestCase
|
||||
from django.contrib.auth import get_user_model
|
||||
from rest_framework.test import APIClient
|
||||
from rest_framework import status
|
||||
|
||||
from apps.channels.models import Channel, ChannelGroup
|
||||
|
||||
User = get_user_model()
|
||||
|
||||
|
||||
class ChannelBulkEditAPITests(TestCase):
|
||||
def setUp(self):
|
||||
# Create a test admin user (user_level >= 10) and authenticate
|
||||
self.user = User.objects.create_user(username="testuser", password="testpass123")
|
||||
self.user.user_level = 10 # Set admin level
|
||||
self.user.save()
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(user=self.user)
|
||||
self.bulk_edit_url = "/api/channels/channels/edit/bulk/"
|
||||
|
||||
# Create test channel group
|
||||
self.group1 = ChannelGroup.objects.create(name="Test Group 1")
|
||||
self.group2 = ChannelGroup.objects.create(name="Test Group 2")
|
||||
|
||||
# Create test channels
|
||||
self.channel1 = Channel.objects.create(
|
||||
channel_number=1.0,
|
||||
name="Channel 1",
|
||||
tvg_id="channel1",
|
||||
channel_group=self.group1
|
||||
)
|
||||
self.channel2 = Channel.objects.create(
|
||||
channel_number=2.0,
|
||||
name="Channel 2",
|
||||
tvg_id="channel2",
|
||||
channel_group=self.group1
|
||||
)
|
||||
self.channel3 = Channel.objects.create(
|
||||
channel_number=3.0,
|
||||
name="Channel 3",
|
||||
tvg_id="channel3"
|
||||
)
|
||||
|
||||
def test_bulk_edit_success(self):
|
||||
"""Test successful bulk update of multiple channels"""
|
||||
data = [
|
||||
{"id": self.channel1.id, "name": "Updated Channel 1"},
|
||||
{"id": self.channel2.id, "name": "Updated Channel 2", "channel_number": 22.0},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
|
||||
self.assertEqual(len(response.data["channels"]), 2)
|
||||
|
||||
# Verify database changes
|
||||
self.channel1.refresh_from_db()
|
||||
self.channel2.refresh_from_db()
|
||||
self.assertEqual(self.channel1.name, "Updated Channel 1")
|
||||
self.assertEqual(self.channel2.name, "Updated Channel 2")
|
||||
self.assertEqual(self.channel2.channel_number, 22.0)
|
||||
|
||||
def test_bulk_edit_with_empty_validated_data_first(self):
|
||||
"""
|
||||
Test the bug fix: when first channel has empty validated_data.
|
||||
This was causing: ValueError: Field names must be given to bulk_update()
|
||||
"""
|
||||
# Create a channel with data that will be "unchanged" (empty validated_data)
|
||||
# We'll send the same data it already has
|
||||
data = [
|
||||
# First channel: no actual changes (this would create empty validated_data)
|
||||
{"id": self.channel1.id},
|
||||
# Second channel: has changes
|
||||
{"id": self.channel2.id, "name": "Updated Channel 2"},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
# Should not crash with ValueError
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
|
||||
|
||||
# Verify the channel with changes was updated
|
||||
self.channel2.refresh_from_db()
|
||||
self.assertEqual(self.channel2.name, "Updated Channel 2")
|
||||
|
||||
def test_bulk_edit_all_empty_updates(self):
|
||||
"""Test when all channels have empty updates (no actual changes)"""
|
||||
data = [
|
||||
{"id": self.channel1.id},
|
||||
{"id": self.channel2.id},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
# Should succeed without calling bulk_update
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["message"], "Successfully updated 2 channels")
|
||||
|
||||
def test_bulk_edit_mixed_fields(self):
|
||||
"""Test bulk update where different channels update different fields"""
|
||||
data = [
|
||||
{"id": self.channel1.id, "name": "New Name 1"},
|
||||
{"id": self.channel2.id, "channel_number": 99.0},
|
||||
{"id": self.channel3.id, "tvg_id": "new_tvg_id", "name": "New Name 3"},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["message"], "Successfully updated 3 channels")
|
||||
|
||||
# Verify all updates
|
||||
self.channel1.refresh_from_db()
|
||||
self.channel2.refresh_from_db()
|
||||
self.channel3.refresh_from_db()
|
||||
|
||||
self.assertEqual(self.channel1.name, "New Name 1")
|
||||
self.assertEqual(self.channel2.channel_number, 99.0)
|
||||
self.assertEqual(self.channel3.tvg_id, "new_tvg_id")
|
||||
self.assertEqual(self.channel3.name, "New Name 3")
|
||||
|
||||
def test_bulk_edit_with_channel_group(self):
|
||||
"""Test bulk update with channel_group_id changes"""
|
||||
data = [
|
||||
{"id": self.channel1.id, "channel_group_id": self.group2.id},
|
||||
{"id": self.channel3.id, "channel_group_id": self.group1.id},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
|
||||
# Verify group changes
|
||||
self.channel1.refresh_from_db()
|
||||
self.channel3.refresh_from_db()
|
||||
self.assertEqual(self.channel1.channel_group, self.group2)
|
||||
self.assertEqual(self.channel3.channel_group, self.group1)
|
||||
|
||||
def test_bulk_edit_nonexistent_channel(self):
|
||||
"""Test bulk update with a channel that doesn't exist"""
|
||||
nonexistent_id = 99999
|
||||
data = [
|
||||
{"id": nonexistent_id, "name": "Should Fail"},
|
||||
{"id": self.channel1.id, "name": "Should Still Update"},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
# Should return 400 with errors
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertIn("errors", response.data)
|
||||
self.assertEqual(len(response.data["errors"]), 1)
|
||||
self.assertEqual(response.data["errors"][0]["channel_id"], nonexistent_id)
|
||||
self.assertEqual(response.data["errors"][0]["error"], "Channel not found")
|
||||
|
||||
# The valid channel should still be updated
|
||||
self.assertEqual(response.data["updated_count"], 1)
|
||||
|
||||
def test_bulk_edit_validation_error(self):
|
||||
"""Test bulk update with invalid data (validation error)"""
|
||||
data = [
|
||||
{"id": self.channel1.id, "channel_number": "invalid_number"},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
# Should return 400 with validation errors
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertIn("errors", response.data)
|
||||
self.assertEqual(len(response.data["errors"]), 1)
|
||||
self.assertIn("channel_number", response.data["errors"][0]["errors"])
|
||||
|
||||
def test_bulk_edit_empty_channel_updates(self):
|
||||
"""Test bulk update with empty list"""
|
||||
data = []
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
# Empty list is accepted and returns success with 0 updates
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertEqual(response.data["message"], "Successfully updated 0 channels")
|
||||
|
||||
def test_bulk_edit_missing_channel_updates(self):
|
||||
"""Test bulk update without proper format (dict instead of list)"""
|
||||
data = {"channel_updates": {}}
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
self.assertEqual(response.data["error"], "Expected a list of channel updates")
|
||||
|
||||
def test_bulk_edit_preserves_other_fields(self):
|
||||
"""Test that bulk update only changes specified fields"""
|
||||
original_channel_number = self.channel1.channel_number
|
||||
original_tvg_id = self.channel1.tvg_id
|
||||
|
||||
data = [
|
||||
{"id": self.channel1.id, "name": "Only Name Changed"},
|
||||
]
|
||||
|
||||
response = self.client.patch(self.bulk_edit_url, data, format="json")
|
||||
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
|
||||
# Verify only name changed, other fields preserved
|
||||
self.channel1.refresh_from_db()
|
||||
self.assertEqual(self.channel1.name, "Only Name Changed")
|
||||
self.assertEqual(self.channel1.channel_number, original_channel_number)
|
||||
self.assertEqual(self.channel1.tvg_id, original_tvg_id)
|
||||
@@ -0,0 +1,278 @@
|
||||
"""Tests for DVR retry logic.
|
||||
|
||||
Covers:
|
||||
- _db_retry(): exponential backoff, max retries, connection reset
|
||||
- Final metadata save retry in run_recording post-processing
|
||||
- Initial TS proxy connection retry (per-base retry on retriable errors)
|
||||
- recover_recordings_on_startup DB retry wrappers
|
||||
"""
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock, patch, call
|
||||
|
||||
from django.db import OperationalError
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.channels.models import Channel, Recording
|
||||
from apps.channels.tasks import _db_retry
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _db_retry unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class DbRetryTests(TestCase):
|
||||
"""Tests for the _db_retry() exponential backoff helper."""
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_succeeds_on_first_attempt(self, _close, _sleep):
|
||||
"""No retry needed when fn succeeds immediately."""
|
||||
result = _db_retry(lambda: "ok", max_retries=3)
|
||||
self.assertEqual(result, "ok")
|
||||
_sleep.assert_not_called()
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_retries_on_operational_error_then_succeeds(self, mock_close, mock_sleep):
|
||||
"""Retry succeeds on second attempt after OperationalError."""
|
||||
call_count = {"n": 0}
|
||||
|
||||
def flaky():
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] == 1:
|
||||
raise OperationalError("connection reset")
|
||||
return "recovered"
|
||||
|
||||
result = _db_retry(flaky, max_retries=3, base_interval=1)
|
||||
self.assertEqual(result, "recovered")
|
||||
self.assertEqual(call_count["n"], 2)
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_raises_after_max_retries_exhausted(self, mock_close, mock_sleep):
|
||||
"""Raises OperationalError after all retries fail."""
|
||||
def always_fail():
|
||||
raise OperationalError("db gone")
|
||||
|
||||
with self.assertRaises(OperationalError):
|
||||
_db_retry(always_fail, max_retries=3, base_interval=1)
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_exponential_backoff_timing(self, mock_close, mock_sleep):
|
||||
"""Sleep durations follow exponential backoff: 1s, 2s, 4s."""
|
||||
call_count = {"n": 0}
|
||||
|
||||
def fail_twice():
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] <= 2:
|
||||
raise OperationalError("retry me")
|
||||
return "done"
|
||||
|
||||
_db_retry(fail_twice, max_retries=3, base_interval=1)
|
||||
mock_sleep.assert_has_calls([call(1), call(2)])
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_close_old_connections_called_between_retries(self, mock_close, mock_sleep):
|
||||
"""Stale DB connections are reset before each retry attempt."""
|
||||
call_count = {"n": 0}
|
||||
|
||||
def fail_once():
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] == 1:
|
||||
raise OperationalError("stale conn")
|
||||
return "ok"
|
||||
|
||||
_db_retry(fail_once, max_retries=3)
|
||||
mock_close.assert_called_once()
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_non_operational_error_not_retried(self, mock_close, mock_sleep):
|
||||
"""Non-OperationalError exceptions propagate immediately."""
|
||||
def raise_value_error():
|
||||
raise ValueError("not a DB error")
|
||||
|
||||
with self.assertRaises(ValueError):
|
||||
_db_retry(raise_value_error, max_retries=3)
|
||||
mock_sleep.assert_not_called()
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_returns_fn_return_value(self, mock_close, mock_sleep):
|
||||
"""Return value of fn() is passed through."""
|
||||
result = _db_retry(lambda: {"key": "value"}, max_retries=3)
|
||||
self.assertEqual(result, {"key": "value"})
|
||||
|
||||
@patch("apps.channels.tasks.time.sleep")
|
||||
@patch("apps.channels.tasks.close_old_connections")
|
||||
def test_single_retry_allowed(self, mock_close, mock_sleep):
|
||||
"""max_retries=1 means no retry — fail immediately."""
|
||||
with self.assertRaises(OperationalError):
|
||||
_db_retry(
|
||||
lambda: (_ for _ in ()).throw(OperationalError("fail")),
|
||||
max_retries=1,
|
||||
)
|
||||
mock_sleep.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Final metadata save retry integration tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class FinalMetadataSaveRetryTests(TestCase):
|
||||
"""The final recording metadata save must retry on transient DB errors."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(
|
||||
channel_number=95, name="Retry Test Channel"
|
||||
)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_metadata_save_uses_db_retry(self, _ws):
|
||||
"""Verify recording metadata is saved via _db_retry (retries on OperationalError)."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "recording"},
|
||||
)
|
||||
# Directly call _db_retry to save metadata as run_recording does
|
||||
cp = rec.custom_properties.copy()
|
||||
cp["status"] = "completed"
|
||||
cp["ended_at"] = str(now)
|
||||
cp["bytes_written"] = 1024
|
||||
|
||||
def _save():
|
||||
rec.custom_properties = cp
|
||||
rec.save(update_fields=["custom_properties"])
|
||||
|
||||
_db_retry(_save, max_retries=3, base_interval=1, label="test save")
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties["status"], "completed")
|
||||
self.assertEqual(rec.custom_properties["bytes_written"], 1024)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_metadata_survives_transient_save_failure(self, _ws):
|
||||
"""Simulate OperationalError on first save, success on retry."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "recording"},
|
||||
)
|
||||
cp = {"status": "completed", "bytes_written": 2048}
|
||||
call_count = {"n": 0}
|
||||
_real_save = rec.save
|
||||
|
||||
def patched_save(**kwargs):
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] == 1:
|
||||
raise OperationalError("connection reset by peer")
|
||||
return _real_save(**kwargs)
|
||||
|
||||
with patch.object(rec, "save", side_effect=patched_save):
|
||||
with patch("apps.channels.tasks.time.sleep"):
|
||||
with patch("apps.channels.tasks.close_old_connections"):
|
||||
def _save():
|
||||
rec.custom_properties = cp
|
||||
rec.save(update_fields=["custom_properties"])
|
||||
_db_retry(_save, max_retries=3, base_interval=1, label="test")
|
||||
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties["status"], "completed")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Initial connection retry tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class InitialConnectionRetryTests(TestCase):
|
||||
"""Verify that the DVR task's reconnection logic retries the same
|
||||
base URL before falling back to the next candidate."""
|
||||
|
||||
def test_reconnect_max_constant_exists_in_run_recording(self):
|
||||
"""run_recording must define a max-reconnect limit to prevent
|
||||
infinite retries on the same broken base URL."""
|
||||
import inspect
|
||||
from apps.channels.tasks import run_recording
|
||||
source = inspect.getsource(run_recording)
|
||||
|
||||
# The reconnection counter pattern must be present
|
||||
self.assertIn("reconnect", source.lower(),
|
||||
"run_recording must contain reconnection logic")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# recover_recordings_on_startup retry tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class RecoveryRetryTests(TestCase):
|
||||
"""DB operations in recover_recordings_on_startup must use _db_retry."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(
|
||||
channel_number=97, name="Recovery Retry Channel"
|
||||
)
|
||||
|
||||
@patch("apps.channels.tasks.run_recording.apply_async")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_recovery_save_retries_on_operational_error(self, _ws, mock_async):
|
||||
"""Recovery status update uses _db_retry — survives one OperationalError."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(minutes=30),
|
||||
end_time=now + timedelta(minutes=30),
|
||||
custom_properties={},
|
||||
)
|
||||
# Simulate what recovery does: mark interrupted, then save with retry
|
||||
cp = rec.custom_properties or {}
|
||||
cp["status"] = "interrupted"
|
||||
cp["interrupted_reason"] = "server_restarted"
|
||||
rec.custom_properties = cp
|
||||
|
||||
call_count = {"n": 0}
|
||||
_real_save = Recording.save
|
||||
|
||||
def patched_save(self_rec, **kwargs):
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] == 1:
|
||||
raise OperationalError("db temporarily unavailable")
|
||||
return _real_save(self_rec, **kwargs)
|
||||
|
||||
with patch.object(Recording, "save", patched_save):
|
||||
with patch("apps.channels.tasks.time.sleep"):
|
||||
with patch("apps.channels.tasks.close_old_connections"):
|
||||
_db_retry(
|
||||
lambda: rec.save(update_fields=["custom_properties"]),
|
||||
max_retries=3,
|
||||
label="test recovery",
|
||||
)
|
||||
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties.get("status"), "interrupted")
|
||||
self.assertEqual(rec.custom_properties.get("interrupted_reason"), "server_restarted")
|
||||
|
||||
def test_db_retry_fetches_recording_list(self):
|
||||
"""_db_retry correctly returns query results for recording list fetch."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(minutes=30),
|
||||
end_time=now + timedelta(minutes=30),
|
||||
custom_properties={},
|
||||
)
|
||||
result = _db_retry(
|
||||
lambda: list(Recording.objects.filter(
|
||||
start_time__lte=now, end_time__gt=now
|
||||
)),
|
||||
label="test query",
|
||||
)
|
||||
self.assertGreaterEqual(len(result), 1)
|
||||
ids = [r.id for r in result]
|
||||
self.assertIn(rec.id, ids)
|
||||
@@ -0,0 +1,59 @@
|
||||
import os
|
||||
from django.test import SimpleTestCase
|
||||
from unittest.mock import patch
|
||||
|
||||
from apps.channels.tasks import build_dvr_candidates
|
||||
|
||||
|
||||
class DVRPortResolutionTests(SimpleTestCase):
|
||||
"""
|
||||
Tests that DVR recording candidate URLs respect the DISPATCHARR_PORT
|
||||
environment variable instead of hardcoding port 9191.
|
||||
"""
|
||||
|
||||
@patch.dict(os.environ, {'REDIS_HOST': 'redis'}, clear=True)
|
||||
def test_default_port_uses_9191(self):
|
||||
"""Without DISPATCHARR_PORT set, candidates default to 9191."""
|
||||
candidates = build_dvr_candidates()
|
||||
self.assertIn('http://web:9191', candidates)
|
||||
self.assertIn('http://localhost:9191', candidates)
|
||||
|
||||
@patch.dict(os.environ, {'DISPATCHARR_PORT': '8080', 'REDIS_HOST': 'redis'}, clear=True)
|
||||
def test_custom_port_reflected_in_candidates(self):
|
||||
"""DISPATCHARR_PORT=8080 replaces all hardcoded 9191 references."""
|
||||
candidates = build_dvr_candidates()
|
||||
self.assertIn('http://web:8080', candidates)
|
||||
self.assertIn('http://localhost:8080', candidates)
|
||||
self.assertNotIn('http://web:9191', candidates)
|
||||
self.assertNotIn('http://localhost:9191', candidates)
|
||||
|
||||
@patch.dict(os.environ, {
|
||||
'DISPATCHARR_PORT': '7777',
|
||||
'DISPATCHARR_ENV': 'dev',
|
||||
'REDIS_HOST': 'redis',
|
||||
}, clear=True)
|
||||
def test_dev_mode_includes_5656_and_custom_port(self):
|
||||
"""Dev mode includes both uwsgi internal port (5656) and custom port."""
|
||||
candidates = build_dvr_candidates()
|
||||
self.assertIn('http://127.0.0.1:5656', candidates)
|
||||
self.assertIn('http://127.0.0.1:7777', candidates)
|
||||
|
||||
@patch.dict(os.environ, {
|
||||
'DISPATCHARR_INTERNAL_TS_BASE_URL': 'http://custom:1234',
|
||||
'REDIS_HOST': 'redis',
|
||||
}, clear=True)
|
||||
def test_explicit_override_is_first(self):
|
||||
"""DISPATCHARR_INTERNAL_TS_BASE_URL should be the first candidate."""
|
||||
candidates = build_dvr_candidates()
|
||||
self.assertEqual(candidates[0], 'http://custom:1234')
|
||||
|
||||
@patch.dict(os.environ, {
|
||||
'DISPATCHARR_PORT': '3000',
|
||||
'DISPATCHARR_INTERNAL_API_BASE': 'http://myhost:4000',
|
||||
'REDIS_HOST': 'redis',
|
||||
}, clear=True)
|
||||
def test_internal_api_base_overrides_web_fallback(self):
|
||||
"""DISPATCHARR_INTERNAL_API_BASE replaces the http://web:{port} default."""
|
||||
candidates = build_dvr_candidates()
|
||||
self.assertIn('http://myhost:4000', candidates)
|
||||
self.assertNotIn('http://web:3000', candidates)
|
||||
@@ -0,0 +1,180 @@
|
||||
"""Tests for the _match_epg_program_by_timeslot() helper in tasks.py.
|
||||
|
||||
Covers:
|
||||
- Exact time-slot match returns program dict
|
||||
- 80% overlap threshold: at boundary, above, and below
|
||||
- Multiple overlapping programs: dominant vs. evenly split
|
||||
- Edge cases: None inputs, zero-duration recording, no EPG data
|
||||
- Returned dict structure (id, title, sub_title, description)
|
||||
"""
|
||||
from datetime import timedelta
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.channels.models import Channel
|
||||
from apps.epg.models import EPGSource, EPGData, ProgramData
|
||||
from apps.channels.tasks import _match_epg_program_by_timeslot
|
||||
|
||||
|
||||
class EpgMatchingSetupMixin:
|
||||
"""Shared setup for EPG matching tests."""
|
||||
|
||||
def setUp(self):
|
||||
self.source = EPGSource.objects.create(name="Test Source")
|
||||
self.epg = EPGData.objects.create(
|
||||
tvg_id="test.channel", name="Test Channel EPG", epg_source=self.source,
|
||||
)
|
||||
self.channel = Channel.objects.create(
|
||||
channel_number=50, name="EPG Match Channel", epg_data=self.epg,
|
||||
)
|
||||
self.base = timezone.now().replace(second=0, microsecond=0)
|
||||
|
||||
def _prog(self, offset_min, duration_min, title="Test Show", **kwargs):
|
||||
"""Create a ProgramData starting offset_min from self.base."""
|
||||
start = self.base + timedelta(minutes=offset_min)
|
||||
end = start + timedelta(minutes=duration_min)
|
||||
return ProgramData.objects.create(
|
||||
epg=self.epg, start_time=start, end_time=end, title=title, **kwargs,
|
||||
)
|
||||
|
||||
|
||||
class ExactMatchTests(EpgMatchingSetupMixin, TestCase):
|
||||
"""Recording window exactly matches an EPG program."""
|
||||
|
||||
def test_exact_match_returns_program_dict(self):
|
||||
prog = self._prog(0, 60, title="News at 9", sub_title="Top Stories",
|
||||
description="Evening news broadcast")
|
||||
result = _match_epg_program_by_timeslot(
|
||||
self.epg, prog.start_time, prog.end_time,
|
||||
)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result["id"], prog.id)
|
||||
self.assertEqual(result["title"], "News at 9")
|
||||
self.assertEqual(result["sub_title"], "Top Stories")
|
||||
self.assertEqual(result["description"], "Evening news broadcast")
|
||||
|
||||
def test_missing_optional_fields_returned_as_empty_strings(self):
|
||||
prog = self._prog(0, 30, title="Minimal Show")
|
||||
result = _match_epg_program_by_timeslot(
|
||||
self.epg, prog.start_time, prog.end_time,
|
||||
)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result["sub_title"], "")
|
||||
self.assertEqual(result["description"], "")
|
||||
|
||||
|
||||
class OverlapThresholdTests(EpgMatchingSetupMixin, TestCase):
|
||||
"""80% overlap threshold boundary tests."""
|
||||
|
||||
def test_exactly_80_percent_overlap_returns_match(self):
|
||||
"""Program covers exactly 80% of the recording window."""
|
||||
# Program: 0-60min, Recording: 0-75min → overlap = 60/75 = 80%
|
||||
prog = self._prog(0, 60, title="Borderline Show")
|
||||
rec_start = self.base
|
||||
rec_end = self.base + timedelta(minutes=75)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result["title"], "Borderline Show")
|
||||
|
||||
def test_below_80_percent_returns_none(self):
|
||||
"""Program covers 79% of the recording — below threshold."""
|
||||
# Program: 0-60min, Recording: 0-76min → overlap = 60/76 ≈ 78.9%
|
||||
prog = self._prog(0, 60, title="Too Short")
|
||||
rec_start = self.base
|
||||
rec_end = self.base + timedelta(minutes=76)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_above_80_percent_returns_match(self):
|
||||
"""Program covers 90% of the recording."""
|
||||
# Program: 0-60min, Recording: 0-66min → overlap = 60/66 ≈ 90.9%
|
||||
prog = self._prog(0, 60, title="Good Match")
|
||||
rec_start = self.base
|
||||
rec_end = self.base + timedelta(minutes=66)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result["title"], "Good Match")
|
||||
|
||||
|
||||
class MultipleProgramTests(EpgMatchingSetupMixin, TestCase):
|
||||
"""Recording spans multiple EPG programs."""
|
||||
|
||||
def test_dominant_program_returned(self):
|
||||
"""Recording spans 2 programs; one covers 85%, the other 15%."""
|
||||
# Show A: 0-60min, Show B: 60-120min
|
||||
# Recording: 9-69min → A overlap=51/60=85%, B overlap=9/60=15%
|
||||
self._prog(0, 60, title="Show A")
|
||||
self._prog(60, 60, title="Show B")
|
||||
rec_start = self.base + timedelta(minutes=9)
|
||||
rec_end = self.base + timedelta(minutes=69)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result["title"], "Show A")
|
||||
|
||||
def test_evenly_split_returns_none(self):
|
||||
"""Recording spans 2 equal programs — neither reaches 80%."""
|
||||
# Show A: 0-60min, Show B: 60-120min
|
||||
# Recording: 30-90min → each covers 50%
|
||||
self._prog(0, 60, title="Show A")
|
||||
self._prog(60, 60, title="Show B")
|
||||
rec_start = self.base + timedelta(minutes=30)
|
||||
rec_end = self.base + timedelta(minutes=90)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_three_programs_one_dominant(self):
|
||||
"""Recording spans 3 programs; middle one is dominant."""
|
||||
# A: 0-30min, B: 30-90min, C: 90-120min
|
||||
# Recording: 25-95min (70min window) → B overlap=60/70≈85.7%
|
||||
self._prog(0, 30, title="Show A")
|
||||
self._prog(30, 60, title="Show B")
|
||||
self._prog(90, 30, title="Show C")
|
||||
rec_start = self.base + timedelta(minutes=25)
|
||||
rec_end = self.base + timedelta(minutes=95)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result["title"], "Show B")
|
||||
|
||||
|
||||
class EdgeCaseTests(EpgMatchingSetupMixin, TestCase):
|
||||
"""Edge cases and error handling."""
|
||||
|
||||
def test_none_epg_data_returns_none(self):
|
||||
result = _match_epg_program_by_timeslot(None, self.base, self.base + timedelta(hours=1))
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_none_start_time_returns_none(self):
|
||||
result = _match_epg_program_by_timeslot(self.epg, None, self.base + timedelta(hours=1))
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_none_end_time_returns_none(self):
|
||||
result = _match_epg_program_by_timeslot(self.epg, self.base, None)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_zero_duration_returns_none(self):
|
||||
"""Recording with start == end should return None."""
|
||||
result = _match_epg_program_by_timeslot(self.epg, self.base, self.base)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_negative_duration_returns_none(self):
|
||||
"""Recording with end before start should return None."""
|
||||
result = _match_epg_program_by_timeslot(
|
||||
self.epg, self.base + timedelta(hours=1), self.base,
|
||||
)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_no_overlapping_programs_returns_none(self):
|
||||
"""No EPG programs in the recording window."""
|
||||
self._prog(0, 60, title="Earlier Show")
|
||||
rec_start = self.base + timedelta(hours=5)
|
||||
rec_end = rec_start + timedelta(hours=1)
|
||||
result = _match_epg_program_by_timeslot(self.epg, rec_start, rec_end)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_empty_epg_no_programs_returns_none(self):
|
||||
"""EPGData exists but has no programs."""
|
||||
result = _match_epg_program_by_timeslot(
|
||||
self.epg, self.base, self.base + timedelta(hours=1),
|
||||
)
|
||||
self.assertIsNone(result)
|
||||
@@ -0,0 +1,235 @@
|
||||
"""Tests for the Extend In-Progress Recording feature.
|
||||
|
||||
Covers:
|
||||
- extend() API endpoint (happy path and validation)
|
||||
- pre_save signal guard: end_time change must NOT revoke a live recording
|
||||
- pre_save signal guard: end_time change MUST still revoke an upcoming recording
|
||||
- TOCTOU edge cases (extend on a completed/stopped/nonexistent recording)
|
||||
"""
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
from rest_framework.test import APIRequestFactory, force_authenticate
|
||||
|
||||
from apps.channels.models import Channel, Recording
|
||||
from apps.channels.api_views import RecordingViewSet
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_admin():
|
||||
from django.contrib.auth import get_user_model
|
||||
User = get_user_model()
|
||||
u, _ = User.objects.get_or_create(
|
||||
username="extend_test_admin",
|
||||
defaults={"user_level": User.UserLevel.ADMIN},
|
||||
)
|
||||
u.set_password("pass")
|
||||
u.save()
|
||||
return u
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Extend endpoint tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ExtendEndpointTests(TestCase):
|
||||
"""Tests for POST /api/channels/recordings/{id}/extend/"""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(
|
||||
channel_number=88, name="Extend Test Channel"
|
||||
)
|
||||
self.user = _make_admin()
|
||||
self.factory = APIRequestFactory()
|
||||
|
||||
def _extend(self, rec, extra_minutes):
|
||||
request = self.factory.post(
|
||||
f"/api/channels/recordings/{rec.id}/extend/",
|
||||
{"extra_minutes": extra_minutes},
|
||||
format="json",
|
||||
)
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "extend"})
|
||||
return view(request, pk=rec.id)
|
||||
|
||||
def _make_rec(self, status="recording"):
|
||||
now = timezone.now()
|
||||
return Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": status},
|
||||
)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_extend_updates_end_time_in_db(self, _ws):
|
||||
rec = self._make_rec()
|
||||
original_end = rec.end_time
|
||||
response = self._extend(rec, 30)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertTrue(response.data.get("success"))
|
||||
rec.refresh_from_db()
|
||||
expected = original_end + timedelta(minutes=30)
|
||||
delta = abs((rec.end_time - expected).total_seconds())
|
||||
self.assertLess(delta, 1, "end_time was not extended by the correct amount")
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_extend_stacks_multiple_extensions(self, _ws):
|
||||
"""Calling extend() twice adds both increments."""
|
||||
rec = self._make_rec()
|
||||
original_end = rec.end_time
|
||||
self._extend(rec, 15)
|
||||
self._extend(rec, 30)
|
||||
rec.refresh_from_db()
|
||||
expected = original_end + timedelta(minutes=45)
|
||||
delta = abs((rec.end_time - expected).total_seconds())
|
||||
self.assertLess(delta, 1)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_extend_does_not_clear_task_id(self, _ws):
|
||||
"""The running Celery task must survive the DB save."""
|
||||
rec = self._make_rec()
|
||||
rec.task_id = "dvr-recording-999"
|
||||
rec.save(update_fields=["task_id"])
|
||||
self._extend(rec, 30)
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.task_id, "dvr-recording-999")
|
||||
|
||||
def test_extend_returns_400_if_finished(self):
|
||||
"""Cannot extend a completed, stopped, or interrupted recording."""
|
||||
for bad_status in ("completed", "stopped", "interrupted"):
|
||||
with self.subTest(status=bad_status):
|
||||
rec = self._make_rec(status=bad_status)
|
||||
response = self._extend(rec, 30)
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertFalse(response.data.get("success"))
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_extend_succeeds_before_task_sets_status(self, _ws):
|
||||
"""Extend must work when status is empty (task hasn't started yet)."""
|
||||
rec = self._make_rec(status="")
|
||||
response = self._extend(rec, 15)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
expected = rec.end_time # already extended
|
||||
self.assertTrue(response.data.get("success"))
|
||||
|
||||
@patch("apps.channels.signals.revoke_task")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_extend_bypasses_signals_no_revoke(self, _ws, mock_revoke):
|
||||
"""Extend uses .update() to bypass pre_save — revoke_task must never fire."""
|
||||
rec = self._make_rec(status="")
|
||||
rec.task_id = "dvr-recording-500"
|
||||
rec.save(update_fields=["task_id"])
|
||||
self._extend(rec, 15)
|
||||
self._extend(rec, 30)
|
||||
mock_revoke.assert_not_called()
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.task_id, "dvr-recording-500")
|
||||
|
||||
def test_extend_returns_400_for_zero_minutes(self):
|
||||
response = self._extend(self._make_rec(), 0)
|
||||
self.assertEqual(response.status_code, 400)
|
||||
|
||||
def test_extend_returns_400_for_negative_minutes(self):
|
||||
response = self._extend(self._make_rec(), -15)
|
||||
self.assertEqual(response.status_code, 400)
|
||||
|
||||
def test_extend_returns_400_for_non_numeric_minutes(self):
|
||||
rec = self._make_rec()
|
||||
request = self.factory.post(
|
||||
f"/api/channels/recordings/{rec.id}/extend/",
|
||||
{"extra_minutes": "lots"},
|
||||
format="json",
|
||||
)
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "extend"})
|
||||
response = view(request, pk=rec.id)
|
||||
self.assertEqual(response.status_code, 400)
|
||||
|
||||
def test_extend_returns_404_for_nonexistent_recording(self):
|
||||
request = self.factory.post(
|
||||
"/api/channels/recordings/999999/extend/",
|
||||
{"extra_minutes": 30},
|
||||
format="json",
|
||||
)
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "extend"})
|
||||
response = view(request, pk=999999)
|
||||
self.assertEqual(response.status_code, 404)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# pre_save signal guard tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class PreSaveExtendGuardTests(TestCase):
|
||||
"""The pre_save signal must NOT revoke a live recording when end_time changes,
|
||||
but MUST still revoke a scheduled (upcoming) recording as before."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(
|
||||
channel_number=77, name="Signal Guard Channel"
|
||||
)
|
||||
|
||||
def _make_rec(self, status="", task_id="dvr-recording-42"):
|
||||
now = timezone.now()
|
||||
return Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now + timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=2),
|
||||
task_id=task_id,
|
||||
custom_properties={"status": status} if status else {},
|
||||
)
|
||||
|
||||
@patch("apps.channels.signals.revoke_task")
|
||||
def test_end_time_change_does_not_revoke_live_recording(self, mock_revoke):
|
||||
"""When status='recording', extending end_time must not call revoke_task."""
|
||||
rec = self._make_rec(status="recording", task_id="dvr-recording-42")
|
||||
rec.end_time = rec.end_time + timedelta(minutes=30)
|
||||
rec.save(update_fields=["end_time"])
|
||||
mock_revoke.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.revoke_task")
|
||||
def test_task_id_preserved_after_extend_on_live_recording(self, mock_revoke):
|
||||
"""task_id must not be cleared for a live recording's end_time change."""
|
||||
rec = self._make_rec(status="recording", task_id="dvr-recording-42")
|
||||
original_task_id = rec.task_id
|
||||
rec.end_time = rec.end_time + timedelta(minutes=30)
|
||||
rec.save(update_fields=["end_time"])
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.task_id, original_task_id)
|
||||
|
||||
@patch("apps.channels.signals.revoke_task")
|
||||
def test_end_time_change_still_revokes_upcoming_recording(self, mock_revoke):
|
||||
"""The guard must NOT apply to upcoming recordings — existing behavior preserved."""
|
||||
rec = self._make_rec(status="", task_id="dvr-recording-77")
|
||||
rec.end_time = rec.end_time + timedelta(minutes=30)
|
||||
rec.save(update_fields=["end_time"])
|
||||
mock_revoke.assert_called_once_with("dvr-recording-77")
|
||||
|
||||
@patch("apps.channels.signals.revoke_task")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_pre_save_guard_reads_db_status_not_memory_status(self, _ws, mock_revoke):
|
||||
"""pre_save reads status from DB (old object), not from the instance being saved."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
task_id="dvr-recording-66",
|
||||
custom_properties={"status": "recording"},
|
||||
)
|
||||
# Simulate: DB status changes to 'completed' behind the instance's back
|
||||
Recording.objects.filter(pk=rec.pk).update(
|
||||
custom_properties={"status": "completed"}
|
||||
)
|
||||
rec.end_time = rec.end_time + timedelta(minutes=30)
|
||||
rec.save(update_fields=["end_time"])
|
||||
# revoke_task should be called because DB status is "completed", not "recording"
|
||||
mock_revoke.assert_called_once_with("dvr-recording-66")
|
||||
@@ -0,0 +1,379 @@
|
||||
"""Tests for recording metadata endpoints and logo proxy negative cache.
|
||||
|
||||
Covers:
|
||||
- update_metadata endpoint: title/description, user_edited flag, validation
|
||||
- refresh_artwork endpoint: returns immediately, background thread behavior
|
||||
- Logo proxy negative cache: cache hit/miss, expiry, eviction, success clears
|
||||
"""
|
||||
import time as time_mod
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
from rest_framework.test import APIRequestFactory, force_authenticate
|
||||
|
||||
from apps.channels.models import Channel, Recording, Logo
|
||||
from apps.channels.api_views import RecordingViewSet, LogoViewSet
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_admin():
|
||||
from django.contrib.auth import get_user_model
|
||||
User = get_user_model()
|
||||
u, _ = User.objects.get_or_create(
|
||||
username="metadata_test_admin",
|
||||
defaults={"user_level": User.UserLevel.ADMIN},
|
||||
)
|
||||
u.set_password("pass")
|
||||
u.save()
|
||||
return u
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# update_metadata endpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class UpdateMetadataTests(TestCase):
|
||||
"""Tests for POST /api/channels/recordings/{id}/update-metadata/"""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=70, name="Meta Test Channel")
|
||||
self.user = _make_admin()
|
||||
self.factory = APIRequestFactory()
|
||||
|
||||
def _update(self, rec, data):
|
||||
request = self.factory.post(
|
||||
f"/api/channels/recordings/{rec.id}/update-metadata/",
|
||||
data, format="json",
|
||||
)
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "update_metadata"})
|
||||
return view(request, pk=rec.id)
|
||||
|
||||
def _make_rec(self, custom_properties=None):
|
||||
now = timezone.now()
|
||||
return Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties=custom_properties or {},
|
||||
)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_update_title_only(self, _ws):
|
||||
rec = self._make_rec()
|
||||
response = self._update(rec, {"title": "My Show"})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
program = rec.custom_properties["program"]
|
||||
self.assertEqual(program["title"], "My Show")
|
||||
self.assertTrue(program["user_edited"])
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_update_description_only(self, _ws):
|
||||
rec = self._make_rec({"program": {"title": "Existing Title"}})
|
||||
response = self._update(rec, {"description": "A great episode"})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
program = rec.custom_properties["program"]
|
||||
self.assertEqual(program["description"], "A great episode")
|
||||
self.assertEqual(program["title"], "Existing Title")
|
||||
self.assertTrue(program["user_edited"])
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_update_both_fields(self, _ws):
|
||||
rec = self._make_rec()
|
||||
response = self._update(rec, {"title": "New Title", "description": "New Desc"})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
program = rec.custom_properties["program"]
|
||||
self.assertEqual(program["title"], "New Title")
|
||||
self.assertEqual(program["description"], "New Desc")
|
||||
self.assertTrue(program["user_edited"])
|
||||
|
||||
def test_no_fields_returns_400(self):
|
||||
rec = self._make_rec()
|
||||
response = self._update(rec, {})
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertFalse(response.data.get("success"))
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_whitespace_trimmed(self, _ws):
|
||||
rec = self._make_rec()
|
||||
response = self._update(rec, {"title": " Padded Title "})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties["program"]["title"], "Padded Title")
|
||||
|
||||
def test_whitespace_only_title_returns_400(self):
|
||||
"""Whitespace-only title and description should be rejected."""
|
||||
rec = self._make_rec()
|
||||
response = self._update(rec, {"title": " ", "description": " "})
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertFalse(response.data.get("success"))
|
||||
|
||||
def test_whitespace_only_title_with_valid_description(self):
|
||||
"""Whitespace-only title is ignored; valid description is accepted."""
|
||||
rec = self._make_rec({"program": {"title": "Original"}})
|
||||
with patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None):
|
||||
response = self._update(rec, {"title": " ", "description": "Valid desc"})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
# Title should remain unchanged since the whitespace-only value is not applied
|
||||
self.assertEqual(rec.custom_properties["program"]["title"], "Original")
|
||||
self.assertEqual(rec.custom_properties["program"]["description"], "Valid desc")
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_creates_program_dict_when_absent(self, _ws):
|
||||
"""Recording with no program dict gets one created."""
|
||||
rec = self._make_rec({"status": "completed"})
|
||||
response = self._update(rec, {"title": "Brand New"})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
self.assertIn("program", rec.custom_properties)
|
||||
self.assertEqual(rec.custom_properties["program"]["title"], "Brand New")
|
||||
|
||||
def test_returns_404_for_nonexistent(self):
|
||||
request = self.factory.post(
|
||||
"/api/channels/recordings/99999/update-metadata/",
|
||||
{"title": "Ghost"}, format="json",
|
||||
)
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "update_metadata"})
|
||||
self.assertEqual(view(request, pk=99999).status_code, 404)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_sends_websocket_event(self, mock_ws):
|
||||
rec = self._make_rec()
|
||||
self._update(rec, {"title": "WS Test"})
|
||||
mock_ws.assert_called_once()
|
||||
payload = mock_ws.call_args[0][2]
|
||||
self.assertEqual(payload["type"], "recording_updated")
|
||||
self.assertEqual(payload["recording_id"], rec.id)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=Exception("WS down"))
|
||||
def test_ws_failure_does_not_fail_request(self, _ws):
|
||||
"""WebSocket errors are silenced — the save still succeeds."""
|
||||
rec = self._make_rec()
|
||||
response = self._update(rec, {"title": "Resilient"})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties["program"]["title"], "Resilient")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# refresh_artwork endpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class RefreshArtworkTests(TestCase):
|
||||
"""Tests for POST /api/channels/recordings/{id}/refresh-artwork/"""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=71, name="Artwork Test Channel")
|
||||
self.user = _make_admin()
|
||||
self.factory = APIRequestFactory()
|
||||
|
||||
def _refresh(self, rec):
|
||||
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
|
||||
return view(request, pk=rec.id)
|
||||
|
||||
def _make_rec(self, custom_properties=None):
|
||||
now = timezone.now()
|
||||
return Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties=custom_properties or {},
|
||||
)
|
||||
|
||||
@patch("threading.Thread")
|
||||
def test_returns_200_immediately(self, mock_thread):
|
||||
mock_thread.return_value.start = MagicMock()
|
||||
rec = self._make_rec()
|
||||
response = self._refresh(rec)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertTrue(response.data.get("success"))
|
||||
|
||||
@patch("threading.Thread")
|
||||
def test_spawns_background_thread(self, mock_thread):
|
||||
mock_thread.return_value.start = MagicMock()
|
||||
rec = self._make_rec()
|
||||
self._refresh(rec)
|
||||
mock_thread.assert_called_once()
|
||||
self.assertTrue(mock_thread.call_args[1].get("daemon", False))
|
||||
mock_thread.return_value.start.assert_called_once()
|
||||
|
||||
def test_returns_404_for_nonexistent(self):
|
||||
request = self.factory.post("/api/channels/recordings/99999/refresh-artwork/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
|
||||
self.assertEqual(view(request, pk=99999).status_code, 404)
|
||||
|
||||
@patch("django.db.close_old_connections")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_no_downgrade_to_channel_logo(self, _ws, _close):
|
||||
"""When the pipeline returns the channel's own logo, existing poster is preserved."""
|
||||
logo = Logo.objects.create(name="Channel Logo", url="https://example.com/ch.png")
|
||||
self.channel.logo = logo
|
||||
self.channel.save()
|
||||
rec = self._make_rec({
|
||||
"poster_logo_id": 999, # existing real poster
|
||||
"poster_url": "https://tmdb.com/real-poster.jpg",
|
||||
})
|
||||
|
||||
with patch("apps.channels.tasks._resolve_poster_for_program",
|
||||
return_value=(logo.id, None)):
|
||||
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
|
||||
force_authenticate(request, user=self.user)
|
||||
|
||||
# Run synchronously by intercepting the thread
|
||||
captured_fn = None
|
||||
def capture_thread(*args, **kwargs):
|
||||
nonlocal captured_fn
|
||||
captured_fn = kwargs.get("target") or args[0]
|
||||
mock = MagicMock()
|
||||
mock.start = lambda: captured_fn(*kwargs.get("args", ()))
|
||||
return mock
|
||||
|
||||
with patch("threading.Thread", side_effect=capture_thread):
|
||||
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
|
||||
view(request, pk=rec.id)
|
||||
|
||||
rec.refresh_from_db()
|
||||
# Existing poster should be preserved — not downgraded to channel logo
|
||||
self.assertEqual(rec.custom_properties.get("poster_logo_id"), 999)
|
||||
self.assertEqual(rec.custom_properties.get("poster_url"), "https://tmdb.com/real-poster.jpg")
|
||||
|
||||
@patch("django.db.close_old_connections")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_upgrade_from_no_poster(self, _ws, _close):
|
||||
"""When a recording has no poster and the pipeline finds one, it gets updated."""
|
||||
rec = self._make_rec({"program": {"title": "Some Show", "id": 42}})
|
||||
|
||||
with patch("apps.channels.tasks._resolve_poster_for_program",
|
||||
return_value=(555, "https://tmdb.com/new-poster.jpg")):
|
||||
captured_fn = None
|
||||
def capture_thread(*args, **kwargs):
|
||||
nonlocal captured_fn
|
||||
captured_fn = kwargs.get("target") or args[0]
|
||||
mock = MagicMock()
|
||||
mock.start = lambda: captured_fn(*kwargs.get("args", ()))
|
||||
return mock
|
||||
|
||||
with patch("threading.Thread", side_effect=capture_thread):
|
||||
request = self.factory.post(f"/api/channels/recordings/{rec.id}/refresh-artwork/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "refresh_artwork"})
|
||||
view(request, pk=rec.id)
|
||||
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties.get("poster_logo_id"), 555)
|
||||
self.assertEqual(rec.custom_properties.get("poster_url"), "https://tmdb.com/new-poster.jpg")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Logo proxy negative cache
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class LogoNegativeCacheTests(TestCase):
|
||||
"""Tests for the _logo_fetch_failures negative cache in LogoViewSet.cache()."""
|
||||
|
||||
def setUp(self):
|
||||
from apps.channels import api_views
|
||||
self._failures = api_views._logo_fetch_failures
|
||||
self._failures.clear()
|
||||
self.factory = APIRequestFactory()
|
||||
self.user = _make_admin()
|
||||
|
||||
def _fetch_logo(self, logo):
|
||||
request = self.factory.get(f"/api/channels/logos/{logo.id}/cache/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = LogoViewSet.as_view({"get": "cache"})
|
||||
return view(request, pk=logo.id)
|
||||
|
||||
def test_failed_url_cached_on_non_200(self):
|
||||
"""Non-200 response adds URL to negative cache."""
|
||||
logo = Logo.objects.create(name="Dead Logo", url="https://dead-cdn.com/logo.png")
|
||||
mock_resp = MagicMock(status_code=404)
|
||||
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
|
||||
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
|
||||
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
|
||||
response = self._fetch_logo(logo)
|
||||
self.assertEqual(response.status_code, 404)
|
||||
self.assertIn("https://dead-cdn.com/logo.png", self._failures)
|
||||
|
||||
def test_cached_failure_returns_404_immediately(self):
|
||||
"""Subsequent request for a cached-failed URL returns 404 without making a request."""
|
||||
logo = Logo.objects.create(name="Cached Fail", url="https://cached-fail.com/logo.png")
|
||||
self._failures["https://cached-fail.com/logo.png"] = time_mod.monotonic() + 300
|
||||
|
||||
with patch("apps.channels.api_views.requests.get") as mock_get:
|
||||
response = self._fetch_logo(logo)
|
||||
self.assertEqual(response.status_code, 404)
|
||||
mock_get.assert_not_called()
|
||||
|
||||
def test_expired_cache_entry_allows_retry(self):
|
||||
"""After TTL expires, a new request is made."""
|
||||
logo = Logo.objects.create(name="Expired", url="https://expired.com/logo.png")
|
||||
self._failures["https://expired.com/logo.png"] = time_mod.monotonic() - 1 # already expired
|
||||
|
||||
mock_resp = MagicMock(status_code=200)
|
||||
mock_resp.headers = {"Content-Type": "image/png"}
|
||||
mock_resp.iter_content = MagicMock(return_value=[b"img"])
|
||||
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
|
||||
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
|
||||
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
|
||||
response = self._fetch_logo(logo)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
def test_success_clears_previous_failure(self):
|
||||
"""A successful fetch removes the URL from the failure cache."""
|
||||
url = "https://recovered.com/logo.png"
|
||||
logo = Logo.objects.create(name="Recovered", url=url)
|
||||
self._failures[url] = time_mod.monotonic() - 1 # expired
|
||||
|
||||
mock_resp = MagicMock(status_code=200)
|
||||
mock_resp.headers = {"Content-Type": "image/png"}
|
||||
mock_resp.iter_content = MagicMock(return_value=[b"img"])
|
||||
with patch("apps.channels.api_views.requests.get", return_value=mock_resp), \
|
||||
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
|
||||
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
|
||||
self._fetch_logo(logo)
|
||||
self.assertNotIn(url, self._failures)
|
||||
|
||||
def test_request_exception_cached(self):
|
||||
"""Network errors are cached the same as non-200 responses."""
|
||||
import requests
|
||||
logo = Logo.objects.create(name="Timeout", url="https://timeout.com/logo.png")
|
||||
with patch("apps.channels.api_views.requests.get", side_effect=requests.Timeout("timed out")), \
|
||||
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
|
||||
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
|
||||
response = self._fetch_logo(logo)
|
||||
self.assertEqual(response.status_code, 404)
|
||||
self.assertIn("https://timeout.com/logo.png", self._failures)
|
||||
|
||||
def test_eviction_when_cache_exceeds_256(self):
|
||||
"""Stale entries are evicted when the cache grows past 256."""
|
||||
now = time_mod.monotonic()
|
||||
# Fill with 257 expired entries
|
||||
for i in range(257):
|
||||
self._failures[f"https://old-{i}.com/x.png"] = now - 1 # already expired
|
||||
|
||||
logo = Logo.objects.create(name="Trigger", url="https://trigger-evict.com/logo.png")
|
||||
import requests
|
||||
with patch("apps.channels.api_views.requests.get", side_effect=requests.ConnectionError("fail")), \
|
||||
patch("apps.channels.api_views.CoreSettings.get_default_user_agent_id", return_value="1"), \
|
||||
patch("apps.channels.api_views.UserAgent.objects.get", return_value=MagicMock(user_agent="Test/1.0")):
|
||||
self._fetch_logo(logo)
|
||||
|
||||
# Expired entries should be evicted
|
||||
old_entries = [k for k in self._failures if k.startswith("https://old-")]
|
||||
self.assertEqual(len(old_entries), 0)
|
||||
# New failure entry should exist
|
||||
self.assertIn("https://trigger-evict.com/logo.png", self._failures)
|
||||
@@ -0,0 +1,524 @@
|
||||
"""Tests for recent DVR fixes.
|
||||
|
||||
Covers:
|
||||
1. Collision avoidance: _build_output_paths checks both .mkv and .ts files
|
||||
2. Logo guard: _resolve_poster_for_program skips external APIs when title ≈ channel name
|
||||
3. Recording status lifecycle: status transitions visible via API
|
||||
4. Concat flags: error-tolerant ffmpeg flags used for segment concatenation
|
||||
5. Recovery skip-list: "recording" status NOT in terminal skip list
|
||||
"""
|
||||
import os
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
from rest_framework.test import APIRequestFactory, force_authenticate
|
||||
|
||||
from apps.channels.models import Channel, Recording
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_admin():
|
||||
from django.contrib.auth import get_user_model
|
||||
User = get_user_model()
|
||||
u, _ = User.objects.get_or_create(
|
||||
username="dvr_fixes_admin",
|
||||
defaults={"user_level": User.UserLevel.ADMIN},
|
||||
)
|
||||
u.set_password("pass")
|
||||
u.save()
|
||||
return u
|
||||
|
||||
|
||||
def _make_channel(name="Test Channel", number=100):
|
||||
return Channel.objects.create(channel_number=number, name=name)
|
||||
|
||||
|
||||
def _make_recording(channel, **overrides):
|
||||
now = timezone.now()
|
||||
defaults = {
|
||||
"channel": channel,
|
||||
"start_time": now - timedelta(hours=1),
|
||||
"end_time": now + timedelta(hours=1),
|
||||
"custom_properties": {},
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return Recording.objects.create(**defaults)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 1. Collision avoidance — _build_output_paths
|
||||
# =========================================================================
|
||||
|
||||
class CollisionAvoidanceTests(TestCase):
|
||||
"""_build_output_paths must increment the filename counter when
|
||||
EITHER the .mkv OR the .ts file already exists with size > 0."""
|
||||
|
||||
def _call(self, channel, program, start, end):
|
||||
from apps.channels.tasks import _build_output_paths
|
||||
return _build_output_paths(channel, program, start, end)
|
||||
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
|
||||
return_value="TV/{show}/{start}.mkv")
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
|
||||
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
|
||||
def test_no_collision_when_nothing_exists(self, _tv, _fb):
|
||||
"""Fresh path — no files exist, counter stays at 1."""
|
||||
ch = MagicMock(name="TestCh")
|
||||
ch.name = "TestCh"
|
||||
program = {"title": "My Show"}
|
||||
now = timezone.now()
|
||||
|
||||
def mock_stat(path):
|
||||
raise OSError("No such file")
|
||||
|
||||
with patch("os.stat", side_effect=mock_stat), \
|
||||
patch("os.makedirs"):
|
||||
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
|
||||
|
||||
# Should NOT have a _2 suffix
|
||||
self.assertNotIn("_2", final)
|
||||
self.assertTrue(final.endswith(".mkv"))
|
||||
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
|
||||
return_value="TV/{show}/{start}.mkv")
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
|
||||
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
|
||||
def test_collision_when_ts_exists_but_mkv_is_zero_bytes(self, _tv, _fb):
|
||||
"""Pre-restart scenario: MKV is 0-byte placeholder, TS has real data.
|
||||
The old code only checked MKV size, so it would reuse the path.
|
||||
The fix also checks TS, so it must increment."""
|
||||
ch = MagicMock(name="TestCh")
|
||||
ch.name = "TestCh"
|
||||
program = {"title": "My Show"}
|
||||
now = timezone.now()
|
||||
|
||||
def mock_stat(path):
|
||||
if "_2" in path:
|
||||
raise OSError("No such file")
|
||||
result = MagicMock()
|
||||
if path.endswith('.mkv'):
|
||||
result.st_size = 0 # MKV is 0-byte placeholder
|
||||
elif path.endswith('.ts'):
|
||||
result.st_size = 5000000 # TS has real data from pre-restart
|
||||
else:
|
||||
result.st_size = 0
|
||||
return result
|
||||
|
||||
with patch("os.stat", side_effect=mock_stat), \
|
||||
patch("os.makedirs"):
|
||||
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
|
||||
|
||||
# Must have incremented to _2
|
||||
self.assertIn("_2", final, "Should increment counter when TS file has data")
|
||||
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
|
||||
return_value="TV/{show}/{start}.mkv")
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
|
||||
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
|
||||
def test_collision_when_mkv_has_data(self, _tv, _fb):
|
||||
"""Standard collision: MKV file has data, should increment."""
|
||||
ch = MagicMock(name="TestCh")
|
||||
ch.name = "TestCh"
|
||||
program = {"title": "My Show"}
|
||||
now = timezone.now()
|
||||
|
||||
def mock_stat(path):
|
||||
if "_2" in path:
|
||||
raise OSError("No such file")
|
||||
result = MagicMock()
|
||||
if path.endswith('.mkv'):
|
||||
result.st_size = 1000000 # MKV has data
|
||||
else:
|
||||
result.st_size = 0
|
||||
return result
|
||||
|
||||
with patch("os.stat", side_effect=mock_stat), \
|
||||
patch("os.makedirs"):
|
||||
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
|
||||
|
||||
self.assertIn("_2", final, "Should increment counter when MKV file has data")
|
||||
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
|
||||
return_value="TV/{show}/{start}.mkv")
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
|
||||
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
|
||||
def test_no_collision_when_both_zero_bytes(self, _tv, _fb):
|
||||
"""Both MKV and TS exist but are 0 bytes — no collision."""
|
||||
ch = MagicMock(name="TestCh")
|
||||
ch.name = "TestCh"
|
||||
program = {"title": "My Show"}
|
||||
now = timezone.now()
|
||||
|
||||
def mock_stat(path):
|
||||
result = MagicMock()
|
||||
result.st_size = 0 # All files empty
|
||||
return result
|
||||
|
||||
with patch("os.stat", side_effect=mock_stat), \
|
||||
patch("os.makedirs"):
|
||||
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
|
||||
|
||||
self.assertNotIn("_2", final, "Should NOT increment when all files are empty")
|
||||
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_fallback_template",
|
||||
return_value="TV/{show}/{start}.mkv")
|
||||
@patch("apps.channels.tasks.CoreSettings.get_dvr_tv_template",
|
||||
return_value="TV/{show}/S{season:02d}E{episode:02d}.mkv")
|
||||
def test_collision_increments_to_3_when_2_also_occupied(self, _tv, _fb):
|
||||
"""When both base and _2 are occupied, should go to _3."""
|
||||
ch = MagicMock(name="TestCh")
|
||||
ch.name = "TestCh"
|
||||
program = {"title": "My Show"}
|
||||
now = timezone.now()
|
||||
|
||||
def mock_stat(path):
|
||||
if "_3" in path:
|
||||
raise OSError("No such file")
|
||||
result = MagicMock()
|
||||
if path.endswith('.ts'):
|
||||
result.st_size = 5000000
|
||||
else:
|
||||
result.st_size = 0
|
||||
return result
|
||||
|
||||
with patch("os.stat", side_effect=mock_stat), \
|
||||
patch("os.makedirs"):
|
||||
final, ts, fname = self._call(ch, program, now, now + timedelta(hours=1))
|
||||
|
||||
self.assertIn("_3", final, "Should increment to _3 when base and _2 are occupied")
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 2. Logo guard — _resolve_poster_for_program
|
||||
# =========================================================================
|
||||
|
||||
class LogoGuardTests(TestCase):
|
||||
"""When the program title matches the channel name, external API
|
||||
searches (VOD, TMDB, OMDb, TVMaze, iTunes) must be skipped."""
|
||||
|
||||
def _call(self, channel_name, program, channel_logo_id=None):
|
||||
from apps.channels.tasks import _resolve_poster_for_program
|
||||
return _resolve_poster_for_program(channel_name, program, channel_logo_id)
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
def test_channel_name_as_title_skips_external_apis(self, mock_get):
|
||||
"""Title = 'USA A&E SD*', channel = 'USA A&E SD*' → no external calls."""
|
||||
program = {"title": "USA A&E SD*"}
|
||||
logo_id, url = self._call("USA A&E SD*", program, channel_logo_id=42)
|
||||
|
||||
# Should NOT have called any external APIs
|
||||
mock_get.assert_not_called()
|
||||
# Should fall back to channel logo
|
||||
self.assertEqual(logo_id, 42)
|
||||
self.assertIsNone(url)
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
def test_channel_name_normalized_match(self, mock_get):
|
||||
"""Title = 'fox news', channel = 'FOX-News*' → normalized match, skip APIs."""
|
||||
program = {"title": "fox news"}
|
||||
logo_id, url = self._call("FOX-News*", program, channel_logo_id=99)
|
||||
|
||||
mock_get.assert_not_called()
|
||||
self.assertEqual(logo_id, 99)
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
def test_real_title_still_searched(self, mock_get):
|
||||
"""Title = 'Breaking Bad' on channel 'AMC' → should try external APIs."""
|
||||
# Mock TVMaze returning a result
|
||||
mock_resp = MagicMock(ok=True, status_code=200)
|
||||
mock_resp.json.return_value = {
|
||||
"image": {"original": "https://tvmaze.com/breaking-bad.jpg"}
|
||||
}
|
||||
mock_get.return_value = mock_resp
|
||||
|
||||
program = {"title": "Breaking Bad"}
|
||||
logo_id, url = self._call("AMC", program)
|
||||
|
||||
# Should have made at least one external API call
|
||||
self.assertTrue(mock_get.called, "Should search external APIs for real titles")
|
||||
self.assertIsNotNone(url)
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
def test_no_title_skips_to_channel_logo(self, mock_get):
|
||||
"""No title at all → falls through to channel logo, no API calls."""
|
||||
program = {}
|
||||
logo_id, url = self._call("SomeChannel", program, channel_logo_id=55)
|
||||
|
||||
mock_get.assert_not_called()
|
||||
self.assertEqual(logo_id, 55)
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
def test_epg_image_still_used_even_when_title_is_channel_name(self, mock_get):
|
||||
"""Even when title = channel name, Stage 1 (EPG images) should still work."""
|
||||
from apps.epg.models import ProgramData, EPGSource, EPGData
|
||||
|
||||
# Create an EPG source + EPGData entry + program with an icon URL
|
||||
epg_source = EPGSource.objects.create(source_type="xmltv", name="Test EPG")
|
||||
epg_data = EPGData.objects.create(tvg_id="test.ch", epg_source=epg_source)
|
||||
prog = ProgramData.objects.create(
|
||||
epg=epg_data,
|
||||
title="Test Channel HD",
|
||||
start_time=timezone.now() - timedelta(hours=1),
|
||||
end_time=timezone.now() + timedelta(hours=1),
|
||||
custom_properties={"icon": "https://epg-cdn.com/test-icon.png"},
|
||||
)
|
||||
|
||||
program = {"title": "Test Channel HD", "id": prog.id}
|
||||
|
||||
# Mock _validate_url to return True for the icon URL
|
||||
with patch("apps.channels.tasks._validate_url", return_value=True):
|
||||
logo_id, url = self._call("Test Channel HD", program, channel_logo_id=10)
|
||||
|
||||
# EPG icon should still be used (Stage 1 doesn't depend on title guard)
|
||||
self.assertEqual(url, "https://epg-cdn.com/test-icon.png")
|
||||
mock_get.assert_not_called()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 3. Recording status lifecycle via API
|
||||
# =========================================================================
|
||||
|
||||
class RecordingStatusLifecycleTests(TestCase):
|
||||
"""Verify recording status transitions and that terminal recordings
|
||||
are properly filterable (supports the red-dot fix in guideUtils)."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = _make_channel("Status Test Channel", 200)
|
||||
self.user = _make_admin()
|
||||
self.factory = APIRequestFactory()
|
||||
|
||||
def _list_recordings(self):
|
||||
from apps.channels.api_views import RecordingViewSet
|
||||
request = self.factory.get("/api/channels/recordings/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"get": "list"})
|
||||
return view(request)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_stopped_recording_has_terminal_status(self, _ws):
|
||||
"""After stop, custom_properties.status = 'stopped'."""
|
||||
from apps.channels.api_views import RecordingViewSet
|
||||
|
||||
rec = _make_recording(self.channel, custom_properties={
|
||||
"status": "recording",
|
||||
"program": {"id": 1, "title": "Live Show"},
|
||||
})
|
||||
|
||||
request = self.factory.post(f"/api/channels/recordings/{rec.id}/stop/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "stop"})
|
||||
|
||||
with patch("apps.channels.signals.revoke_task"):
|
||||
response = view(request, pk=rec.id)
|
||||
|
||||
self.assertIn(response.status_code, [200, 204])
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties.get("status"), "stopped")
|
||||
|
||||
def test_listing_includes_status_in_custom_properties(self):
|
||||
"""API listing returns custom_properties with status field."""
|
||||
_make_recording(self.channel, custom_properties={
|
||||
"status": "recording",
|
||||
"program": {"id": 1, "title": "Recording Show"},
|
||||
})
|
||||
_make_recording(self.channel, custom_properties={
|
||||
"status": "stopped",
|
||||
"program": {"id": 2, "title": "Stopped Show"},
|
||||
})
|
||||
|
||||
response = self._list_recordings()
|
||||
self.assertEqual(response.status_code, 200)
|
||||
|
||||
statuses = [r["custom_properties"].get("status") for r in response.data]
|
||||
self.assertIn("recording", statuses)
|
||||
self.assertIn("stopped", statuses)
|
||||
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_delete_recording_removes_from_listing(self, _ws):
|
||||
"""Deleting a recording removes it from the listing entirely."""
|
||||
from apps.channels.api_views import RecordingViewSet
|
||||
|
||||
rec = _make_recording(self.channel, custom_properties={
|
||||
"status": "stopped",
|
||||
"program": {"id": 3, "title": "To Delete"},
|
||||
})
|
||||
rec_id = rec.id
|
||||
|
||||
request = self.factory.delete(f"/api/channels/recordings/{rec_id}/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"delete": "destroy"})
|
||||
|
||||
with patch("apps.channels.signals.revoke_task"):
|
||||
response = view(request, pk=rec_id)
|
||||
|
||||
self.assertIn(response.status_code, [200, 204])
|
||||
self.assertFalse(Recording.objects.filter(id=rec_id).exists())
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 4. Concat flags — error-tolerant ffmpeg
|
||||
# =========================================================================
|
||||
|
||||
class ConcatFlagsTests(TestCase):
|
||||
"""Verify that the finalize phase uses error-tolerant ffmpeg flags
|
||||
when concatenating pre-restart segments."""
|
||||
|
||||
def test_concat_command_includes_error_tolerant_flags(self):
|
||||
"""Inspect the source code to confirm error-tolerant flags are present.
|
||||
This is a static analysis test — no ffmpeg execution needed."""
|
||||
import inspect
|
||||
from apps.channels.tasks import run_recording
|
||||
source = inspect.getsource(run_recording)
|
||||
|
||||
# The concat subprocess.run call must include these flags
|
||||
self.assertIn("+genpts+igndts+discardcorrupt", source,
|
||||
"Concat must use +genpts+igndts+discardcorrupt fflags")
|
||||
self.assertIn("ignore_err", source,
|
||||
"Concat must use -err_detect ignore_err")
|
||||
self.assertIn("-f", source)
|
||||
self.assertIn("concat", source)
|
||||
|
||||
def test_concat_goes_directly_to_mkv(self):
|
||||
"""Concat must produce MKV directly (not intermediate .ts) to
|
||||
preserve timestamp boundaries and avoid playback freeze at splice."""
|
||||
import inspect
|
||||
from apps.channels.tasks import run_recording
|
||||
source = inspect.getsource(run_recording)
|
||||
|
||||
# Must contain reset_timestamps for proper segment boundary handling
|
||||
self.assertIn("reset_timestamps", source,
|
||||
"Concat must use -reset_timestamps 1 for seamless seeking")
|
||||
# Must write directly to final_path (MKV), not an intermediate .ts
|
||||
self.assertIn("_concat_did_remux", source,
|
||||
"Concat path must set flag to skip separate remux step")
|
||||
|
||||
def test_segment_time_metadata_present(self):
|
||||
"""Verify concat uses -segment_time_metadata for boundary awareness."""
|
||||
import inspect
|
||||
from apps.channels.tasks import run_recording
|
||||
source = inspect.getsource(run_recording)
|
||||
|
||||
self.assertIn("segment_time_metadata", source,
|
||||
"Concat must use -segment_time_metadata 1 for segment boundary handling")
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 5. Recovery skip-list
|
||||
# =========================================================================
|
||||
|
||||
class RecoverySkipListTests(TestCase):
|
||||
"""Verify that the recovery function does NOT skip 'recording' status,
|
||||
since that's the exact status recordings have when the server crashes."""
|
||||
|
||||
def test_recording_status_not_in_skip_list(self):
|
||||
"""Inspect recover_recordings_on_startup to ensure 'recording' is
|
||||
NOT treated as a terminal/skip state."""
|
||||
import inspect
|
||||
from apps.channels.tasks import recover_recordings_on_startup
|
||||
source = inspect.getsource(recover_recordings_on_startup)
|
||||
|
||||
# Find the skip condition line
|
||||
# It should be: if current_status in ("completed", "stopped"):
|
||||
# NOT: if current_status in ("completed", "stopped", "recording"):
|
||||
lines = source.split('\n')
|
||||
skip_line = None
|
||||
for line in lines:
|
||||
if 'current_status in' in line and ('completed' in line or 'stopped' in line):
|
||||
skip_line = line.strip()
|
||||
break
|
||||
|
||||
self.assertIsNotNone(skip_line, "Should find the skip-list condition")
|
||||
self.assertNotIn('"recording"', skip_line,
|
||||
"Skip list must NOT contain 'recording' — "
|
||||
"that's the status of crashed mid-stream recordings that need recovery")
|
||||
|
||||
@patch("core.utils.RedisClient")
|
||||
@patch("apps.channels.tasks.run_recording")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_recovery_processes_recording_status(self, _ws, mock_run, mock_redis_cls):
|
||||
"""A recording with status='recording' should be recovered, not skipped."""
|
||||
mock_redis_conn = MagicMock()
|
||||
mock_redis_conn.set.return_value = True # Acquire lock
|
||||
mock_redis_cls.get_client.return_value = mock_redis_conn
|
||||
|
||||
channel = _make_channel("Recovery Test", 300)
|
||||
now = timezone.now()
|
||||
rec = _make_recording(channel, custom_properties={
|
||||
"status": "recording",
|
||||
"program": {"title": "Crashed Show"},
|
||||
}, end_time=now + timedelta(hours=2))
|
||||
|
||||
from apps.channels.tasks import recover_recordings_on_startup
|
||||
|
||||
with patch("apps.channels.signals.revoke_task"):
|
||||
result = recover_recordings_on_startup()
|
||||
|
||||
# The recording should have been dispatched for recovery
|
||||
self.assertTrue(mock_run.apply_async.called,
|
||||
"Recording with status='recording' should be dispatched for recovery")
|
||||
|
||||
@patch("core.utils.RedisClient")
|
||||
@patch("apps.channels.tasks.run_recording")
|
||||
@patch("core.utils.send_websocket_update", side_effect=lambda *a, **kw: None)
|
||||
def test_recovery_skips_stopped_recordings(self, _ws, mock_run, mock_redis_cls):
|
||||
"""A recording with status='stopped' should be skipped by recovery."""
|
||||
mock_redis_conn = MagicMock()
|
||||
mock_redis_conn.set.return_value = True
|
||||
mock_redis_cls.get_client.return_value = mock_redis_conn
|
||||
|
||||
channel = _make_channel("Recovery Skip Test", 301)
|
||||
now = timezone.now()
|
||||
rec = _make_recording(channel, custom_properties={
|
||||
"status": "stopped",
|
||||
"program": {"title": "Finished Show"},
|
||||
}, end_time=now + timedelta(hours=2))
|
||||
|
||||
from apps.channels.tasks import recover_recordings_on_startup
|
||||
with patch("apps.channels.signals.revoke_task"):
|
||||
recover_recordings_on_startup()
|
||||
|
||||
# Should NOT have dispatched a recovery task
|
||||
mock_run.apply_async.assert_not_called()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# 6. Frontend red-dot filter (guideUtils.mapRecordingsByProgramId)
|
||||
# =========================================================================
|
||||
|
||||
class MapRecordingsByProgramIdTests(TestCase):
|
||||
"""These test the BACKEND side — confirming that recording status
|
||||
is preserved in the API response so the frontend can filter on it.
|
||||
|
||||
The actual frontend filtering is covered by frontend/src/pages/__tests__/DVR.test.jsx
|
||||
and the guideUtils code, but we verify the data contract here."""
|
||||
|
||||
def test_recording_custom_properties_status_persisted(self):
|
||||
"""Recording status in custom_properties survives save/load cycle."""
|
||||
channel = _make_channel("Red Dot Test", 400)
|
||||
rec = _make_recording(channel, custom_properties={
|
||||
"status": "stopped",
|
||||
"program": {"id": 42, "title": "A Show"},
|
||||
})
|
||||
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties["status"], "stopped")
|
||||
|
||||
def test_terminal_statuses_are_well_defined(self):
|
||||
"""Verify the terminal status set matches what the frontend uses."""
|
||||
# These are the statuses that should NOT show a red dot in the Guide
|
||||
terminal = {"stopped", "completed", "interrupted", "failed"}
|
||||
channel = _make_channel("Terminal Status Test", 410)
|
||||
|
||||
# Verify each status is a valid recording status
|
||||
for status in terminal:
|
||||
rec = _make_recording(channel, custom_properties={
|
||||
"status": status,
|
||||
"program": {"id": 100, "title": "Test"},
|
||||
})
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties["status"], status)
|
||||
@@ -0,0 +1,585 @@
|
||||
"""Tests for DVR recording scheduling with ClockedSchedule.
|
||||
|
||||
Uses ClockedSchedule instead of apply_async with countdown because Redis
|
||||
visibility_timeout (default 3600s) causes task redelivery for long countdowns,
|
||||
leading to duplicate recordings.
|
||||
"""
|
||||
from datetime import timedelta
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
from django_celery_beat.models import ClockedSchedule, PeriodicTask
|
||||
|
||||
from apps.channels.models import Channel, Recording
|
||||
from apps.channels.signals import (
|
||||
schedule_recording_task,
|
||||
revoke_task,
|
||||
_dvr_task_name,
|
||||
)
|
||||
|
||||
|
||||
class ScheduleRecordingTaskTests(TestCase):
|
||||
"""Tests for schedule_recording_task()."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
|
||||
def tearDown(self):
|
||||
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
|
||||
ClockedSchedule.objects.all().delete()
|
||||
|
||||
@patch("apps.channels.signals.run_recording")
|
||||
def test_future_recording_creates_periodic_task(self, mock_run_recording):
|
||||
"""Recordings in the future create a ClockedSchedule + PeriodicTask."""
|
||||
future_time = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future_time,
|
||||
end_time=future_time + timedelta(hours=1),
|
||||
)
|
||||
|
||||
task_id = schedule_recording_task(rec, eta=future_time)
|
||||
|
||||
expected_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(task_id, expected_name)
|
||||
|
||||
pt = PeriodicTask.objects.get(name=expected_name)
|
||||
self.assertTrue(pt.one_off)
|
||||
self.assertTrue(pt.enabled)
|
||||
self.assertEqual(pt.task, "apps.channels.tasks.run_recording")
|
||||
self.assertIsNotNone(pt.clocked)
|
||||
|
||||
# apply_async should not have been called
|
||||
mock_run_recording.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.run_recording")
|
||||
def test_immediate_recording_creates_periodic_task(self, mock_run_recording):
|
||||
"""Recordings starting now also use ClockedSchedule for consistency."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now,
|
||||
end_time=now + timedelta(hours=1),
|
||||
)
|
||||
|
||||
task_id = schedule_recording_task(rec, eta=now)
|
||||
|
||||
expected_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(task_id, expected_name)
|
||||
self.assertTrue(PeriodicTask.objects.filter(name=expected_name).exists())
|
||||
|
||||
@patch("apps.channels.signals.run_recording")
|
||||
def test_past_start_time_clamps_to_now(self, mock_run_recording):
|
||||
"""Recordings with past start_time get clamped to now."""
|
||||
past_time = timezone.now() - timedelta(minutes=5)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=past_time,
|
||||
end_time=timezone.now() + timedelta(hours=1),
|
||||
)
|
||||
|
||||
task_id = schedule_recording_task(rec, eta=past_time)
|
||||
|
||||
expected_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(task_id, expected_name)
|
||||
pt = PeriodicTask.objects.get(name=expected_name)
|
||||
# Clocked time should be >= now
|
||||
self.assertGreaterEqual(pt.clocked.clocked_time, past_time)
|
||||
|
||||
@patch("apps.channels.signals.run_recording")
|
||||
def test_reschedule_updates_existing_periodic_task(self, mock_run_recording):
|
||||
"""Calling schedule_recording_task twice updates the existing PeriodicTask."""
|
||||
future_time = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future_time,
|
||||
end_time=future_time + timedelta(hours=1),
|
||||
)
|
||||
|
||||
schedule_recording_task(rec, eta=future_time)
|
||||
|
||||
# Reschedule with a different time
|
||||
new_eta = future_time + timedelta(hours=1)
|
||||
schedule_recording_task(rec, eta=new_eta)
|
||||
|
||||
# Should still be exactly one PeriodicTask
|
||||
task_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(PeriodicTask.objects.filter(name=task_name).count(), 1)
|
||||
|
||||
@patch("apps.channels.signals.run_recording")
|
||||
def test_naive_eta_is_made_aware(self, mock_run_recording):
|
||||
"""A naive (timezone-unaware) eta is made timezone-aware."""
|
||||
from datetime import datetime
|
||||
naive_eta = datetime(2030, 6, 15, 14, 0, 0)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=timezone.now() + timedelta(hours=1),
|
||||
end_time=timezone.now() + timedelta(hours=2),
|
||||
)
|
||||
|
||||
task_id = schedule_recording_task(rec, eta=naive_eta)
|
||||
|
||||
expected_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(task_id, expected_name)
|
||||
pt = PeriodicTask.objects.get(name=expected_name)
|
||||
self.assertTrue(timezone.is_aware(pt.clocked.clocked_time))
|
||||
|
||||
|
||||
class RevokeTaskTests(TestCase):
|
||||
"""Tests for revoke_task()."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
|
||||
def tearDown(self):
|
||||
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
|
||||
ClockedSchedule.objects.all().delete()
|
||||
|
||||
def test_revoke_deletes_periodic_task_and_clocked_schedule(self):
|
||||
"""revoke_task deletes the PeriodicTask and orphaned ClockedSchedule."""
|
||||
eta = timezone.now() + timedelta(hours=5)
|
||||
clocked = ClockedSchedule.objects.create(clocked_time=eta)
|
||||
PeriodicTask.objects.create(
|
||||
name="dvr-recording-10",
|
||||
task="apps.channels.tasks.run_recording",
|
||||
clocked=clocked,
|
||||
one_off=True,
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
revoke_task("dvr-recording-10")
|
||||
|
||||
self.assertFalse(PeriodicTask.objects.filter(name="dvr-recording-10").exists())
|
||||
self.assertFalse(ClockedSchedule.objects.filter(id=clocked.id).exists())
|
||||
|
||||
def test_revoke_keeps_shared_clocked_schedule(self):
|
||||
"""ClockedSchedule is kept if another PeriodicTask still references it."""
|
||||
eta = timezone.now() + timedelta(hours=5)
|
||||
clocked = ClockedSchedule.objects.create(clocked_time=eta)
|
||||
PeriodicTask.objects.create(
|
||||
name="dvr-recording-10",
|
||||
task="apps.channels.tasks.run_recording",
|
||||
clocked=clocked,
|
||||
one_off=True,
|
||||
)
|
||||
PeriodicTask.objects.create(
|
||||
name="dvr-recording-11",
|
||||
task="apps.channels.tasks.run_recording",
|
||||
clocked=clocked,
|
||||
one_off=True,
|
||||
)
|
||||
|
||||
revoke_task("dvr-recording-10")
|
||||
|
||||
self.assertFalse(PeriodicTask.objects.filter(name="dvr-recording-10").exists())
|
||||
self.assertTrue(ClockedSchedule.objects.filter(id=clocked.id).exists())
|
||||
|
||||
@patch("apps.channels.signals.AsyncResult")
|
||||
def test_revoke_falls_back_to_async_result_for_legacy_ids(self, mock_async_result):
|
||||
"""revoke_task falls back to AsyncResult.revoke() for old-style UUIDs."""
|
||||
revoke_task("550e8400-e29b-41d4-a716-446655440000")
|
||||
|
||||
mock_async_result.assert_called_once_with("550e8400-e29b-41d4-a716-446655440000")
|
||||
mock_async_result.return_value.revoke.assert_called_once()
|
||||
|
||||
def test_revoke_none_is_noop(self):
|
||||
"""revoke_task(None) does nothing."""
|
||||
revoke_task(None) # Should not raise
|
||||
|
||||
def test_revoke_empty_string_is_noop(self):
|
||||
"""revoke_task('') does nothing."""
|
||||
revoke_task("") # Should not raise
|
||||
|
||||
|
||||
class DvrTaskNameTests(TestCase):
|
||||
"""Tests for the naming convention helper."""
|
||||
|
||||
def test_task_name_format(self):
|
||||
self.assertEqual(_dvr_task_name(42), "dvr-recording-42")
|
||||
|
||||
def test_task_name_fits_in_charfield(self):
|
||||
name = _dvr_task_name(999999999)
|
||||
self.assertLessEqual(len(name), 255)
|
||||
|
||||
|
||||
class SignalIntegrationTests(TestCase):
|
||||
"""Integration tests for the post_save / post_delete signal handlers."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
|
||||
def tearDown(self):
|
||||
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
|
||||
ClockedSchedule.objects.all().delete()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_post_save_creates_periodic_task_for_future_recording(self, mock_artwork):
|
||||
"""Saving a future Recording creates a PeriodicTask via post_save signal."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
)
|
||||
|
||||
rec.refresh_from_db()
|
||||
task_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(rec.task_id, task_name)
|
||||
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_post_delete_removes_periodic_task(self, mock_artwork):
|
||||
"""Deleting a Recording removes its PeriodicTask."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
)
|
||||
|
||||
rec.refresh_from_db()
|
||||
task_name = rec.task_id
|
||||
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
|
||||
|
||||
rec.delete()
|
||||
self.assertFalse(PeriodicTask.objects.filter(name=task_name).exists())
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_bulk_delete_cleans_up_all_periodic_tasks(self, mock_artwork):
|
||||
"""Bulk deleting recordings cleans up all their PeriodicTasks."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec_ids = []
|
||||
for i in range(5):
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future + timedelta(hours=i),
|
||||
end_time=future + timedelta(hours=i + 1),
|
||||
)
|
||||
rec_ids.append(rec.id)
|
||||
|
||||
for rid in rec_ids:
|
||||
self.assertTrue(
|
||||
PeriodicTask.objects.filter(name=f"dvr-recording-{rid}").exists()
|
||||
)
|
||||
|
||||
Recording.objects.filter(channel=self.channel).delete()
|
||||
|
||||
self.assertEqual(
|
||||
PeriodicTask.objects.filter(name__startswith="dvr-recording-").count(), 0
|
||||
)
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_post_save_schedules_currently_playing_recording(self, mock_artwork):
|
||||
"""A recording with past start_time but future end_time schedules immediately."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
past_start = timezone.now() - timedelta(minutes=30)
|
||||
future_end = timezone.now() + timedelta(minutes=30)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=past_start,
|
||||
end_time=future_end,
|
||||
)
|
||||
|
||||
rec.refresh_from_db()
|
||||
task_name = f"dvr-recording-{rec.id}"
|
||||
self.assertEqual(rec.task_id, task_name)
|
||||
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_post_save_skips_fully_past_recording(self, mock_artwork):
|
||||
"""A recording with both start_time and end_time in the past is not scheduled."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
past_start = timezone.now() - timedelta(hours=2)
|
||||
past_end = timezone.now() - timedelta(hours=1)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=past_start,
|
||||
end_time=past_end,
|
||||
)
|
||||
|
||||
rec.refresh_from_db()
|
||||
self.assertIsNone(rec.task_id)
|
||||
self.assertFalse(
|
||||
PeriodicTask.objects.filter(name=f"dvr-recording-{rec.id}").exists()
|
||||
)
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_pre_save_revokes_on_time_change(self, mock_artwork):
|
||||
"""Changing a recording's start_time revokes the old task and creates a new one."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
)
|
||||
|
||||
rec.refresh_from_db()
|
||||
old_task_name = rec.task_id
|
||||
self.assertTrue(PeriodicTask.objects.filter(name=old_task_name).exists())
|
||||
|
||||
# Change the start time — pre_save clears task_id, post_save reschedules
|
||||
new_future = future + timedelta(hours=3)
|
||||
rec.start_time = new_future
|
||||
rec.end_time = new_future + timedelta(hours=1)
|
||||
rec.save()
|
||||
|
||||
rec.refresh_from_db()
|
||||
# Old PeriodicTask should be deleted; new one should exist
|
||||
self.assertIsNotNone(rec.task_id)
|
||||
self.assertTrue(
|
||||
PeriodicTask.objects.filter(name=f"dvr-recording-{rec.id}").exists()
|
||||
)
|
||||
|
||||
|
||||
class IdempotencyGuardTests(TestCase):
|
||||
"""Tests for the idempotency guard in run_recording()."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
|
||||
@patch("apps.channels.tasks.get_channel_layer")
|
||||
def test_skips_if_already_recording(self, mock_layer):
|
||||
"""run_recording returns early if status is already 'recording'."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now,
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "recording", "started_at": str(now)},
|
||||
)
|
||||
|
||||
from apps.channels.tasks import run_recording as run_rec_task
|
||||
result = run_rec_task(rec.id, self.channel.id, str(now), str(now + timedelta(hours=1)))
|
||||
|
||||
self.assertIsNone(result)
|
||||
# get_channel_layer should not have been called (returned before)
|
||||
mock_layer.assert_not_called()
|
||||
|
||||
@patch("apps.channels.tasks.get_channel_layer")
|
||||
def test_skips_if_already_completed(self, mock_layer):
|
||||
"""run_recording returns early if status is already 'completed'."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=2),
|
||||
end_time=now - timedelta(hours=1),
|
||||
custom_properties={"status": "completed"},
|
||||
)
|
||||
|
||||
from apps.channels.tasks import run_recording as run_rec_task
|
||||
result = run_rec_task(rec.id, self.channel.id, str(rec.start_time), str(rec.end_time))
|
||||
|
||||
self.assertIsNone(result)
|
||||
mock_layer.assert_not_called()
|
||||
|
||||
@patch("apps.channels.tasks.get_channel_layer")
|
||||
def test_skips_if_already_stopped(self, mock_layer):
|
||||
"""run_recording returns early if status is already 'stopped' (user stopped it early)."""
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "stopped", "stopped_at": str(now)},
|
||||
)
|
||||
|
||||
from apps.channels.tasks import run_recording as run_rec_task
|
||||
result = run_rec_task(rec.id, self.channel.id, str(rec.start_time), str(rec.end_time))
|
||||
|
||||
self.assertIsNone(result)
|
||||
mock_layer.assert_not_called()
|
||||
|
||||
|
||||
class ArtworkPrefetchSignalGuardTests(TestCase):
|
||||
"""Tests that the post_save signal does not schedule artwork prefetch when
|
||||
the recording is in an active or terminal state."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
|
||||
def tearDown(self):
|
||||
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
|
||||
ClockedSchedule.objects.all().delete()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_artwork_prefetch_not_scheduled_when_status_recording(self, mock_artwork):
|
||||
"""post_save must NOT schedule artwork prefetch when status='recording'
|
||||
to prevent a race that overwrites the running task's status updates."""
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
custom_properties={"status": "recording"},
|
||||
)
|
||||
|
||||
# Simulate a save that run_recording itself might do mid-recording
|
||||
rec.custom_properties = {"status": "recording", "file_path": "/data/recordings/test.mkv"}
|
||||
rec.save(update_fields=["custom_properties"])
|
||||
|
||||
# apply_async was not called for the "recording" save
|
||||
mock_artwork.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_artwork_prefetch_not_scheduled_when_status_completed(self, mock_artwork):
|
||||
"""post_save must NOT schedule artwork prefetch when status='completed'."""
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
custom_properties={"status": "completed"},
|
||||
)
|
||||
|
||||
rec.custom_properties = {"status": "completed"}
|
||||
rec.save(update_fields=["custom_properties"])
|
||||
|
||||
mock_artwork.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_artwork_prefetch_not_scheduled_when_status_stopped(self, mock_artwork):
|
||||
"""post_save must NOT schedule artwork prefetch when status='stopped'."""
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
custom_properties={"status": "stopped"},
|
||||
)
|
||||
|
||||
rec.custom_properties = {"status": "stopped"}
|
||||
rec.save(update_fields=["custom_properties"])
|
||||
|
||||
mock_artwork.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_artwork_prefetch_scheduled_for_new_upcoming_recording(self, mock_artwork):
|
||||
"""post_save SHOULD schedule artwork prefetch for a newly created upcoming recording."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
custom_properties={}, # no status yet — should trigger prefetch
|
||||
)
|
||||
|
||||
self.assertTrue(mock_artwork.apply_async.called)
|
||||
|
||||
|
||||
class DestroyDvrClientIsolationTests(TestCase):
|
||||
"""Tests that deleting a recording only stops DVR clients when the
|
||||
recording is actively streaming — never for completed/upcoming recordings
|
||||
that could share a channel with an unrelated in-progress recording."""
|
||||
|
||||
def setUp(self):
|
||||
from django.contrib.auth import get_user_model
|
||||
from rest_framework.test import APIRequestFactory, force_authenticate
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
User = get_user_model()
|
||||
self.user = User.objects.create_user(
|
||||
username="dvr_test_admin", password="pass",
|
||||
user_level=User.UserLevel.ADMIN,
|
||||
)
|
||||
self.factory = APIRequestFactory()
|
||||
self.force_authenticate = force_authenticate
|
||||
|
||||
def _delete_recording(self, rec):
|
||||
from apps.channels.api_views import RecordingViewSet
|
||||
request = self.factory.delete(f"/api/channels/recordings/{rec.id}/")
|
||||
self.force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"delete": "destroy"})
|
||||
return view(request, pk=rec.id)
|
||||
|
||||
@patch("apps.channels.api_views._stop_dvr_clients")
|
||||
def test_destroy_completed_recording_does_not_stop_dvr_clients(self, mock_stop):
|
||||
"""Deleting a completed recording must NOT call _stop_dvr_clients."""
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=timezone.now() - timedelta(hours=2),
|
||||
end_time=timezone.now() - timedelta(hours=1),
|
||||
custom_properties={"status": "completed", "file_path": "/data/recordings/test.mkv"},
|
||||
)
|
||||
self._delete_recording(rec)
|
||||
mock_stop.assert_not_called()
|
||||
|
||||
@patch("apps.channels.api_views._stop_dvr_clients")
|
||||
def test_destroy_upcoming_recording_does_not_stop_dvr_clients(self, mock_stop):
|
||||
"""Deleting an upcoming (scheduled) recording must NOT call _stop_dvr_clients."""
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
custom_properties={},
|
||||
)
|
||||
self._delete_recording(rec)
|
||||
mock_stop.assert_not_called()
|
||||
|
||||
@patch("apps.channels.api_views._stop_dvr_clients")
|
||||
def test_destroy_active_recording_does_stop_dvr_clients(self, mock_stop):
|
||||
"""Deleting an in-progress recording MUST call _stop_dvr_clients."""
|
||||
mock_stop.return_value = 1
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(minutes=5),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "recording"},
|
||||
)
|
||||
self._delete_recording(rec)
|
||||
mock_stop.assert_called_once_with(str(self.channel.uuid), recording_id=rec.id)
|
||||
|
||||
|
||||
class PeriodicTaskCleanupOnExecutionTests(TestCase):
|
||||
"""Tests for PeriodicTask cleanup when run_recording starts."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=1, name="Test Channel")
|
||||
|
||||
def tearDown(self):
|
||||
PeriodicTask.objects.filter(name__startswith="dvr-recording-").delete()
|
||||
ClockedSchedule.objects.all().delete()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
@patch("apps.channels.tasks.get_channel_layer")
|
||||
def test_periodic_task_cleaned_up_on_execution(self, mock_layer, mock_artwork):
|
||||
"""When run_recording executes, it deletes its own PeriodicTask."""
|
||||
mock_layer.return_value = MagicMock()
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=future,
|
||||
end_time=future + timedelta(hours=1),
|
||||
custom_properties={},
|
||||
)
|
||||
|
||||
# post_save signal should have created the PeriodicTask
|
||||
task_name = f"dvr-recording-{rec.id}"
|
||||
self.assertTrue(PeriodicTask.objects.filter(name=task_name).exists())
|
||||
pt = PeriodicTask.objects.get(name=task_name)
|
||||
clocked_id = pt.clocked_id
|
||||
|
||||
from apps.channels.tasks import run_recording as run_rec_task
|
||||
# This will proceed past guards, clean up the PeriodicTask, then
|
||||
# eventually fail on the actual stream connection (expected)
|
||||
try:
|
||||
run_rec_task(rec.id, self.channel.id, str(future), str(future + timedelta(hours=1)))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
self.assertFalse(PeriodicTask.objects.filter(name=task_name).exists())
|
||||
self.assertFalse(ClockedSchedule.objects.filter(id=clocked_id).exists())
|
||||
@@ -0,0 +1,356 @@
|
||||
"""Tests for the DVR Stop/Cancel feature set.
|
||||
|
||||
Covers:
|
||||
- stop() endpoint
|
||||
- destroy() was_in_progress field in recording_cancelled WebSocket event
|
||||
- signals.py update_fields re-entrancy guard
|
||||
- run_recording race guard before status write
|
||||
- _stop_dvr_clients() DVR client isolation
|
||||
"""
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock, AsyncMock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
from rest_framework.test import APIRequestFactory, force_authenticate
|
||||
|
||||
from apps.channels.models import Channel, Recording
|
||||
from apps.channels.api_views import RecordingViewSet, _stop_dvr_clients
|
||||
|
||||
|
||||
def _make_admin():
|
||||
from django.contrib.auth import get_user_model
|
||||
User = get_user_model()
|
||||
u, _ = User.objects.get_or_create(
|
||||
username="stop_test_admin",
|
||||
defaults={"user_level": User.UserLevel.ADMIN},
|
||||
)
|
||||
u.set_password("pass")
|
||||
u.save()
|
||||
return u
|
||||
|
||||
|
||||
def _async_channel_layer_mock():
|
||||
layer = MagicMock()
|
||||
layer.group_send = AsyncMock()
|
||||
return layer
|
||||
|
||||
|
||||
class StopEndpointTests(TestCase):
|
||||
"""Tests for POST /api/channels/recordings/{id}/stop/"""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=99, name="Stop Test Channel")
|
||||
self.user = _make_admin()
|
||||
self.factory = APIRequestFactory()
|
||||
|
||||
def _stop(self, rec):
|
||||
request = self.factory.post(f"/api/channels/recordings/{rec.id}/stop/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "stop"})
|
||||
return view(request, pk=rec.id)
|
||||
|
||||
def _make_rec(self, status="recording"):
|
||||
now = timezone.now()
|
||||
return Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(hours=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": status},
|
||||
)
|
||||
|
||||
@patch("core.utils.send_websocket_update")
|
||||
@patch("threading.Thread")
|
||||
def test_stop_writes_status_to_db_before_returning(self, mock_thread, mock_ws):
|
||||
"""DB write is synchronous — run_recording polls for this."""
|
||||
mock_thread.return_value.start = MagicMock()
|
||||
rec = self._make_rec()
|
||||
response = self._stop(rec)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertTrue(response.data.get("success"))
|
||||
rec.refresh_from_db()
|
||||
self.assertEqual(rec.custom_properties.get("status"), "stopped")
|
||||
|
||||
@patch("core.utils.send_websocket_update")
|
||||
@patch("threading.Thread")
|
||||
def test_stop_writes_stopped_at_timestamp(self, mock_thread, mock_ws):
|
||||
mock_thread.return_value.start = MagicMock()
|
||||
rec = self._make_rec()
|
||||
self._stop(rec)
|
||||
rec.refresh_from_db()
|
||||
self.assertIn("stopped_at", rec.custom_properties)
|
||||
|
||||
def test_stop_calls_stop_dvr_clients_in_background(self):
|
||||
"""stop() spawns a background thread whose target calls _stop_dvr_clients."""
|
||||
rec = self._make_rec()
|
||||
|
||||
with patch("apps.channels.api_views._stop_dvr_clients", return_value=1) as mock_stop, \
|
||||
patch("core.utils.send_websocket_update"), \
|
||||
patch("threading.Thread") as mock_thread:
|
||||
mock_thread.return_value.start = MagicMock()
|
||||
self._stop(rec)
|
||||
|
||||
# Verify a daemon thread was spawned
|
||||
mock_thread.assert_called_once()
|
||||
thread_kwargs = mock_thread.call_args[1]
|
||||
self.assertTrue(thread_kwargs.get("daemon"), "Thread must be daemon")
|
||||
|
||||
# Execute the captured target with DB connection close patched out
|
||||
target = thread_kwargs["target"]
|
||||
with patch("apps.channels.api_views._stop_dvr_clients", return_value=1) as mock_stop2, \
|
||||
patch("apps.channels.signals.revoke_task", side_effect=Exception("skip")), \
|
||||
patch("django.db.connection") as mock_conn:
|
||||
target()
|
||||
|
||||
self.assertTrue(mock_stop2.called)
|
||||
args, kwargs = mock_stop2.call_args
|
||||
actual_rec_id = kwargs.get("recording_id") or (args[1] if len(args) > 1 else None)
|
||||
self.assertEqual(actual_rec_id, rec.id)
|
||||
|
||||
def test_stop_returns_404_for_nonexistent(self):
|
||||
request = self.factory.post("/api/channels/recordings/99999/stop/")
|
||||
force_authenticate(request, user=self.user)
|
||||
view = RecordingViewSet.as_view({"post": "stop"})
|
||||
self.assertEqual(view(request, pk=99999).status_code, 404)
|
||||
|
||||
@patch("core.utils.send_websocket_update")
|
||||
@patch("threading.Thread")
|
||||
def test_stop_idempotent_on_already_stopped(self, mock_thread, mock_ws):
|
||||
mock_thread.return_value.start = MagicMock()
|
||||
rec = self._make_rec(status="stopped")
|
||||
self.assertEqual(self._stop(rec).status_code, 200)
|
||||
|
||||
|
||||
class CancelDestroyWasInProgressTests(TestCase):
|
||||
"""was_in_progress field in the recording_cancelled WebSocket event."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=98, name="Cancel Test Channel")
|
||||
self.user = _make_admin()
|
||||
self.factory = APIRequestFactory()
|
||||
|
||||
def _delete(self, rec):
|
||||
request = self.factory.delete(f"/api/channels/recordings/{rec.id}/")
|
||||
force_authenticate(request, user=self.user)
|
||||
return RecordingViewSet.as_view({"delete": "destroy"})(request, pk=rec.id)
|
||||
|
||||
@patch("apps.channels.api_views._stop_dvr_clients", return_value=1)
|
||||
@patch("core.utils.send_websocket_update")
|
||||
def test_in_progress_sends_was_in_progress_true(self, mock_ws, _):
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(minutes=10),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "recording"},
|
||||
)
|
||||
self._delete(rec)
|
||||
payload = mock_ws.call_args[0][2]
|
||||
self.assertEqual(payload["type"], "recording_cancelled")
|
||||
self.assertTrue(payload["was_in_progress"])
|
||||
|
||||
@patch("core.utils.send_websocket_update")
|
||||
def test_completed_sends_was_in_progress_false(self, mock_ws):
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=timezone.now() - timedelta(hours=2),
|
||||
end_time=timezone.now() - timedelta(hours=1),
|
||||
custom_properties={"status": "completed"},
|
||||
)
|
||||
self._delete(rec)
|
||||
self.assertFalse(mock_ws.call_args[0][2]["was_in_progress"])
|
||||
|
||||
|
||||
class SignalUpdateFieldsReentrancyGuardTests(TestCase):
|
||||
"""update_fields guard in schedule_task_on_save prevents redundant WS events."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=97, name="Signal Guard Channel")
|
||||
|
||||
def _create_upcoming(self):
|
||||
future = timezone.now() + timedelta(hours=2)
|
||||
return Recording.objects.create(
|
||||
channel=self.channel, start_time=future,
|
||||
end_time=future + timedelta(hours=1), custom_properties={},
|
||||
)
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_custom_properties_save_skips_artwork(self, mock_artwork):
|
||||
rec = self._create_upcoming()
|
||||
mock_artwork.reset_mock()
|
||||
rec.custom_properties = {"poster_url": "https://example.com/p.jpg"}
|
||||
rec.save(update_fields=["custom_properties"])
|
||||
mock_artwork.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_task_id_save_skips_artwork(self, mock_artwork):
|
||||
rec = self._create_upcoming()
|
||||
mock_artwork.reset_mock()
|
||||
rec.task_id = "dvr-recording-999"
|
||||
rec.save(update_fields=["task_id"])
|
||||
mock_artwork.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_combined_metadata_save_skips_artwork(self, mock_artwork):
|
||||
rec = self._create_upcoming()
|
||||
mock_artwork.reset_mock()
|
||||
rec.task_id = "dvr-recording-1000"
|
||||
rec.custom_properties = {"poster_url": "x"}
|
||||
rec.save(update_fields=["custom_properties", "task_id"])
|
||||
mock_artwork.apply_async.assert_not_called()
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_creation_dispatches_artwork(self, mock_artwork):
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
self._create_upcoming()
|
||||
self.assertTrue(mock_artwork.apply_async.called)
|
||||
|
||||
@patch("apps.channels.signals.prefetch_recording_artwork")
|
||||
def test_scheduling_field_update_dispatches_artwork(self, mock_artwork):
|
||||
"""save(update_fields=['start_time']) is not a metadata save — dispatch runs."""
|
||||
mock_artwork.apply_async.return_value = MagicMock()
|
||||
rec = self._create_upcoming()
|
||||
mock_artwork.reset_mock()
|
||||
future = timezone.now() + timedelta(hours=3)
|
||||
rec.start_time = future
|
||||
rec.end_time = future + timedelta(hours=1)
|
||||
rec.save(update_fields=["start_time", "end_time"])
|
||||
mock_artwork.apply_async.assert_called()
|
||||
|
||||
|
||||
class RunRecordingRaceGuardTests(TestCase):
|
||||
"""Race guard: stop() fires between idempotency check and status write."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=96, name="Race Guard Channel")
|
||||
|
||||
def test_race_guard_exits_when_stopped_at_db_read(self):
|
||||
"""If Recording.objects.get() shows 'stopped', the task must exit
|
||||
without writing 'recording' to the DB."""
|
||||
from apps.channels.tasks import run_recording as run_rec
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(minutes=1),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={},
|
||||
)
|
||||
mock_layer = _async_channel_layer_mock()
|
||||
original_get = Recording.objects.get
|
||||
|
||||
def patched_get(*args, **kwargs):
|
||||
obj = original_get(*args, **kwargs)
|
||||
if kwargs.get("id") == rec.id or (args and args[0] == rec.id):
|
||||
obj.custom_properties = {"status": "stopped"}
|
||||
return obj
|
||||
|
||||
with patch("apps.channels.tasks.get_channel_layer", return_value=mock_layer), \
|
||||
patch("core.utils.log_system_event", side_effect=Exception("skip")), \
|
||||
patch.object(Recording.objects, "get", side_effect=patched_get):
|
||||
result = run_rec(
|
||||
rec.id, self.channel.id, str(rec.start_time), str(rec.end_time),
|
||||
)
|
||||
|
||||
self.assertIsNone(result)
|
||||
rec.refresh_from_db()
|
||||
self.assertNotEqual(
|
||||
rec.custom_properties.get("status"), "recording",
|
||||
"Race guard failed: task overwrote 'stopped' with 'recording'",
|
||||
)
|
||||
|
||||
def test_idempotency_guard_catches_stopped_before_channel_layer(self):
|
||||
"""When status='stopped' at the idempotency check, get_channel_layer is never called."""
|
||||
from apps.channels.tasks import run_recording as run_rec
|
||||
now = timezone.now()
|
||||
rec = Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=now - timedelta(minutes=5),
|
||||
end_time=now + timedelta(hours=1),
|
||||
custom_properties={"status": "stopped"},
|
||||
)
|
||||
with patch("apps.channels.tasks.get_channel_layer") as mock_get_layer:
|
||||
result = run_rec(
|
||||
rec.id, self.channel.id, str(rec.start_time), str(rec.end_time),
|
||||
)
|
||||
self.assertIsNone(result)
|
||||
mock_get_layer.assert_not_called()
|
||||
|
||||
|
||||
class StopDvrClientsTests(TestCase):
|
||||
"""_stop_dvr_clients() DVR client isolation."""
|
||||
|
||||
def setUp(self):
|
||||
self.channel = Channel.objects.create(channel_number=95, name="DVR Clients Channel")
|
||||
self._redis = "core.utils.RedisClient"
|
||||
self._sc = "apps.proxy.ts_proxy.services.channel_service.ChannelService.stop_client"
|
||||
self._sch = "apps.proxy.ts_proxy.services.channel_service.ChannelService.stop_channel"
|
||||
|
||||
def _mock_redis(self, client_ids, ua_map):
|
||||
r = MagicMock()
|
||||
r.smembers.return_value = {c.encode() for c in client_ids}
|
||||
def hget_side(key, field):
|
||||
ks = key if isinstance(key, str) else key.decode("utf-8", errors="replace")
|
||||
for cid, ua in ua_map.items():
|
||||
if cid in ks:
|
||||
return ua.encode() if isinstance(ua, str) else ua
|
||||
return b""
|
||||
r.hget.side_effect = hget_side
|
||||
return r
|
||||
|
||||
def test_returns_zero_when_redis_none(self):
|
||||
with patch(self._redis) as rc:
|
||||
rc.get_client.return_value = None
|
||||
self.assertEqual(_stop_dvr_clients(str(self.channel.uuid)), 0)
|
||||
|
||||
def test_stops_only_matching_client_when_recording_id_given(self):
|
||||
r = self._mock_redis(
|
||||
["client-a", "client-b"],
|
||||
{"client-a": "Dispatcharr-DVR/recording-42",
|
||||
"client-b": "Dispatcharr-DVR/recording-99"},
|
||||
)
|
||||
with patch(self._redis) as rc, patch(self._sc) as sc:
|
||||
rc.get_client.return_value = r
|
||||
result = _stop_dvr_clients(str(self.channel.uuid), recording_id=42)
|
||||
self.assertEqual(result, 1)
|
||||
stopped = [c[0][1] for c in sc.call_args_list]
|
||||
self.assertIn("client-a", stopped)
|
||||
self.assertNotIn("client-b", stopped)
|
||||
|
||||
def test_stops_all_dvr_clients_without_recording_id(self):
|
||||
r = self._mock_redis(
|
||||
["client-a", "client-b"],
|
||||
{"client-a": "Dispatcharr-DVR/recording-42",
|
||||
"client-b": "Dispatcharr-DVR/recording-99"},
|
||||
)
|
||||
with patch(self._redis) as rc, patch(self._sc) as sc:
|
||||
rc.get_client.return_value = r
|
||||
result = _stop_dvr_clients(str(self.channel.uuid))
|
||||
self.assertEqual(result, 2)
|
||||
|
||||
def test_skips_non_dvr_clients(self):
|
||||
r = self._mock_redis(
|
||||
["viewer", "dvr-client"],
|
||||
{"viewer": "Mozilla/5.0", "dvr-client": "Dispatcharr-DVR/recording-1"},
|
||||
)
|
||||
with patch(self._redis) as rc, patch(self._sc) as sc:
|
||||
rc.get_client.return_value = r
|
||||
result = _stop_dvr_clients(str(self.channel.uuid))
|
||||
self.assertEqual(result, 1)
|
||||
stopped = [c[0][1] for c in sc.call_args_list]
|
||||
self.assertNotIn("viewer", stopped)
|
||||
|
||||
def test_returns_zero_for_empty_channel(self):
|
||||
r = MagicMock()
|
||||
r.smembers.return_value = set()
|
||||
with patch(self._redis) as rc, patch(self._sc) as sc:
|
||||
rc.get_client.return_value = r
|
||||
self.assertEqual(_stop_dvr_clients(str(self.channel.uuid)), 0)
|
||||
sc.assert_not_called()
|
||||
|
||||
def test_never_calls_stop_channel(self):
|
||||
"""Must not stop the whole channel proxy — only individual clients."""
|
||||
r = self._mock_redis(["dvr-1"], {"dvr-1": "Dispatcharr-DVR/recording-1"})
|
||||
with patch(self._redis) as rc, patch(self._sc), patch(self._sch) as sch:
|
||||
rc.get_client.return_value = r
|
||||
_stop_dvr_clients(str(self.channel.uuid))
|
||||
sch.assert_not_called()
|
||||
@@ -0,0 +1,40 @@
|
||||
from datetime import datetime, timedelta
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.channels.models import Channel, RecurringRecordingRule, Recording
|
||||
from apps.channels.tasks import sync_recurring_rule_impl, purge_recurring_rule_impl
|
||||
|
||||
|
||||
class RecurringRecordingRuleTasksTests(TestCase):
|
||||
def test_sync_recurring_rule_creates_and_purges_recordings(self):
|
||||
now = timezone.now()
|
||||
channel = Channel.objects.create(channel_number=1, name='Test Channel')
|
||||
|
||||
start_time = (now + timedelta(minutes=15)).time().replace(second=0, microsecond=0)
|
||||
end_time = (now + timedelta(minutes=75)).time().replace(second=0, microsecond=0)
|
||||
|
||||
rule = RecurringRecordingRule.objects.create(
|
||||
channel=channel,
|
||||
days_of_week=[now.weekday()],
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
created = sync_recurring_rule_impl(rule.id, drop_existing=True, horizon_days=1)
|
||||
self.assertEqual(created, 1)
|
||||
|
||||
recording = Recording.objects.filter(custom_properties__rule__id=rule.id).first()
|
||||
self.assertIsNotNone(recording)
|
||||
self.assertEqual(recording.channel, channel)
|
||||
self.assertEqual(recording.custom_properties.get('rule', {}).get('id'), rule.id)
|
||||
|
||||
expected_start = timezone.make_aware(
|
||||
datetime.combine(recording.start_time.date(), start_time),
|
||||
timezone.get_current_timezone(),
|
||||
)
|
||||
self.assertLess(abs((recording.start_time - expected_start).total_seconds()), 60)
|
||||
|
||||
removed = purge_recurring_rule_impl(rule.id)
|
||||
self.assertEqual(removed, 1)
|
||||
self.assertFalse(Recording.objects.filter(custom_properties__rule__id=rule.id).exists())
|
||||
@@ -0,0 +1,718 @@
|
||||
"""Tests for series rule evaluation deduplication.
|
||||
|
||||
Unit tests verify the dedup logic in evaluate_series_rules_impl.
|
||||
Integration tests exercise the full path: EPG refresh → series rule
|
||||
evaluation → Recording creation → post_save signal chain.
|
||||
"""
|
||||
from datetime import timedelta
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
from django.test import TestCase
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.channels.models import Channel, Recording
|
||||
from apps.epg.models import EPGSource, EPGData, ProgramData
|
||||
from core.models import CoreSettings
|
||||
|
||||
|
||||
def _set_series_rules(rules):
|
||||
"""Helper to store series rules in CoreSettings."""
|
||||
CoreSettings.set_dvr_series_rules(rules)
|
||||
|
||||
|
||||
def _set_dvr_offsets(pre_min=0, post_min=0):
|
||||
"""Helper to store DVR pre/post offsets."""
|
||||
CoreSettings._update_group("dvr_settings", "DVR Settings", {
|
||||
"pre_offset_minutes": pre_min,
|
||||
"post_offset_minutes": post_min,
|
||||
})
|
||||
|
||||
|
||||
class SeriesRuleDedupBaseTestCase(TestCase):
|
||||
"""Shared setup for series rule dedup tests."""
|
||||
|
||||
def setUp(self):
|
||||
self.now = timezone.now()
|
||||
self.epg_source = EPGSource.objects.create(
|
||||
name="Test EPG", source_type="xmltv"
|
||||
)
|
||||
self.epg = EPGData.objects.create(
|
||||
tvg_id="test.channel.1",
|
||||
name="Test Channel EPG",
|
||||
epg_source=self.epg_source,
|
||||
)
|
||||
self.channel = Channel.objects.create(
|
||||
channel_number=1, name="Test Channel", epg_data=self.epg
|
||||
)
|
||||
|
||||
_set_series_rules([{
|
||||
"tvg_id": "test.channel.1",
|
||||
"mode": "all",
|
||||
"title": "Test Show",
|
||||
}])
|
||||
_set_dvr_offsets(pre_min=0, post_min=0)
|
||||
|
||||
def _create_program(self, hours_from_now=1, title="Test Show",
|
||||
sub_title="Episode 1", tvg_id="test.channel.1"):
|
||||
"""Create a ProgramData at the given offset."""
|
||||
start = self.now + timedelta(hours=hours_from_now)
|
||||
end = start + timedelta(hours=1)
|
||||
return ProgramData.objects.create(
|
||||
epg=self.epg,
|
||||
tvg_id=tvg_id,
|
||||
start_time=start,
|
||||
end_time=end,
|
||||
title=title,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
|
||||
def _simulate_epg_refresh(self, programs_data):
|
||||
"""Delete all ProgramData and recreate with new IDs (simulates EPG refresh)."""
|
||||
ProgramData.objects.filter(epg=self.epg).delete()
|
||||
new_programs = []
|
||||
for data in programs_data:
|
||||
prog = ProgramData.objects.create(epg=self.epg, **data)
|
||||
new_programs.append(prog)
|
||||
return new_programs
|
||||
|
||||
def _program_data_for_refresh(self, prog):
|
||||
"""Build the dict needed by _simulate_epg_refresh from a ProgramData."""
|
||||
return {
|
||||
"tvg_id": prog.tvg_id,
|
||||
"start_time": prog.start_time,
|
||||
"end_time": prog.end_time,
|
||||
"title": prog.title,
|
||||
"sub_title": prog.sub_title,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit tests: dedup logic in evaluate_series_rules_impl
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@patch("apps.channels.tasks.prefetch_recording_artwork")
|
||||
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
|
||||
class ProgramIdStabilityTests(SeriesRuleDedupBaseTestCase):
|
||||
"""Verify dedup works after EPG refresh changes ProgramData IDs."""
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_no_duplicate_after_epg_refresh(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Same program should not be recorded twice after EPG refresh."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
old_id = prog.id
|
||||
result1 = evaluate_series_rules_impl()
|
||||
self.assertEqual(result1["scheduled"], 1)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
new_programs = self._simulate_epg_refresh(
|
||||
[self._program_data_for_refresh(prog)]
|
||||
)
|
||||
self.assertNotEqual(old_id, new_programs[0].id)
|
||||
|
||||
result2 = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
self.assertEqual(result2["scheduled"], 0)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_no_duplicate_with_offsets_after_refresh(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Dedup works when DVR offsets shift Recording times away from program times."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
_set_dvr_offsets(pre_min=5, post_min=5)
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
result1 = evaluate_series_rules_impl()
|
||||
self.assertEqual(result1["scheduled"], 1)
|
||||
|
||||
rec = Recording.objects.first()
|
||||
self.assertEqual(rec.start_time, prog.start_time - timedelta(minutes=5))
|
||||
self.assertEqual(rec.end_time, prog.end_time + timedelta(minutes=5))
|
||||
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
result2 = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_different_episodes_still_recorded(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Different episodes on the same channel should each get a recording."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
self._create_program(hours_from_now=2, sub_title="Episode 1")
|
||||
self._create_program(hours_from_now=4, sub_title="Episode 2")
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(result["scheduled"], 2)
|
||||
self.assertEqual(Recording.objects.count(), 2)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_new_episode_after_refresh_is_recorded(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""A genuinely new episode appearing after EPG refresh should be recorded."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
prog = self._create_program(hours_from_now=2, sub_title="Episode 1")
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
self._simulate_epg_refresh([
|
||||
self._program_data_for_refresh(prog),
|
||||
{
|
||||
"tvg_id": "test.channel.1",
|
||||
"start_time": prog.end_time,
|
||||
"end_time": prog.end_time + timedelta(hours=1),
|
||||
"title": "Test Show",
|
||||
"sub_title": "Episode 2",
|
||||
},
|
||||
])
|
||||
|
||||
result2 = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 2)
|
||||
self.assertEqual(result2["scheduled"], 1)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_multiple_epg_refreshes_no_duplicates(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Multiple consecutive EPG refreshes should not accumulate duplicates."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
for _ in range(5):
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
evaluate_series_rules_impl()
|
||||
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
|
||||
@patch("apps.channels.tasks.prefetch_recording_artwork")
|
||||
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
|
||||
class ConcurrencyGuardTests(SeriesRuleDedupBaseTestCase):
|
||||
"""Verify the task lock prevents concurrent evaluation."""
|
||||
|
||||
def test_lock_acquired_and_released(self, mock_schedule, mock_artwork):
|
||||
"""evaluate_series_rules_impl acquires and releases the task lock."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
self._create_program(hours_from_now=2)
|
||||
|
||||
with patch("apps.channels.tasks.acquire_task_lock", return_value=True) as mock_lock, \
|
||||
patch("apps.channels.tasks.release_task_lock") as mock_release:
|
||||
evaluate_series_rules_impl()
|
||||
mock_lock.assert_called_once_with('evaluate_series_rules', 'all')
|
||||
mock_release.assert_called_once_with('evaluate_series_rules', 'all')
|
||||
|
||||
def test_skips_when_lock_held(self, mock_schedule, mock_artwork):
|
||||
"""Returns early with skip reason when lock is already held."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
self._create_program(hours_from_now=2)
|
||||
|
||||
with patch("apps.channels.tasks.acquire_task_lock", return_value=False):
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(result["scheduled"], 0)
|
||||
self.assertTrue(
|
||||
any(d.get("reason") == "concurrent evaluation in progress"
|
||||
for d in result["details"]),
|
||||
)
|
||||
self.assertEqual(Recording.objects.count(), 0)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_lock_released_on_exception(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Lock is released even if the inner implementation raises."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
with patch("apps.channels.tasks._evaluate_series_rules_locked",
|
||||
side_effect=RuntimeError("test error")):
|
||||
with self.assertRaises(RuntimeError):
|
||||
evaluate_series_rules_impl()
|
||||
mock_release.assert_called_once_with('evaluate_series_rules', 'all')
|
||||
|
||||
|
||||
@patch("apps.channels.tasks.prefetch_recording_artwork")
|
||||
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
|
||||
class SecondaryGuardTests(SeriesRuleDedupBaseTestCase):
|
||||
"""Verify the secondary DB guard uses stable program attributes."""
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_secondary_guard_catches_duplicate_with_offsets(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Secondary guard works with stale program IDs and DVR offsets."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
_set_dvr_offsets(pre_min=10, post_min=10)
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
|
||||
# Pre-existing recording with a stale program ID (from previous EPG refresh)
|
||||
Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=prog.start_time - timedelta(minutes=10),
|
||||
end_time=prog.end_time + timedelta(minutes=10),
|
||||
custom_properties={
|
||||
"program": {
|
||||
"id": 99999,
|
||||
"tvg_id": prog.tvg_id,
|
||||
"title": prog.title,
|
||||
"start_time": prog.start_time.isoformat(),
|
||||
"end_time": prog.end_time.isoformat(),
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
self.assertEqual(result["scheduled"], 0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Integration tests: full path from EPG refresh through recording creation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@patch("apps.channels.tasks.prefetch_recording_artwork")
|
||||
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
|
||||
class IntegrationEPGRefreshTests(SeriesRuleDedupBaseTestCase):
|
||||
"""End-to-end tests simulating the EPG refresh → evaluate → record flow.
|
||||
|
||||
These exercise the full signal chain: evaluate_series_rules_impl creates
|
||||
a Recording, the post_save signal fires schedule_recording_task, and
|
||||
subsequent evaluations (after EPG refresh) must not create duplicates.
|
||||
"""
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_single_episode_no_duplicates(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Simulate: create rule → evaluate → EPG refresh → re-evaluate.
|
||||
|
||||
The full recording lifecycle must result in exactly 1 recording.
|
||||
"""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
# Initial EPG data
|
||||
prog = self._create_program(hours_from_now=2, sub_title="Pilot")
|
||||
|
||||
# First evaluation creates the recording
|
||||
result1 = evaluate_series_rules_impl()
|
||||
self.assertEqual(result1["scheduled"], 1)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
# Verify the recording was created with correct program metadata
|
||||
rec = Recording.objects.first()
|
||||
self.assertEqual(rec.custom_properties["program"]["tvg_id"], "test.channel.1")
|
||||
self.assertEqual(rec.custom_properties["program"]["title"], "Test Show")
|
||||
self.assertEqual(
|
||||
rec.custom_properties["program"]["start_time"],
|
||||
prog.start_time.isoformat()
|
||||
)
|
||||
|
||||
# Verify the post_save signal scheduled a task
|
||||
mock_schedule.assert_called()
|
||||
initial_schedule_count = mock_schedule.call_count
|
||||
|
||||
# Simulate EPG refresh (programs get new DB IDs)
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
|
||||
# Re-evaluate after refresh (this is what EPG refresh triggers)
|
||||
result2 = evaluate_series_rules_impl()
|
||||
self.assertEqual(result2["scheduled"], 0)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
# No additional task scheduling should have occurred
|
||||
self.assertEqual(mock_schedule.call_count, initial_schedule_count)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_with_offsets_no_duplicates(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Full flow with DVR offsets: recording times differ from program times."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
_set_dvr_offsets(pre_min=5, post_min=10)
|
||||
prog = self._create_program(hours_from_now=3, sub_title="Episode 1")
|
||||
|
||||
result1 = evaluate_series_rules_impl()
|
||||
self.assertEqual(result1["scheduled"], 1)
|
||||
|
||||
rec = Recording.objects.first()
|
||||
# Verify offset-adjusted recording times
|
||||
self.assertEqual(rec.start_time, prog.start_time - timedelta(minutes=5))
|
||||
self.assertEqual(rec.end_time, prog.end_time + timedelta(minutes=10))
|
||||
# Verify original (unadjusted) program times in custom_properties
|
||||
self.assertEqual(
|
||||
rec.custom_properties["program"]["start_time"],
|
||||
prog.start_time.isoformat()
|
||||
)
|
||||
self.assertEqual(
|
||||
rec.custom_properties["program"]["end_time"],
|
||||
prog.end_time.isoformat()
|
||||
)
|
||||
|
||||
# EPG refresh + re-evaluate
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
result2 = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
self.assertEqual(result2["scheduled"], 0)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_multiple_episodes_across_refreshes(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""New episodes appear across multiple EPG refreshes; each recorded once."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
ep1 = self._create_program(hours_from_now=2, sub_title="Episode 1")
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
# EPG refresh adds episode 2 alongside episode 1
|
||||
ep1_data = self._program_data_for_refresh(ep1)
|
||||
ep2_start = ep1.end_time
|
||||
ep2_data = {
|
||||
"tvg_id": "test.channel.1",
|
||||
"start_time": ep2_start,
|
||||
"end_time": ep2_start + timedelta(hours=1),
|
||||
"title": "Test Show",
|
||||
"sub_title": "Episode 2",
|
||||
}
|
||||
self._simulate_epg_refresh([ep1_data, ep2_data])
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 2)
|
||||
|
||||
# Another EPG refresh adds episode 3
|
||||
ep3_start = ep2_start + timedelta(hours=1)
|
||||
ep3_data = {
|
||||
"tvg_id": "test.channel.1",
|
||||
"start_time": ep3_start,
|
||||
"end_time": ep3_start + timedelta(hours=1),
|
||||
"title": "Test Show",
|
||||
"sub_title": "Episode 3",
|
||||
}
|
||||
self._simulate_epg_refresh([ep1_data, ep2_data, ep3_data])
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 3)
|
||||
|
||||
# Final EPG refresh with no new episodes — count must stay at 3
|
||||
self._simulate_epg_refresh([ep1_data, ep2_data, ep3_data])
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 3)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_multiple_series_rules(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Multiple series rules on different channels, each evaluated correctly."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
# Second channel with its own EPG
|
||||
epg2 = EPGData.objects.create(
|
||||
tvg_id="test.channel.2",
|
||||
name="Channel 2 EPG",
|
||||
epg_source=self.epg_source,
|
||||
)
|
||||
channel2 = Channel.objects.create(
|
||||
channel_number=2, name="Test Channel 2", epg_data=epg2
|
||||
)
|
||||
|
||||
_set_series_rules([
|
||||
{"tvg_id": "test.channel.1", "mode": "all", "title": "Show A"},
|
||||
{"tvg_id": "test.channel.2", "mode": "all", "title": "Show B"},
|
||||
])
|
||||
|
||||
# Programs on both channels
|
||||
start1 = self.now + timedelta(hours=2)
|
||||
prog1 = ProgramData.objects.create(
|
||||
epg=self.epg, tvg_id="test.channel.1",
|
||||
start_time=start1, end_time=start1 + timedelta(hours=1),
|
||||
title="Show A", sub_title="Episode 1",
|
||||
)
|
||||
start2 = self.now + timedelta(hours=3)
|
||||
prog2 = ProgramData.objects.create(
|
||||
epg=epg2, tvg_id="test.channel.2",
|
||||
start_time=start2, end_time=start2 + timedelta(hours=1),
|
||||
title="Show B", sub_title="Episode 1",
|
||||
)
|
||||
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 2)
|
||||
self.assertEqual(Recording.objects.filter(channel=self.channel).count(), 1)
|
||||
self.assertEqual(Recording.objects.filter(channel=channel2).count(), 1)
|
||||
|
||||
# EPG refresh for both channels
|
||||
ProgramData.objects.filter(epg=self.epg).delete()
|
||||
ProgramData.objects.filter(epg=epg2).delete()
|
||||
ProgramData.objects.create(
|
||||
epg=self.epg, tvg_id="test.channel.1",
|
||||
start_time=start1, end_time=start1 + timedelta(hours=1),
|
||||
title="Show A", sub_title="Episode 1",
|
||||
)
|
||||
ProgramData.objects.create(
|
||||
epg=epg2, tvg_id="test.channel.2",
|
||||
start_time=start2, end_time=start2 + timedelta(hours=1),
|
||||
title="Show B", sub_title="Episode 1",
|
||||
)
|
||||
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 2,
|
||||
"No duplicates across multiple series rules after EPG refresh")
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_rapid_epg_refreshes_simulate_user_report(
|
||||
self, mock_release, mock_lock, mock_schedule, mock_artwork
|
||||
):
|
||||
"""Reproduce the user-reported scenario: series rule + multiple EPG refreshes
|
||||
causing count to balloon from 6 to 25 and 5 simultaneous recordings.
|
||||
|
||||
Simulates 6 episodes with 5 EPG refreshes (each assigning new ProgramData IDs).
|
||||
"""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
# Create 6 episodes (the user had "next of 6")
|
||||
episodes = []
|
||||
for i in range(6):
|
||||
start = self.now + timedelta(hours=2 + i * 2)
|
||||
episodes.append({
|
||||
"tvg_id": "test.channel.1",
|
||||
"start_time": start,
|
||||
"end_time": start + timedelta(hours=1),
|
||||
"title": "Test Show",
|
||||
"sub_title": f"Episode {i + 1}",
|
||||
})
|
||||
|
||||
# Create initial ProgramData
|
||||
for ep in episodes:
|
||||
ProgramData.objects.create(epg=self.epg, **ep)
|
||||
|
||||
# First evaluation: should create exactly 6 recordings
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 6)
|
||||
|
||||
# Simulate 5 EPG refreshes (the user saw count balloon to 25)
|
||||
for refresh_num in range(5):
|
||||
self._simulate_epg_refresh(episodes)
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(
|
||||
Recording.objects.count(), 6,
|
||||
f"After EPG refresh #{refresh_num + 1}, expected 6 recordings "
|
||||
f"but got {Recording.objects.count()}"
|
||||
)
|
||||
self.assertEqual(result["scheduled"], 0)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_recording_survives_program_removal_and_readd(
|
||||
self, mock_release, mock_lock, mock_schedule, mock_artwork
|
||||
):
|
||||
"""Program temporarily disappears from EPG then reappears — no duplicate."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
prog = self._create_program(hours_from_now=2, sub_title="Episode 1")
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
# EPG refresh removes the program entirely
|
||||
self._simulate_epg_refresh([])
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1,
|
||||
"Existing recording preserved when program disappears from EPG")
|
||||
|
||||
# EPG refresh adds the program back (new ID)
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1,
|
||||
"No duplicate when program reappears with new ID")
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_celery_task_wrapper_calls_impl(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""The @shared_task evaluate_series_rules delegates to _impl correctly."""
|
||||
from apps.channels.tasks import evaluate_series_rules
|
||||
|
||||
self._create_program(hours_from_now=2)
|
||||
result = evaluate_series_rules()
|
||||
self.assertEqual(result["scheduled"], 1)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
# Call again (simulating a second EPG refresh trigger)
|
||||
result2 = evaluate_series_rules()
|
||||
self.assertEqual(result2["scheduled"], 0)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_tvg_id_scoped_evaluation(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Scoped evaluation (tvg_id parameter) still prevents duplicates."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
result1 = evaluate_series_rules_impl(tvg_id="test.channel.1")
|
||||
self.assertEqual(result1["scheduled"], 1)
|
||||
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
result2 = evaluate_series_rules_impl(tvg_id="test.channel.1")
|
||||
self.assertEqual(result2["scheduled"], 0)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_full_flow_offset_change_between_refreshes(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Changing DVR offsets between EPG refreshes doesn't create duplicates.
|
||||
|
||||
Even though Recording.start_time/end_time change when offsets change,
|
||||
the dedup key uses the original program times from custom_properties.
|
||||
"""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
_set_dvr_offsets(pre_min=5, post_min=5)
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
rec = Recording.objects.first()
|
||||
original_start = rec.start_time
|
||||
original_end = rec.end_time
|
||||
|
||||
# Change offsets
|
||||
_set_dvr_offsets(pre_min=10, post_min=15)
|
||||
|
||||
# EPG refresh
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1,
|
||||
"Changing offsets between refreshes should not create duplicates")
|
||||
self.assertEqual(result["scheduled"], 0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Edge case tests: Redis unavailability, non-series recordings, robustness
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@patch("apps.channels.tasks.prefetch_recording_artwork")
|
||||
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
|
||||
class RedisUnavailabilityTests(SeriesRuleDedupBaseTestCase):
|
||||
"""Verify evaluation works when Redis is unavailable (lock cannot be acquired)."""
|
||||
|
||||
def test_proceeds_when_redis_down(self, mock_schedule, mock_artwork):
|
||||
"""Evaluation succeeds (with dedup guards) when Redis raises on lock acquire."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
self._create_program(hours_from_now=2)
|
||||
|
||||
with patch("apps.channels.tasks.acquire_task_lock",
|
||||
side_effect=ConnectionError("Redis unavailable")):
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(result["scheduled"], 1)
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
def test_dedup_still_works_without_lock(self, mock_schedule, mock_artwork):
|
||||
"""Dedup guards prevent duplicates even when the lock is unavailable."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
|
||||
# First call: Redis down, proceeds without lock
|
||||
with patch("apps.channels.tasks.acquire_task_lock",
|
||||
side_effect=ConnectionError("Redis unavailable")):
|
||||
evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1)
|
||||
|
||||
# EPG refresh
|
||||
self._simulate_epg_refresh([self._program_data_for_refresh(prog)])
|
||||
|
||||
# Second call: Redis still down
|
||||
with patch("apps.channels.tasks.acquire_task_lock",
|
||||
side_effect=ConnectionError("Redis unavailable")):
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(Recording.objects.count(), 1,
|
||||
"Dedup guards prevent duplicates even without lock")
|
||||
self.assertEqual(result["scheduled"], 0)
|
||||
|
||||
def test_lock_not_released_when_not_acquired(self, mock_schedule, mock_artwork):
|
||||
"""release_task_lock is not called if acquire raised an exception."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
self._create_program(hours_from_now=2)
|
||||
|
||||
with patch("apps.channels.tasks.acquire_task_lock",
|
||||
side_effect=ConnectionError("Redis unavailable")), \
|
||||
patch("apps.channels.tasks.release_task_lock") as mock_release:
|
||||
evaluate_series_rules_impl()
|
||||
mock_release.assert_not_called()
|
||||
|
||||
|
||||
@patch("apps.channels.tasks.prefetch_recording_artwork")
|
||||
@patch("apps.channels.signals.schedule_recording_task", return_value="mock-task-id")
|
||||
class NonSeriesRecordingTests(SeriesRuleDedupBaseTestCase):
|
||||
"""Verify non-series recordings don't interfere with series rule dedup."""
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_manual_recording_without_program_data_ignored(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Recordings without custom_properties.program are skipped by dedup key builder."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
# Manual recording with no program metadata
|
||||
Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=self.now + timedelta(hours=2),
|
||||
end_time=self.now + timedelta(hours=3),
|
||||
custom_properties={},
|
||||
)
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(result["scheduled"], 1)
|
||||
self.assertEqual(Recording.objects.count(), 2)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_recurring_rule_recording_does_not_interfere(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Recordings from recurring rules (custom_properties.rule) don't block series rules."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=self.now + timedelta(hours=2),
|
||||
end_time=self.now + timedelta(hours=3),
|
||||
custom_properties={"rule": {"id": 1, "name": "Daily News"}},
|
||||
)
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(result["scheduled"], 1)
|
||||
self.assertEqual(Recording.objects.count(), 2)
|
||||
|
||||
@patch("apps.channels.tasks.acquire_task_lock", return_value=True)
|
||||
@patch("apps.channels.tasks.release_task_lock")
|
||||
def test_recording_with_null_custom_properties_ignored(self, mock_release, mock_lock,
|
||||
mock_schedule, mock_artwork):
|
||||
"""Recordings with None custom_properties don't crash the dedup key builder."""
|
||||
from apps.channels.tasks import evaluate_series_rules_impl
|
||||
|
||||
Recording.objects.create(
|
||||
channel=self.channel,
|
||||
start_time=self.now + timedelta(hours=2),
|
||||
end_time=self.now + timedelta(hours=3),
|
||||
custom_properties=None,
|
||||
)
|
||||
|
||||
prog = self._create_program(hours_from_now=2)
|
||||
result = evaluate_series_rules_impl()
|
||||
self.assertEqual(result["scheduled"], 1)
|
||||
@@ -0,0 +1,385 @@
|
||||
"""Tests for ghost client detection and cleanup.
|
||||
|
||||
Covers:
|
||||
- ClientManager.remove_ghost_clients() pipelined EXISTS logic
|
||||
- channel_status detailed stats path removes ghost clients from Redis SET
|
||||
- channel_status basic stats path removes ghost clients and corrects count
|
||||
- _check_orphaned_metadata() validates client SET entries and cleans up
|
||||
channels where all clients are ghosts
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch, PropertyMock
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.proxy.ts_proxy.client_manager import ClientManager
|
||||
from apps.proxy.ts_proxy.constants import ChannelMetadataField, ChannelState
|
||||
from apps.proxy.ts_proxy.redis_keys import RedisKeys
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
CHANNEL_ID = "00000000-0000-0000-0000-000000000001"
|
||||
|
||||
|
||||
def _make_proxy_server(redis_client=None):
|
||||
"""Create a minimal mock ProxyServer with a redis_client."""
|
||||
server = MagicMock()
|
||||
server.redis_client = redis_client or MagicMock()
|
||||
server.stream_managers = {}
|
||||
server.client_managers = {}
|
||||
server.worker_id = "test-worker-1"
|
||||
return server
|
||||
|
||||
|
||||
def _metadata_for_channel(state="active"):
|
||||
"""Return a plausible channel metadata dict (bytes keys/values)."""
|
||||
return {
|
||||
ChannelMetadataField.STATE.encode(): state.encode(),
|
||||
ChannelMetadataField.URL.encode(): b"http://example.com/stream",
|
||||
ChannelMetadataField.STREAM_PROFILE.encode(): b"default",
|
||||
ChannelMetadataField.OWNER.encode(): b"test-worker-1",
|
||||
ChannelMetadataField.INIT_TIME.encode(): b"1773500000.0",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit tests for ClientManager.remove_ghost_clients()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class RemoveGhostClientsTests(TestCase):
|
||||
"""Directly exercises the static method that all callers rely on."""
|
||||
|
||||
def test_ghost_removed_and_returned(self):
|
||||
"""Client ID in SET with no metadata hash should be SREM'd."""
|
||||
redis = MagicMock()
|
||||
redis.smembers.return_value = {b"ghost_001"}
|
||||
|
||||
pipe = MagicMock()
|
||||
redis.pipeline.return_value = pipe
|
||||
pipe.execute.return_value = [False] # EXISTS → False
|
||||
|
||||
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
|
||||
|
||||
self.assertEqual(result, [b"ghost_001"])
|
||||
redis.srem.assert_called_once()
|
||||
|
||||
def test_live_client_preserved(self):
|
||||
"""Client with valid metadata hash should NOT be removed."""
|
||||
redis = MagicMock()
|
||||
redis.smembers.return_value = {b"live_001"}
|
||||
|
||||
pipe = MagicMock()
|
||||
redis.pipeline.return_value = pipe
|
||||
pipe.execute.return_value = [True] # EXISTS → True
|
||||
|
||||
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
|
||||
|
||||
self.assertEqual(result, [])
|
||||
redis.srem.assert_not_called()
|
||||
|
||||
def test_mixed_ghost_and_live(self):
|
||||
"""Only ghost clients should be removed; live ones preserved."""
|
||||
redis = MagicMock()
|
||||
redis.smembers.return_value = {b"ghost_001", b"live_001"}
|
||||
|
||||
pipe = MagicMock()
|
||||
redis.pipeline.return_value = pipe
|
||||
# Order matches list(smembers), which is non-deterministic —
|
||||
# map both IDs so the test is stable regardless of iteration order.
|
||||
client_id_list = list(redis.smembers.return_value)
|
||||
|
||||
def exists_results():
|
||||
return [
|
||||
b"ghost_001" not in cid.decode() == False
|
||||
for cid in client_id_list
|
||||
]
|
||||
|
||||
# Simpler: mock based on key content
|
||||
def pipe_exists(key):
|
||||
pass # just enqueued; results come from execute()
|
||||
|
||||
pipe.exists.side_effect = pipe_exists
|
||||
pipe.execute.return_value = [
|
||||
"live" in cid.decode() for cid in client_id_list
|
||||
]
|
||||
|
||||
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
|
||||
|
||||
self.assertEqual(len(result), 1)
|
||||
self.assertTrue(any(b"ghost" in cid for cid in result))
|
||||
redis.srem.assert_called_once()
|
||||
|
||||
def test_empty_set_returns_empty(self):
|
||||
"""No clients means nothing to clean."""
|
||||
redis = MagicMock()
|
||||
redis.smembers.return_value = set()
|
||||
|
||||
result = ClientManager.remove_ghost_clients(redis, CHANNEL_ID)
|
||||
|
||||
self.assertEqual(result, [])
|
||||
redis.pipeline.assert_not_called()
|
||||
|
||||
def test_pre_fetched_client_ids_skips_smembers(self):
|
||||
"""When client_ids is passed, SMEMBERS should not be called."""
|
||||
redis = MagicMock()
|
||||
pipe = MagicMock()
|
||||
redis.pipeline.return_value = pipe
|
||||
pipe.execute.return_value = [False]
|
||||
|
||||
pre_fetched = {b"ghost_001"}
|
||||
result = ClientManager.remove_ghost_clients(
|
||||
redis, CHANNEL_ID, client_ids=pre_fetched
|
||||
)
|
||||
|
||||
redis.smembers.assert_not_called()
|
||||
self.assertEqual(len(result), 1)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Detailed stats path: exercises get_detailed_channel_info()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
|
||||
class DetailedStatsGhostClientTests(TestCase):
|
||||
"""get_detailed_channel_info() should remove ghost clients whose metadata
|
||||
hash has expired from the Redis client SET."""
|
||||
|
||||
def _setup_redis(self, mock_proxy_cls, client_ids, hgetall_side_effect):
|
||||
"""Wire up a mock ProxyServer with controlled Redis responses."""
|
||||
redis = MagicMock()
|
||||
server = _make_proxy_server(redis)
|
||||
mock_proxy_cls.get_instance.return_value = server
|
||||
|
||||
redis.hgetall.side_effect = hgetall_side_effect
|
||||
redis.smembers.return_value = client_ids
|
||||
# buffer_index, ttl, exists all need safe defaults
|
||||
redis.get.return_value = b"10"
|
||||
redis.ttl.return_value = 300
|
||||
redis.exists.return_value = True
|
||||
return redis
|
||||
|
||||
def test_ghost_client_removed_from_set(self, mock_proxy_cls):
|
||||
"""Ghost client should be SREM'd and excluded from result."""
|
||||
from apps.proxy.ts_proxy.channel_status import ChannelStatus
|
||||
|
||||
def hgetall_side_effect(key):
|
||||
if "clients:" in key:
|
||||
return {} # ghost — metadata expired
|
||||
return _metadata_for_channel()
|
||||
|
||||
redis = self._setup_redis(
|
||||
mock_proxy_cls, {b"ghost_001"}, hgetall_side_effect
|
||||
)
|
||||
|
||||
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
|
||||
|
||||
self.assertEqual(result['client_count'], 0)
|
||||
self.assertEqual(len(result['clients']), 0)
|
||||
redis.srem.assert_called_once()
|
||||
|
||||
def test_live_client_preserved(self, mock_proxy_cls):
|
||||
"""Client with valid metadata should appear in results."""
|
||||
from apps.proxy.ts_proxy.channel_status import ChannelStatus
|
||||
|
||||
def hgetall_side_effect(key):
|
||||
if "clients:" in key:
|
||||
return {
|
||||
b'user_agent': b'VLC/3.0',
|
||||
b'worker_id': b'test-worker-1',
|
||||
b'connected_at': b'1773500000.0',
|
||||
}
|
||||
return _metadata_for_channel()
|
||||
|
||||
redis = self._setup_redis(
|
||||
mock_proxy_cls, {b"live_001"}, hgetall_side_effect
|
||||
)
|
||||
|
||||
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
|
||||
|
||||
self.assertEqual(result['client_count'], 1)
|
||||
self.assertEqual(len(result['clients']), 1)
|
||||
redis.srem.assert_not_called()
|
||||
|
||||
def test_mixed_ghost_and_live(self, mock_proxy_cls):
|
||||
"""Only ghost clients should be removed; live ones preserved."""
|
||||
from apps.proxy.ts_proxy.channel_status import ChannelStatus
|
||||
|
||||
def hgetall_side_effect(key):
|
||||
if "clients:" in key:
|
||||
if "ghost" in key:
|
||||
return {}
|
||||
return {
|
||||
b'user_agent': b'VLC/3.0',
|
||||
b'worker_id': b'test-worker-1',
|
||||
}
|
||||
return _metadata_for_channel()
|
||||
|
||||
redis = self._setup_redis(
|
||||
mock_proxy_cls, {b"ghost_001", b"live_001"}, hgetall_side_effect
|
||||
)
|
||||
|
||||
result = ChannelStatus.get_detailed_channel_info(CHANNEL_ID)
|
||||
|
||||
self.assertEqual(result['client_count'], 1)
|
||||
self.assertEqual(len(result['clients']), 1)
|
||||
redis.srem.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Basic stats path: exercises get_basic_channel_info()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
|
||||
class BasicStatsGhostClientTests(TestCase):
|
||||
"""get_basic_channel_info() should call remove_ghost_clients(), skip
|
||||
ghosts from display, and correct client_count."""
|
||||
|
||||
def _setup_redis(self, mock_proxy_cls, client_ids, ghost_ids):
|
||||
"""Wire up mock ProxyServer. ghost_ids controls which EXISTS return False."""
|
||||
redis = MagicMock()
|
||||
server = _make_proxy_server(redis)
|
||||
mock_proxy_cls.get_instance.return_value = server
|
||||
|
||||
redis.hgetall.return_value = _metadata_for_channel()
|
||||
redis.get.return_value = b"10" # buffer_index
|
||||
redis.scard.return_value = len(client_ids)
|
||||
redis.smembers.return_value = client_ids
|
||||
redis.hget.return_value = None # individual field lookups
|
||||
|
||||
# Pipeline for remove_ghost_clients
|
||||
pipe = MagicMock()
|
||||
redis.pipeline.return_value = pipe
|
||||
client_id_list = list(client_ids)
|
||||
pipe.execute.return_value = [
|
||||
cid not in ghost_ids for cid in client_id_list
|
||||
]
|
||||
|
||||
return redis
|
||||
|
||||
def test_ghost_removed_and_count_corrected(self, mock_proxy_cls):
|
||||
"""Ghost client should be cleaned and client_count decremented."""
|
||||
from apps.proxy.ts_proxy.channel_status import ChannelStatus
|
||||
|
||||
redis = self._setup_redis(
|
||||
mock_proxy_cls,
|
||||
client_ids={b"ghost_001"},
|
||||
ghost_ids={b"ghost_001"},
|
||||
)
|
||||
|
||||
result = ChannelStatus.get_basic_channel_info(CHANNEL_ID)
|
||||
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result['client_count'], 0)
|
||||
redis.srem.assert_called_once()
|
||||
|
||||
def test_live_client_count_preserved(self, mock_proxy_cls):
|
||||
"""Live clients should be counted correctly."""
|
||||
from apps.proxy.ts_proxy.channel_status import ChannelStatus
|
||||
|
||||
redis = self._setup_redis(
|
||||
mock_proxy_cls,
|
||||
client_ids={b"live_001"},
|
||||
ghost_ids=set(),
|
||||
)
|
||||
|
||||
result = ChannelStatus.get_basic_channel_info(CHANNEL_ID)
|
||||
|
||||
self.assertIsNotNone(result)
|
||||
self.assertEqual(result['client_count'], 1)
|
||||
redis.srem.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Orphaned channel cleanup: exercises _check_orphaned_metadata()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@patch("apps.proxy.ts_proxy.channel_status.ProxyServer")
|
||||
class OrphanedChannelGhostValidationTests(TestCase):
|
||||
"""_check_orphaned_metadata() should validate client SET entries when
|
||||
owner is dead and client_count > 0. If all clients are ghosts, it
|
||||
should clean up the channel."""
|
||||
|
||||
def _make_server_for_orphan_check(self, mock_proxy_cls, channel_id,
|
||||
client_ids, ghost_ids, owner="dead-worker"):
|
||||
"""Build a mock ProxyServer whose Redis state simulates an orphaned channel."""
|
||||
redis = MagicMock()
|
||||
server = _make_proxy_server(redis)
|
||||
mock_proxy_cls.get_instance.return_value = server
|
||||
|
||||
metadata_key = RedisKeys.channel_metadata(channel_id)
|
||||
metadata = _metadata_for_channel()
|
||||
metadata[ChannelMetadataField.OWNER.encode()] = owner.encode()
|
||||
|
||||
# scan returns the one channel metadata key
|
||||
redis.scan.return_value = (0, [metadata_key.encode()])
|
||||
redis.hgetall.return_value = metadata
|
||||
redis.scard.return_value = len(client_ids)
|
||||
redis.smembers.return_value = client_ids
|
||||
# Owner heartbeat is dead
|
||||
redis.exists.side_effect = lambda key: (
|
||||
False if "heartbeat" in key else True
|
||||
)
|
||||
|
||||
# Pipeline for remove_ghost_clients
|
||||
pipe = MagicMock()
|
||||
redis.pipeline.return_value = pipe
|
||||
client_id_list = list(client_ids)
|
||||
pipe.execute.return_value = [
|
||||
cid not in ghost_ids for cid in client_id_list
|
||||
]
|
||||
|
||||
return server, redis
|
||||
|
||||
def test_all_ghosts_triggers_cleanup(self, mock_proxy_cls):
|
||||
"""When all clients are ghosts, channel should be cleaned up."""
|
||||
from apps.proxy.ts_proxy.server import ProxyServer
|
||||
|
||||
channel_id = "00000000-0000-0000-0000-000000000005"
|
||||
server, redis = self._make_server_for_orphan_check(
|
||||
mock_proxy_cls, channel_id,
|
||||
client_ids={b"ghost_001", b"ghost_002"},
|
||||
ghost_ids={b"ghost_001", b"ghost_002"},
|
||||
)
|
||||
|
||||
# Call the real method on a real-ish ProxyServer
|
||||
# The method lives on the server instance, so invoke it directly.
|
||||
# We need to call _check_orphaned_metadata on the actual server mock,
|
||||
# but it's a MagicMock. Instead, test via remove_ghost_clients directly
|
||||
# and verify the cleanup decision logic.
|
||||
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
|
||||
real_count = max(0, len({b"ghost_001", b"ghost_002"}) - len(stale_ids))
|
||||
|
||||
self.assertEqual(len(stale_ids), 2)
|
||||
self.assertEqual(real_count, 0)
|
||||
redis.srem.assert_called_once()
|
||||
|
||||
def test_mixed_preserves_live_clients(self, mock_proxy_cls):
|
||||
"""When some clients are live, real_count should be > 0."""
|
||||
channel_id = "00000000-0000-0000-0000-000000000006"
|
||||
server, redis = self._make_server_for_orphan_check(
|
||||
mock_proxy_cls, channel_id,
|
||||
client_ids={b"ghost_001", b"live_001"},
|
||||
ghost_ids={b"ghost_001"},
|
||||
)
|
||||
|
||||
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
|
||||
real_count = max(0, 2 - len(stale_ids))
|
||||
|
||||
self.assertEqual(len(stale_ids), 1)
|
||||
self.assertEqual(real_count, 1)
|
||||
|
||||
def test_no_ghosts_no_cleanup(self, mock_proxy_cls):
|
||||
"""When all clients are live, no SREM should be called."""
|
||||
channel_id = "00000000-0000-0000-0000-000000000007"
|
||||
server, redis = self._make_server_for_orphan_check(
|
||||
mock_proxy_cls, channel_id,
|
||||
client_ids={b"live_001"},
|
||||
ghost_ids=set(),
|
||||
)
|
||||
|
||||
stale_ids = ClientManager.remove_ghost_clients(redis, channel_id)
|
||||
|
||||
self.assertEqual(len(stale_ids), 0)
|
||||
redis.srem.assert_not_called()
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Tests for stuck INITIALIZING state fix.
|
||||
|
||||
Covers:
|
||||
- stream_manager.run() finally block: ownership check + state guard fallback
|
||||
- ChannelState.PRE_ACTIVE contains the correct states
|
||||
- INITIALIZING is included in the cleanup task grace period check
|
||||
"""
|
||||
import time
|
||||
import threading
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.proxy.ts_proxy.constants import ChannelMetadataField, ChannelState
|
||||
from apps.proxy.ts_proxy.redis_keys import RedisKeys
|
||||
from apps.proxy.ts_proxy.stream_manager import StreamManager
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
CHANNEL_ID = "00000000-0000-0000-0000-000000000001"
|
||||
|
||||
|
||||
def _make_stream_manager(tried_stream_ids=None, max_retries=3):
|
||||
"""Build a StreamManager via __new__ (bypasses __init__) with the
|
||||
minimum attributes required by the run() finally block."""
|
||||
sm = StreamManager.__new__(StreamManager)
|
||||
sm.channel_id = CHANNEL_ID
|
||||
sm.worker_id = "worker-1"
|
||||
sm.max_retries = max_retries
|
||||
sm.tried_stream_ids = tried_stream_ids if tried_stream_ids is not None else set()
|
||||
sm.running = False # while-loop exits immediately
|
||||
sm.connected = False
|
||||
sm.transcode_process_active = False
|
||||
sm._buffer_check_timers = []
|
||||
sm.url = "http://example.com/stream"
|
||||
sm.url_switching = False
|
||||
sm.url_switch_start_time = 0
|
||||
sm.url_switch_timeout = 30
|
||||
sm.stop_requested = False
|
||||
sm.stopping = False
|
||||
sm.socket = None
|
||||
sm.transcode_process = None
|
||||
sm.current_response = None
|
||||
sm.current_session = None
|
||||
sm.current_stream_id = None
|
||||
|
||||
buffer = MagicMock()
|
||||
buffer.redis_client = MagicMock()
|
||||
buffer.channel_id = CHANNEL_ID
|
||||
sm.buffer = buffer
|
||||
|
||||
return sm
|
||||
|
||||
|
||||
def _run_finally_block(sm, owner_value, current_state):
|
||||
"""Invoke StreamManager.run() so its finally block executes against real code.
|
||||
|
||||
Patches threading.Thread and ConfigHelper so the try-block is inert
|
||||
(self.running=False makes the while-loop exit immediately).
|
||||
|
||||
Returns True if the finally block wrote ERROR to Redis.
|
||||
"""
|
||||
redis = sm.buffer.redis_client
|
||||
|
||||
# Mock the owner key GET — the finally block calls redis.get(owner_key)
|
||||
def get_side_effect(key):
|
||||
if "owner" in key:
|
||||
return owner_value
|
||||
return None
|
||||
|
||||
redis.get.side_effect = get_side_effect
|
||||
|
||||
# Mock hget for state field lookup in the PRE_ACTIVE guard
|
||||
if current_state is not None:
|
||||
redis.hget.return_value = current_state.encode('utf-8')
|
||||
else:
|
||||
redis.hget.return_value = None
|
||||
|
||||
# Reset hset so we can detect whether ERROR was written
|
||||
redis.hset.reset_mock()
|
||||
redis.setex.reset_mock()
|
||||
|
||||
with patch.object(threading, 'Thread', return_value=MagicMock()):
|
||||
with patch('apps.proxy.ts_proxy.stream_manager.ConfigHelper') as mock_cfg:
|
||||
mock_cfg.max_stream_switches.return_value = 0
|
||||
mock_cfg.max_retries.return_value = sm.max_retries
|
||||
sm.run()
|
||||
|
||||
# Check if hset was called with ERROR state
|
||||
if redis.hset.called:
|
||||
mapping = redis.hset.call_args[1].get('mapping', {})
|
||||
return mapping.get(ChannelMetadataField.STATE) == ChannelState.ERROR
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# stream_manager.run() finally block: ownership + state guard behavior
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class StreamManagerFinallyBlockTests(TestCase):
|
||||
"""The run() finally block writes ERROR if the worker is still the owner
|
||||
(normal case) OR if ownership expired and the channel is still in a
|
||||
pre-active state (no new owner has taken over)."""
|
||||
|
||||
# --- Owner still valid: always write ERROR ---
|
||||
|
||||
def test_owner_writes_error_regardless_of_state(self):
|
||||
"""When we're still the owner, always write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
owner = sm.worker_id.encode('utf-8')
|
||||
self.assertTrue(_run_finally_block(sm, owner, ChannelState.ACTIVE))
|
||||
|
||||
def test_owner_writes_error_on_initializing(self):
|
||||
"""Owner + INITIALIZING = write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
owner = sm.worker_id.encode('utf-8')
|
||||
self.assertTrue(_run_finally_block(sm, owner, ChannelState.INITIALIZING))
|
||||
|
||||
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
|
||||
self.assertEqual(mapping[ChannelMetadataField.STATE], ChannelState.ERROR)
|
||||
|
||||
# --- Ownership expired, no new owner: use state guard ---
|
||||
|
||||
def test_no_owner_initializing_writes_error(self):
|
||||
"""Ownership expired + INITIALIZING = write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertTrue(_run_finally_block(sm, None, ChannelState.INITIALIZING))
|
||||
|
||||
def test_no_owner_connecting_writes_error(self):
|
||||
"""Ownership expired + CONNECTING = write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertTrue(_run_finally_block(sm, None, ChannelState.CONNECTING))
|
||||
|
||||
def test_no_owner_buffering_writes_error(self):
|
||||
"""Ownership expired + BUFFERING = write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertTrue(_run_finally_block(sm, None, ChannelState.BUFFERING))
|
||||
|
||||
def test_no_owner_waiting_for_clients_writes_error(self):
|
||||
"""Ownership expired + WAITING_FOR_CLIENTS = write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertTrue(_run_finally_block(sm, None, ChannelState.WAITING_FOR_CLIENTS))
|
||||
|
||||
def test_no_owner_active_does_not_write(self):
|
||||
"""Ownership expired + ACTIVE = do NOT write ERROR."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertFalse(_run_finally_block(sm, None, ChannelState.ACTIVE))
|
||||
|
||||
def test_no_owner_error_does_not_write(self):
|
||||
"""Ownership expired + already ERROR = do NOT write again."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertFalse(_run_finally_block(sm, None, ChannelState.ERROR))
|
||||
|
||||
def test_no_owner_no_state_does_not_write(self):
|
||||
"""Ownership expired + no state metadata = do NOT write."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertFalse(_run_finally_block(sm, None, None))
|
||||
|
||||
# --- New owner took over: never clobber ---
|
||||
|
||||
def test_new_owner_initializing_does_not_write(self):
|
||||
"""Another worker owns the channel — do NOT clobber."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertFalse(_run_finally_block(sm, b"other-worker", ChannelState.INITIALIZING))
|
||||
|
||||
def test_new_owner_active_does_not_write(self):
|
||||
"""Another worker owns the channel and is ACTIVE — do NOT write."""
|
||||
sm = _make_stream_manager()
|
||||
self.assertFalse(_run_finally_block(sm, b"other-worker", ChannelState.ACTIVE))
|
||||
|
||||
# --- Stopping key and error messages ---
|
||||
|
||||
def test_stopping_key_set_on_error_update(self):
|
||||
"""When ERROR is written, stopping key must also be set."""
|
||||
sm = _make_stream_manager()
|
||||
_run_finally_block(sm, None, ChannelState.INITIALIZING)
|
||||
|
||||
sm.buffer.redis_client.setex.assert_called_once()
|
||||
args = sm.buffer.redis_client.setex.call_args[0]
|
||||
self.assertIn("stopping", args[0])
|
||||
self.assertEqual(args[1], 60)
|
||||
|
||||
def test_error_message_includes_stream_count(self):
|
||||
"""When multiple streams were tried, error message reflects that."""
|
||||
sm = _make_stream_manager(tried_stream_ids={1, 2, 3})
|
||||
_run_finally_block(sm, None, ChannelState.INITIALIZING)
|
||||
|
||||
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
|
||||
error_msg = mapping[ChannelMetadataField.ERROR_MESSAGE]
|
||||
self.assertIn("3 stream options failed", error_msg)
|
||||
|
||||
def test_error_message_with_no_streams_tried(self):
|
||||
"""When no alternate streams were tried, shows retry count."""
|
||||
sm = _make_stream_manager(tried_stream_ids=set(), max_retries=5)
|
||||
_run_finally_block(sm, None, ChannelState.INITIALIZING)
|
||||
|
||||
mapping = sm.buffer.redis_client.hset.call_args[1]['mapping']
|
||||
error_msg = mapping[ChannelMetadataField.ERROR_MESSAGE]
|
||||
self.assertIn("5", error_msg)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ChannelState.PRE_ACTIVE: verify contents and immutability
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class PreActiveStateTests(TestCase):
|
||||
"""Verify PRE_ACTIVE contains the correct states and is immutable."""
|
||||
|
||||
def test_initializing_in_pre_active(self):
|
||||
self.assertIn(ChannelState.INITIALIZING, ChannelState.PRE_ACTIVE)
|
||||
|
||||
def test_connecting_in_pre_active(self):
|
||||
self.assertIn(ChannelState.CONNECTING, ChannelState.PRE_ACTIVE)
|
||||
|
||||
def test_buffering_in_pre_active(self):
|
||||
self.assertIn(ChannelState.BUFFERING, ChannelState.PRE_ACTIVE)
|
||||
|
||||
def test_waiting_for_clients_in_pre_active(self):
|
||||
self.assertIn(ChannelState.WAITING_FOR_CLIENTS, ChannelState.PRE_ACTIVE)
|
||||
|
||||
def test_active_not_in_pre_active(self):
|
||||
self.assertNotIn(ChannelState.ACTIVE, ChannelState.PRE_ACTIVE)
|
||||
|
||||
def test_error_not_in_pre_active(self):
|
||||
self.assertNotIn(ChannelState.ERROR, ChannelState.PRE_ACTIVE)
|
||||
|
||||
def test_pre_active_is_frozenset(self):
|
||||
self.assertIsInstance(ChannelState.PRE_ACTIVE, frozenset)
|
||||
@@ -0,0 +1,331 @@
|
||||
"""Tests for ts_proxy keepalive and stats-update behavior.
|
||||
|
||||
Covers:
|
||||
- stream_generator._should_send_keepalive() owner vs non-owner worker paths
|
||||
- stream_generator._should_send_keepalive() Redis last_data health check
|
||||
- client_manager._do_stats_update() error handling and WebSocket dispatch
|
||||
- client_manager.remove_client() non-blocking stats update
|
||||
- Keepalive/DVR-timeout timing invariants
|
||||
"""
|
||||
import threading
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _should_send_keepalive: owner worker path
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class OwnerWorkerKeepaliveTests(TestCase):
|
||||
"""Owner worker has a stream_manager; keepalive logic uses it directly."""
|
||||
|
||||
def _make_generator(self, healthy, at_buffer_head, consecutive_empty):
|
||||
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
|
||||
gen = StreamGenerator.__new__(StreamGenerator)
|
||||
gen.channel_id = "00000000-0000-0000-0000-000000000001"
|
||||
gen.client_id = "test-client"
|
||||
|
||||
buffer = MagicMock()
|
||||
buffer.index = 10 if at_buffer_head else 100
|
||||
gen.local_index = 10
|
||||
gen.buffer = buffer
|
||||
|
||||
stream_manager = MagicMock()
|
||||
stream_manager.healthy = healthy
|
||||
gen.stream_manager = stream_manager
|
||||
|
||||
gen.consecutive_empty = consecutive_empty
|
||||
return gen
|
||||
|
||||
def test_owner_healthy_returns_false(self):
|
||||
"""Owner worker, healthy stream -> no keepalive."""
|
||||
gen = self._make_generator(healthy=True, at_buffer_head=True, consecutive_empty=10)
|
||||
self.assertFalse(gen._should_send_keepalive(gen.local_index))
|
||||
|
||||
def test_owner_unhealthy_at_head_returns_true(self):
|
||||
"""Owner worker, unhealthy stream, at buffer head -> send keepalive."""
|
||||
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=10)
|
||||
self.assertTrue(gen._should_send_keepalive(gen.local_index))
|
||||
|
||||
def test_owner_unhealthy_not_at_head_returns_false(self):
|
||||
"""Owner worker, unhealthy stream, but NOT at buffer head -> no keepalive."""
|
||||
gen = self._make_generator(healthy=False, at_buffer_head=False, consecutive_empty=10)
|
||||
self.assertFalse(gen._should_send_keepalive(gen.local_index))
|
||||
|
||||
def test_owner_insufficient_consecutive_empty_returns_false(self):
|
||||
"""Owner worker, unhealthy, at head but consecutive_empty < 5 -> no keepalive."""
|
||||
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=3)
|
||||
self.assertFalse(gen._should_send_keepalive(gen.local_index))
|
||||
|
||||
def test_owner_exactly_5_consecutive_empty_returns_true(self):
|
||||
"""consecutive_empty == 5 is the minimum threshold."""
|
||||
gen = self._make_generator(healthy=False, at_buffer_head=True, consecutive_empty=5)
|
||||
self.assertTrue(gen._should_send_keepalive(gen.local_index))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _should_send_keepalive: non-owner worker path
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class NonOwnerWorkerKeepaliveTests(TestCase):
|
||||
"""Non-owner worker has stream_manager=None; health determined from Redis."""
|
||||
|
||||
def _make_generator(self, consecutive_empty=10):
|
||||
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
|
||||
gen = StreamGenerator.__new__(StreamGenerator)
|
||||
gen.channel_id = "00000000-0000-0000-0000-000000000002"
|
||||
gen.client_id = "test-client-nonowner"
|
||||
|
||||
buffer = MagicMock()
|
||||
buffer.index = 10
|
||||
gen.local_index = 10
|
||||
gen.buffer = buffer
|
||||
|
||||
gen.stream_manager = None # non-owner worker
|
||||
gen.consecutive_empty = consecutive_empty
|
||||
|
||||
# Attributes added by health-check throttling (set in __init__)
|
||||
gen._last_health_check_time = 0.0
|
||||
gen._last_health_check_result = False
|
||||
gen._health_check_interval = 2.0
|
||||
gen.proxy_server = None
|
||||
|
||||
return gen
|
||||
|
||||
def _mock_proxy_server(self, last_data_value):
|
||||
"""Return a mock ProxyServer with a redis_client pre-configured."""
|
||||
server = MagicMock()
|
||||
redis_client = MagicMock()
|
||||
server.redis_client = redis_client
|
||||
redis_client.get.return_value = last_data_value
|
||||
return server
|
||||
|
||||
def test_non_owner_fresh_data_returns_false(self):
|
||||
"""Non-owner, last_data < 10s ago -> stream healthy -> no keepalive."""
|
||||
gen = self._make_generator()
|
||||
fresh_ts = str(time.time() - 2.0).encode()
|
||||
server = self._mock_proxy_server(fresh_ts)
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertFalse(result, "Fresh data should NOT trigger keepalive")
|
||||
|
||||
def test_non_owner_stale_data_returns_true(self):
|
||||
"""Non-owner, last_data >= 10s ago -> stream unhealthy -> send keepalive."""
|
||||
gen = self._make_generator()
|
||||
stale_ts = str(time.time() - 12.0).encode()
|
||||
server = self._mock_proxy_server(stale_ts)
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertTrue(result, "Stale data (12s) should trigger keepalive")
|
||||
|
||||
def test_non_owner_exactly_at_timeout_returns_true(self):
|
||||
"""Data age exactly equal to CONNECTION_TIMEOUT (10s) -> send keepalive."""
|
||||
gen = self._make_generator()
|
||||
ts = str(time.time() - 10.0).encode()
|
||||
server = self._mock_proxy_server(ts)
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertTrue(result, "Data at exactly timeout threshold should trigger keepalive")
|
||||
|
||||
def test_non_owner_no_redis_key_returns_true(self):
|
||||
"""Non-owner, last_data key missing from Redis -> assume unhealthy."""
|
||||
gen = self._make_generator()
|
||||
server = self._mock_proxy_server(None)
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertTrue(result, "Missing last_data key should trigger keepalive")
|
||||
|
||||
def test_non_owner_redis_client_none_returns_false(self):
|
||||
"""Non-owner, redis_client is None (disconnected) -> conservative, no keepalive."""
|
||||
gen = self._make_generator()
|
||||
server = MagicMock()
|
||||
server.redis_client = None
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertFalse(result, "No redis_client -> conservative, no keepalive")
|
||||
|
||||
def test_non_owner_redis_exception_returns_false(self):
|
||||
"""Non-owner, Redis raises an exception -> conservative, no keepalive."""
|
||||
gen = self._make_generator()
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.side_effect = Exception("Redis error")
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertFalse(result, "Redis error -> conservative, no keepalive")
|
||||
|
||||
def test_non_owner_not_at_buffer_head_returns_false(self):
|
||||
"""Non-owner, NOT at buffer head -> no keepalive regardless of Redis."""
|
||||
gen = self._make_generator()
|
||||
gen.buffer.index = 100 # far ahead of local_index=10
|
||||
server = self._mock_proxy_server(None)
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertFalse(result)
|
||||
|
||||
def test_non_owner_insufficient_consecutive_empty_returns_false(self):
|
||||
"""Non-owner, at head, but consecutive_empty < 5 -> no keepalive."""
|
||||
gen = self._make_generator(consecutive_empty=2)
|
||||
stale_ts = str(time.time() - 30.0).encode()
|
||||
server = self._mock_proxy_server(stale_ts)
|
||||
|
||||
with patch("apps.proxy.ts_proxy.stream_generator.ProxyServer") as MockPS:
|
||||
MockPS.get_instance.return_value = server
|
||||
result = gen._should_send_keepalive(gen.local_index)
|
||||
|
||||
self.assertFalse(result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _do_stats_update: error handling and WebSocket dispatch
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class DoStatsUpdateTests(TestCase):
|
||||
"""_do_stats_update runs the actual Redis scan + WebSocket call."""
|
||||
|
||||
def _make_client_manager(self):
|
||||
from apps.proxy.ts_proxy.client_manager import ClientManager
|
||||
cm = ClientManager.__new__(ClientManager)
|
||||
cm.channel_id = "00000000-0000-0000-0000-000000000004"
|
||||
cm._heartbeat_running = False
|
||||
return cm
|
||||
|
||||
def test_do_stats_update_calls_send_websocket_update(self):
|
||||
"""_do_stats_update must call send_websocket_update with channel_stats."""
|
||||
cm = self._make_client_manager()
|
||||
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.scan.return_value = (0, [])
|
||||
|
||||
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update") as mock_ws, \
|
||||
patch("redis.Redis.from_url", return_value=mock_redis):
|
||||
cm._do_stats_update()
|
||||
|
||||
mock_ws.assert_called_once()
|
||||
event_type = mock_ws.call_args[0][1]
|
||||
self.assertEqual(event_type, "update")
|
||||
payload = mock_ws.call_args[0][2]
|
||||
self.assertEqual(payload["type"], "channel_stats")
|
||||
|
||||
def test_do_stats_update_does_not_raise_on_redis_error(self):
|
||||
"""Redis failure must be swallowed (logged), not propagated."""
|
||||
cm = self._make_client_manager()
|
||||
|
||||
with patch("redis.Redis.from_url", side_effect=Exception("Redis down")):
|
||||
try:
|
||||
cm._do_stats_update()
|
||||
except Exception as e:
|
||||
self.fail(f"_do_stats_update raised an exception: {e}")
|
||||
|
||||
def test_do_stats_update_scans_channel_client_keys(self):
|
||||
"""Must scan for ts_proxy:channel:*:clients pattern."""
|
||||
cm = self._make_client_manager()
|
||||
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.scan.return_value = (0, [])
|
||||
|
||||
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update"), \
|
||||
patch("redis.Redis.from_url", return_value=mock_redis):
|
||||
cm._do_stats_update()
|
||||
|
||||
scan_call = mock_redis.scan.call_args
|
||||
self.assertIn("ts_proxy:channel:*:clients", str(scan_call))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Integration: remove_client must not block on WebSocket
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ClientRemoveIntegrationTests(TestCase):
|
||||
"""When remove_client() fires, _trigger_stats_update must not block."""
|
||||
|
||||
def test_remove_client_does_not_block_on_websocket(self):
|
||||
"""remove_client() must return quickly even if WebSocket is slow."""
|
||||
from apps.proxy.ts_proxy.client_manager import ClientManager
|
||||
|
||||
cm = ClientManager.__new__(ClientManager)
|
||||
cm.channel_id = "00000000-0000-0000-0000-000000000005"
|
||||
cm._heartbeat_running = False
|
||||
cm.clients = {"test-client-1"}
|
||||
cm.last_heartbeat_time = {"test-client-1": time.time()}
|
||||
cm.last_active_time = time.time()
|
||||
cm.client_set_key = f"ts_proxy:channel:{cm.channel_id}:clients"
|
||||
cm.client_ttl = 60
|
||||
cm.worker_id = "worker-1"
|
||||
cm.proxy_server = MagicMock()
|
||||
cm.proxy_server.am_i_owner.return_value = False
|
||||
cm.lock = threading.Lock()
|
||||
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.hgetall.return_value = {b"ip_address": b"127.0.0.1"}
|
||||
mock_redis.scard.return_value = 1
|
||||
cm.redis_client = mock_redis
|
||||
|
||||
slow_ws_called = threading.Event()
|
||||
|
||||
def slow_websocket(*args, **kwargs):
|
||||
time.sleep(2.0)
|
||||
slow_ws_called.set()
|
||||
|
||||
start = time.time()
|
||||
with patch("apps.proxy.ts_proxy.client_manager.send_websocket_update", side_effect=slow_websocket):
|
||||
cm.remove_client("test-client-1")
|
||||
elapsed = time.time() - start
|
||||
|
||||
self.assertLess(elapsed, 1.0,
|
||||
f"remove_client() blocked for {elapsed:.2f}s waiting for WebSocket "
|
||||
f"(should dispatch to background thread and return immediately)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DVR timeout threshold vs keepalive timing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class KeepaliveTimingTests(TestCase):
|
||||
"""Verify that keepalive threshold gives sufficient margin before DVR timeout."""
|
||||
|
||||
def test_keepalive_threshold_less_than_dvr_timeout(self):
|
||||
"""CONNECTION_TIMEOUT (keepalive trigger) must be < DVR read timeout (15s)."""
|
||||
from apps.proxy.config import TSConfig as Config
|
||||
connection_timeout = getattr(Config, "CONNECTION_TIMEOUT", 10)
|
||||
dvr_read_timeout = 15 # hard-coded in run_recording: timeout=(10, 15)
|
||||
self.assertLess(
|
||||
connection_timeout,
|
||||
dvr_read_timeout,
|
||||
f"CONNECTION_TIMEOUT ({connection_timeout}s) must be < DVR timeout ({dvr_read_timeout}s) "
|
||||
f"so keepalives fire before DVR times out",
|
||||
)
|
||||
|
||||
def test_keepalive_interval_is_short(self):
|
||||
"""KEEPALIVE_INTERVAL must be short enough to send multiple keepalives in the gap."""
|
||||
from apps.proxy.config import TSConfig as Config
|
||||
interval = getattr(Config, "KEEPALIVE_INTERVAL", 0.5)
|
||||
connection_timeout = getattr(Config, "CONNECTION_TIMEOUT", 10)
|
||||
remaining_window = 15 - connection_timeout
|
||||
self.assertGreater(
|
||||
remaining_window / interval,
|
||||
3,
|
||||
f"KEEPALIVE_INTERVAL ({interval}s) is too long: only "
|
||||
f"{remaining_window/interval:.1f} keepalives would fit in the "
|
||||
f"{remaining_window}s window before DVR timeout",
|
||||
)
|
||||
@@ -0,0 +1,195 @@
|
||||
"""
|
||||
Unit tests for the keepalive duration cap in StreamGenerator._stream_data_generator.
|
||||
|
||||
Verifies that a client held in keepalive mode is disconnected after
|
||||
MAX_KEEPALIVE_DURATION seconds, and that the timer resets when real data resumes.
|
||||
"""
|
||||
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch, call
|
||||
from django.test import TestCase
|
||||
|
||||
|
||||
def _make_generator(consecutive_empty=10, local_index=10, buffer_index=10):
|
||||
"""Minimal StreamGenerator stub for testing _stream_data_generator logic."""
|
||||
from apps.proxy.ts_proxy.stream_generator import StreamGenerator
|
||||
|
||||
gen = StreamGenerator.__new__(StreamGenerator)
|
||||
gen.channel_id = "00000000-0000-0000-0000-000000000099"
|
||||
gen.client_id = "test-client-duration"
|
||||
gen.consecutive_empty = consecutive_empty
|
||||
gen.empty_reads = 0
|
||||
gen.local_index = local_index
|
||||
gen.bytes_sent = 0
|
||||
gen.chunks_sent = 0
|
||||
gen.last_yield_time = time.time()
|
||||
gen.stream_start_time = time.time()
|
||||
gen.last_stats_time = time.time()
|
||||
gen.last_stats_bytes = 0
|
||||
gen.current_rate = 0.0
|
||||
gen.last_ttl_refresh = time.time()
|
||||
gen.ttl_refresh_interval = 3
|
||||
gen.is_owner_worker = False
|
||||
gen.stream_manager = None
|
||||
gen._last_health_check_time = 0.0
|
||||
gen._last_health_check_result = False
|
||||
gen._health_check_interval = 2.0
|
||||
gen.proxy_server = None
|
||||
|
||||
buffer = MagicMock()
|
||||
buffer.index = buffer_index
|
||||
buffer.get_optimized_client_data.return_value = ([], local_index)
|
||||
buffer.find_oldest_available_chunk.return_value = None
|
||||
gen.buffer = buffer
|
||||
|
||||
return gen
|
||||
|
||||
|
||||
class KeepaliveDurationCapTests(TestCase):
|
||||
"""MAX_KEEPALIVE_DURATION cap disconnects clients stuck in keepalive mode."""
|
||||
|
||||
def _run_generator_to_break(self, gen, max_iterations=20):
|
||||
"""Drive _stream_data_generator until it breaks or hits iteration limit."""
|
||||
iterations = 0
|
||||
for _ in gen._stream_data_generator():
|
||||
iterations += 1
|
||||
if iterations >= max_iterations:
|
||||
break
|
||||
return iterations
|
||||
|
||||
def test_cap_fires_after_max_duration_exceeded(self):
|
||||
"""Generator exits when keepalive has run longer than MAX_KEEPALIVE_DURATION."""
|
||||
gen = _make_generator()
|
||||
|
||||
with patch.object(gen, '_check_resources', return_value=True), \
|
||||
patch.object(gen, '_should_send_keepalive', return_value=True), \
|
||||
patch.object(gen, '_is_ghost_client', return_value=False), \
|
||||
patch.object(gen, '_is_timeout', return_value=False), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.gevent') as mock_gevent, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
|
||||
|
||||
MockPS.get_instance.return_value = None
|
||||
MockConfig.KEEPALIVE_INTERVAL = 0
|
||||
MockConfig.MAX_KEEPALIVE_DURATION = 30
|
||||
|
||||
# First call: keepalive_start_time not yet set (returns current)
|
||||
# Second call: inside the cap check — simulate time elapsed > 30s
|
||||
mock_time.time.side_effect = [
|
||||
1000.0, # keepalive_start_time assignment
|
||||
1031.0, # cap check: 31s elapsed > 30s limit
|
||||
]
|
||||
|
||||
packets = list(gen._stream_data_generator())
|
||||
|
||||
# No packets should be yielded — cap fires before yield
|
||||
self.assertEqual(len(packets), 0)
|
||||
|
||||
def test_cap_does_not_fire_before_max_duration(self):
|
||||
"""Generator yields keepalive packets while within MAX_KEEPALIVE_DURATION."""
|
||||
gen = _make_generator()
|
||||
|
||||
call_count = 0
|
||||
|
||||
def time_side_effect():
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
# keepalive_start_time set at t=1000; cap checks always see <30s elapsed
|
||||
if call_count == 1:
|
||||
return 1000.0 # keepalive_start_time
|
||||
return 1010.0 # always 10s elapsed — under the 30s cap
|
||||
|
||||
with patch.object(gen, '_check_resources', side_effect=[True, True, False]), \
|
||||
patch.object(gen, '_should_send_keepalive', return_value=True), \
|
||||
patch.object(gen, '_is_ghost_client', return_value=False), \
|
||||
patch.object(gen, '_is_timeout', return_value=False), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
|
||||
|
||||
MockPS.get_instance.return_value = None
|
||||
MockConfig.KEEPALIVE_INTERVAL = 0
|
||||
MockConfig.MAX_KEEPALIVE_DURATION = 30
|
||||
mock_time.time.side_effect = time_side_effect
|
||||
|
||||
packets = list(gen._stream_data_generator())
|
||||
|
||||
# Two iterations with _check_resources=True should yield two keepalive packets
|
||||
self.assertGreater(len(packets), 0)
|
||||
|
||||
def test_timer_resets_when_real_data_resumes(self):
|
||||
"""keepalive_start_time is cleared to None when real chunks are received."""
|
||||
gen = _make_generator()
|
||||
|
||||
chunk = b'\x47' * 188
|
||||
real_chunks = ([chunk], gen.local_index + 1)
|
||||
no_chunks = ([], gen.local_index)
|
||||
|
||||
# Sequence: no data (keepalive), then real data, then stop
|
||||
gen.buffer.get_optimized_client_data.side_effect = [
|
||||
no_chunks, # iteration 1: keepalive
|
||||
real_chunks, # iteration 2: real data — should reset timer
|
||||
no_chunks, # iteration 3: keepalive again — timer restarts fresh
|
||||
]
|
||||
|
||||
captured_start_times = []
|
||||
|
||||
original_gen = gen
|
||||
|
||||
with patch.object(gen, '_check_resources', side_effect=[True, True, True, False]), \
|
||||
patch.object(gen, '_should_send_keepalive', return_value=True), \
|
||||
patch.object(gen, '_is_ghost_client', return_value=False), \
|
||||
patch.object(gen, '_is_timeout', return_value=False), \
|
||||
patch.object(gen, '_process_chunks', return_value=iter([chunk])), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
|
||||
|
||||
MockPS.get_instance.return_value = None
|
||||
MockConfig.KEEPALIVE_INTERVAL = 0
|
||||
MockConfig.MAX_KEEPALIVE_DURATION = 300
|
||||
mock_time.time.return_value = 1000.0
|
||||
|
||||
list(gen._stream_data_generator())
|
||||
|
||||
# Test passes if no exception and generator completes normally —
|
||||
# if the timer were NOT reset, the second keepalive block would
|
||||
# carry over the old start time rather than starting fresh.
|
||||
|
||||
def test_cap_uses_config_value(self):
|
||||
"""Cap threshold reads MAX_KEEPALIVE_DURATION from Config, not a hardcoded value."""
|
||||
gen = _make_generator()
|
||||
|
||||
with patch.object(gen, '_check_resources', return_value=True), \
|
||||
patch.object(gen, '_should_send_keepalive', return_value=True), \
|
||||
patch.object(gen, '_is_ghost_client', return_value=False), \
|
||||
patch.object(gen, '_is_timeout', return_value=False), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.create_ts_packet', return_value=b'\x00' * 188), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.ProxyServer') as MockPS, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.Config') as MockConfig, \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.gevent'), \
|
||||
patch('apps.proxy.ts_proxy.stream_generator.time') as mock_time:
|
||||
|
||||
MockPS.get_instance.return_value = None
|
||||
MockConfig.KEEPALIVE_INTERVAL = 0
|
||||
# Set a custom cap of 60s
|
||||
MockConfig.MAX_KEEPALIVE_DURATION = 60
|
||||
|
||||
mock_time.time.side_effect = [
|
||||
1000.0, # keepalive_start_time
|
||||
1050.0, # cap check: 50s elapsed — under 60s, should NOT fire
|
||||
1000.0, # last_yield_time update
|
||||
1070.0, # cap check on next iteration: 70s elapsed — fires
|
||||
]
|
||||
|
||||
packets = list(gen._stream_data_generator())
|
||||
|
||||
# First iteration: 50s < 60s cap — one keepalive yielded
|
||||
# Second iteration: 70s > 60s cap — generator exits
|
||||
self.assertEqual(len(packets), 1)
|
||||
@@ -0,0 +1,169 @@
|
||||
"""Tests for the _validate_url() helper in tasks.py.
|
||||
|
||||
Covers:
|
||||
- Rejection of None, empty, and non-string inputs
|
||||
- Non-HTTP URLs pass through without network requests
|
||||
- HTTP(S) URLs validated via HEAD request (2xx/3xx pass, 4xx/5xx fail)
|
||||
- Network errors (timeout, connection) treated as failures
|
||||
- Per-worker result cache: hits, expiry, eviction
|
||||
"""
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.channels.tasks import _validate_url, _url_validation_cache, _URL_CACHE_TTL
|
||||
|
||||
|
||||
class ValidateUrlInputTests(TestCase):
|
||||
"""Input validation — no network requests should be made."""
|
||||
|
||||
def setUp(self):
|
||||
_url_validation_cache.clear()
|
||||
|
||||
def test_none_returns_false(self):
|
||||
self.assertFalse(_validate_url(None))
|
||||
|
||||
def test_empty_string_returns_false(self):
|
||||
self.assertFalse(_validate_url(""))
|
||||
|
||||
def test_non_string_returns_false(self):
|
||||
self.assertFalse(_validate_url(123))
|
||||
self.assertFalse(_validate_url(["http://example.com"]))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_non_http_url_returns_true_without_request(self, mock_head):
|
||||
"""file:// and other non-HTTP schemes skip validation."""
|
||||
self.assertTrue(_validate_url("file:///local/path.jpg"))
|
||||
self.assertTrue(_validate_url("/data/images/poster.jpg"))
|
||||
mock_head.assert_not_called()
|
||||
|
||||
|
||||
class ValidateUrlNetworkTests(TestCase):
|
||||
"""HTTP(S) URL validation via HEAD request."""
|
||||
|
||||
def setUp(self):
|
||||
_url_validation_cache.clear()
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_200_returns_true(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=200)
|
||||
self.assertTrue(_validate_url("https://example.com/poster.jpg"))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_302_redirect_returns_true(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=302)
|
||||
self.assertTrue(_validate_url("https://example.com/redirect"))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_404_returns_false(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=404)
|
||||
self.assertFalse(_validate_url("https://dead-cdn.com/missing.jpg"))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_500_returns_false(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=500)
|
||||
self.assertFalse(_validate_url("https://broken.com/error"))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_timeout_returns_false(self, mock_head):
|
||||
import requests
|
||||
mock_head.side_effect = requests.Timeout("timed out")
|
||||
self.assertFalse(_validate_url("https://slow-cdn.com/poster.jpg"))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_connection_error_returns_false(self, mock_head):
|
||||
import requests
|
||||
mock_head.side_effect = requests.ConnectionError("refused")
|
||||
self.assertFalse(_validate_url("https://unreachable.com/poster.jpg"))
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_custom_timeout_passed_to_head(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=200)
|
||||
_validate_url("https://example.com/img.jpg", timeout=10)
|
||||
mock_head.assert_called_once_with(
|
||||
"https://example.com/img.jpg", timeout=10, allow_redirects=True
|
||||
)
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_405_falls_back_to_get(self, mock_head, mock_get):
|
||||
"""When HEAD returns 405, fall back to a ranged GET request."""
|
||||
mock_head.return_value = MagicMock(status_code=405)
|
||||
mock_resp = MagicMock(status_code=200)
|
||||
mock_get.return_value = mock_resp
|
||||
self.assertTrue(_validate_url("https://no-head.com/poster.jpg"))
|
||||
mock_get.assert_called_once()
|
||||
mock_resp.close.assert_called_once()
|
||||
|
||||
@patch("apps.channels.tasks.requests.get")
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_405_fallback_get_also_fails(self, mock_head, mock_get):
|
||||
"""When HEAD returns 405 and GET also fails, return False."""
|
||||
mock_head.return_value = MagicMock(status_code=405)
|
||||
mock_get.return_value = MagicMock(status_code=403)
|
||||
self.assertFalse(_validate_url("https://blocked.com/poster.jpg"))
|
||||
|
||||
|
||||
class ValidateUrlCacheTests(TestCase):
|
||||
"""Per-worker result caching."""
|
||||
|
||||
def setUp(self):
|
||||
_url_validation_cache.clear()
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_cache_hit_avoids_second_request(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=200)
|
||||
url = "https://cached.com/poster.jpg"
|
||||
self.assertTrue(_validate_url(url))
|
||||
self.assertTrue(_validate_url(url))
|
||||
mock_head.assert_called_once()
|
||||
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_cache_hit_returns_false_for_failed_url(self, mock_head):
|
||||
mock_head.return_value = MagicMock(status_code=404)
|
||||
url = "https://dead.com/missing.jpg"
|
||||
self.assertFalse(_validate_url(url))
|
||||
self.assertFalse(_validate_url(url))
|
||||
mock_head.assert_called_once()
|
||||
|
||||
@patch("apps.channels.tasks.time.monotonic")
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_cache_expiry_triggers_new_request(self, mock_head, mock_time):
|
||||
"""After TTL expires, a new HEAD request is made."""
|
||||
mock_head.return_value = MagicMock(status_code=200)
|
||||
url = "https://expiring.com/poster.jpg"
|
||||
|
||||
mock_time.return_value = 1000.0
|
||||
self.assertTrue(_validate_url(url))
|
||||
self.assertEqual(mock_head.call_count, 1)
|
||||
|
||||
# Within TTL — cache hit
|
||||
mock_time.return_value = 1000.0 + _URL_CACHE_TTL - 1
|
||||
self.assertTrue(_validate_url(url))
|
||||
self.assertEqual(mock_head.call_count, 1)
|
||||
|
||||
# Past TTL — new request
|
||||
mock_time.return_value = 1000.0 + _URL_CACHE_TTL + 1
|
||||
self.assertTrue(_validate_url(url))
|
||||
self.assertEqual(mock_head.call_count, 2)
|
||||
|
||||
@patch("apps.channels.tasks.time.monotonic")
|
||||
@patch("apps.channels.tasks.requests.head")
|
||||
def test_eviction_when_cache_exceeds_limit(self, mock_head, mock_time):
|
||||
"""Expired entries are evicted when cache grows past 512 entries."""
|
||||
mock_head.return_value = MagicMock(status_code=200)
|
||||
|
||||
# Fill cache with 513 entries at time 0
|
||||
mock_time.return_value = 0.0
|
||||
for i in range(513):
|
||||
_url_validation_cache[f"https://fill-{i}.com/img.jpg"] = (True, 0.0)
|
||||
|
||||
# Advance past TTL and add one more — triggers eviction
|
||||
mock_time.return_value = _URL_CACHE_TTL + 1
|
||||
_validate_url("https://trigger-eviction.com/img.jpg")
|
||||
|
||||
# All 513 old entries expired and should be evicted
|
||||
remaining = [k for k in _url_validation_cache if k.startswith("https://fill-")]
|
||||
self.assertEqual(len(remaining), 0)
|
||||
# The new entry should remain
|
||||
self.assertIn("https://trigger-eviction.com/img.jpg", _url_validation_cache)
|
||||
@@ -0,0 +1,12 @@
|
||||
from django.urls import path
|
||||
from .views import StreamDashboardView, channels_dashboard_view
|
||||
|
||||
app_name = 'channels_dashboard'
|
||||
|
||||
urlpatterns = [
|
||||
# Example “dashboard” routes for streams
|
||||
path('streams/', StreamDashboardView.as_view(), name='stream_dashboard'),
|
||||
|
||||
# Example “dashboard” route for channels
|
||||
path('channels/', channels_dashboard_view, name='channels_dashboard'),
|
||||
]
|
||||
@@ -0,0 +1,25 @@
|
||||
import threading
|
||||
|
||||
lock = threading.Lock()
|
||||
# Dictionary to track usage: {account_id: current_usage}
|
||||
active_streams_map = {}
|
||||
|
||||
def increment_stream_count(account):
|
||||
with lock:
|
||||
current_usage = active_streams_map.get(account.id, 0)
|
||||
current_usage += 1
|
||||
active_streams_map[account.id] = current_usage
|
||||
account.active_streams = current_usage
|
||||
account.save(update_fields=['active_streams'])
|
||||
|
||||
def decrement_stream_count(account):
|
||||
with lock:
|
||||
current_usage = active_streams_map.get(account.id, 0)
|
||||
if current_usage > 0:
|
||||
current_usage -= 1
|
||||
if current_usage == 0:
|
||||
del active_streams_map[account.id]
|
||||
else:
|
||||
active_streams_map[account.id] = current_usage
|
||||
account.active_streams = current_usage
|
||||
account.save(update_fields=['active_streams'])
|
||||
@@ -0,0 +1,41 @@
|
||||
from django.views import View
|
||||
from django.http import JsonResponse
|
||||
from django.utils.decorators import method_decorator
|
||||
from django.contrib.auth.decorators import login_required
|
||||
from django.views.decorators.csrf import csrf_exempt
|
||||
from django.shortcuts import render
|
||||
|
||||
from .models import Stream
|
||||
|
||||
@method_decorator(csrf_exempt, name='dispatch')
|
||||
@method_decorator(login_required, name='dispatch')
|
||||
class StreamDashboardView(View):
|
||||
"""
|
||||
Example “dashboard” style view for Streams
|
||||
"""
|
||||
def get(self, request, *args, **kwargs):
|
||||
streams = Stream.objects.values(
|
||||
'id', 'name', 'url',
|
||||
'channel_group', 'current_viewers'
|
||||
)
|
||||
return JsonResponse({'data': list(streams)}, safe=False)
|
||||
|
||||
def post(self, request, *args, **kwargs):
|
||||
"""
|
||||
Creates a new Stream from JSON data
|
||||
"""
|
||||
import json
|
||||
try:
|
||||
data = json.loads(request.body)
|
||||
new_stream = Stream.objects.create(**data)
|
||||
return JsonResponse({
|
||||
'id': new_stream.id,
|
||||
'message': 'Stream created successfully!'
|
||||
}, status=201)
|
||||
except Exception as e:
|
||||
return JsonResponse({'error': str(e)}, status=400)
|
||||
|
||||
|
||||
@login_required
|
||||
def channels_dashboard_view(request):
|
||||
return render(request, 'channels/channels.html')
|
||||
@@ -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
|
||||
@@ -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
Reference in New Issue
Block a user