This commit is contained in:
Nirvana
2026-01-27 11:57:39 +01:00
parent 0036344f66
commit 502b1d517f
2 changed files with 343 additions and 59 deletions
@@ -179,7 +179,7 @@ class DRMOperations:
return None
manifest_url = provider.get_manifest(channel_id, **kwargs)
logger.debug(f" GENERIC plugin: pssh from '{manifest_url}'")
logger.debug(f"GENERIC plugin: pssh from '{manifest_url}'")
if manifest_url:
pssh_data_list = self._extract_pssh_from_manifest(manifest_url, provider_name)
if pssh_data_list:
@@ -2,10 +2,11 @@
import base64
import xml.etree.ElementTree as ET
import re
from typing import Optional, Tuple, Set
from typing import Optional, Tuple, Set, Dict
from urllib.parse import urljoin, urlparse, quote, urlencode
from datetime import datetime, timezone
from email.utils import parsedate_to_datetime
from dataclasses import dataclass, field
from .logger import logger
@@ -15,18 +16,55 @@ ISO_8601_PERIOD_RE = re.compile(
)
@dataclass
class KeyConfiguration:
"""Configuration for DRM key management with validation and normalization."""
keys: Dict[str, str] = field(default_factory=dict)
single_key_mode: bool = field(init=False)
default_kid: Optional[str] = field(init=False, default=None)
default_key: Optional[str] = field(init=False, default=None)
def __post_init__(self):
"""Normalize and validate keys on initialization."""
normalized = {}
for kid, key in self.keys.items():
norm_kid = kid.replace("-", "").lower()
norm_key = key.replace("-", "").lower()
# Validate hex format (32 characters = 16 bytes)
if len(norm_kid) != 32 or not all(c in '0123456789abcdef' for c in norm_kid):
logger.warning(f"Invalid KID format (expected 32 hex chars): {kid}")
continue
if len(norm_key) != 32 or not all(c in '0123456789abcdef' for c in norm_key):
logger.warning(f"Invalid key format for KID {kid}: {key}")
continue
normalized[norm_kid] = norm_key
self.keys = normalized
self.single_key_mode = len(self.keys) <= 1
if self.single_key_mode and self.keys:
self.default_kid, self.default_key = next(iter(self.keys.items()))
logger.debug(f"Single key mode: KID={self.default_kid[:8]}...")
elif self.keys:
logger.debug(f"Multi-key mode: {len(self.keys)} keys available")
class MPDRewriter:
MPD_NAMESPACE = {"mpd": "urn:mpeg:dash:schema:mpd:2011"}
CENC_NAMESPACE = {"cenc": "urn:mpeg:cenc:2013"}
def __init__(
self,
media_proxy_url: str,
provider_proxy_url: Optional[str] = None,
clearkey_keyids: Optional[dict] = None,
self,
media_proxy_url: str,
provider_proxy_url: Optional[str] = None,
clearkey_keyids: Optional[dict] = None,
):
self.media_proxy_url = media_proxy_url.rstrip("/")
self.provider_proxy_url = provider_proxy_url
self.clearkey_keyids = clearkey_keyids or {}
self.key_config = KeyConfiguration(clearkey_keyids or {})
# Pre-calculate query params that don't change to save cycles during rewrite
self._static_params = {}
if self.provider_proxy_url:
@@ -44,26 +82,49 @@ class MPDRewriter:
return base64.urlsafe_b64decode(encoded.encode("utf-8")).decode("utf-8")
def build_proxy_url(
self,
original_url: str,
template_pattern: Optional[str] = None,
segment_type: Optional[str] = None,
is_encrypted: bool = False,
self,
original_url: str,
template_pattern: Optional[str] = None,
segment_type: Optional[str] = None,
is_encrypted: bool = False,
kid: Optional[str] = None, # Specific KID for this AdaptationSet
) -> str:
params = {"url": original_url, **self._static_params}
if self.clearkey_keyids and is_encrypted:
# Preserved Logic: Specific keys for specific segment types
if segment_type == "initialization":
params["kid"] = next(iter(self.clearkey_keyids.keys()))
elif segment_type == "media":
params["key"] = next(iter(self.clearkey_keyids.values()))
if self.key_config.keys and is_encrypted:
if self.key_config.single_key_mode:
# Single key mode: use the default key for everything
if segment_type == "initialization":
params["kid"] = self.key_config.default_kid
elif segment_type == "media":
params["key"] = self.key_config.default_key
else:
params["kid"] = self.key_config.default_kid
params["key"] = self.key_config.default_key
else:
kid, key = next(iter(self.clearkey_keyids.items()))
params["kid"], params["key"] = kid, key
# Multi-key mode: use specific KID if provided
if kid and kid in self.key_config.keys:
key = self.key_config.keys[kid]
if segment_type == "initialization":
params["kid"] = kid
elif segment_type == "media":
params["key"] = key
else:
params["kid"] = kid
params["key"] = key
else:
# Fallback to first key (should only happen if we couldn't extract KID)
if segment_type == "initialization":
params["kid"] = self.key_config.default_kid
elif segment_type == "media":
params["key"] = self.key_config.default_key
else:
params["kid"] = self.key_config.default_kid
params["key"] = self.key_config.default_key
logger.warning(f"No KID provided for encrypted segment, using fallback key")
encoded = self.encode_url(urlencode(params))
endpoint = "decrypt" if (self.clearkey_keyids and is_encrypted) else "proxy"
endpoint = "decrypt" if (self.key_config.keys and is_encrypted) else "proxy"
proxy_url = f"{self.media_proxy_url}/api/{endpoint}/{encoded}"
if template_pattern:
@@ -79,21 +140,31 @@ class MPDRewriter:
last_slash_before_template = url.rfind("/", 0, first_template_pos)
if last_slash_before_template == -1:
return "", url
return url[:last_slash_before_template], url[last_slash_before_template + 1 :]
return url[:last_slash_before_template], url[last_slash_before_template + 1:]
def rewrite_mpd(self, mpd_content: str, manifest_url: str) -> str:
try:
root = ET.fromstring(mpd_content)
ET.register_namespace("", self.MPD_NAMESPACE["mpd"])
# Single-pass tree preparation
encrypted_ids, as_id_to_kid = self._prepare_tree_and_extract_kids(root)
# Filter out encrypted AdaptationSets without available keys
if self.key_config.keys:
self._remove_adaptationsets_without_keys(root, as_id_to_kid)
else:
self._remove_all_encrypted_adaptationsets(root)
# Verify we have playable content remaining
remaining_sets = root.findall(".//mpd:AdaptationSet", self.MPD_NAMESPACE)
if not remaining_sets:
raise ValueError("No AdaptationSets remain after key filtering - manifest would be empty")
base_url = self._extract_base_url(root, manifest_url)
# Optimization: Single pass to clean tree and map encryption IDs
# Eliminates need for _is_descendant and O(N^2) searches
encrypted_as_ids = self._prepare_tree_and_get_encrypted_ids(root)
# Recursive rewrite with state-passing
self._rewrite_node(root, base_url, encrypted_as_ids, False)
# Rewrite URLs with appropriate keys
self._rewrite_node(root, base_url, encrypted_ids, as_id_to_kid, False, None, "")
rewritten = ET.tostring(root, encoding="unicode", method="xml")
if not rewritten.startswith("<?xml"):
@@ -103,40 +174,250 @@ class MPDRewriter:
logger.error(f"Failed to rewrite MPD: {e}")
raise
def _prepare_tree_and_get_encrypted_ids(self, root: ET.Element) -> Set[str]:
"""Cleans BaseURL/ContentProtection and identifies encrypted sets in one pass."""
def _prepare_tree_and_extract_kids(self, root: ET.Element) -> Tuple[Set[str], Dict[str, str]]:
"""
Single-pass optimization: clean tree, identify encrypted sets, extract KIDs.
Returns: (encrypted_adaptation_set_ids, as_id_to_kid_mapping)
"""
encrypted_ids = set()
for parent in root.iter():
# Remove BaseURL elements
for bu in list(parent.findall("mpd:BaseURL", self.MPD_NAMESPACE)):
parent.remove(bu)
as_id_to_kid = {}
# Process all Periods (handles multi-period manifests correctly)
for period in root.findall(".//mpd:Period", self.MPD_NAMESPACE):
period_id = period.get("id", "")
# Remove Period-level BaseURL elements
for bu in list(period.findall("mpd:BaseURL", self.MPD_NAMESPACE)):
period.remove(bu)
for adaptation_set in period.findall("mpd:AdaptationSet", self.MPD_NAMESPACE):
as_id = adaptation_set.get("id")
if not as_id:
as_id = str(id(adaptation_set))
# Make ID unique across periods
unique_id = f"{period_id}_{as_id}" if period_id else as_id
# Remove AdaptationSet-level BaseURL elements
for bu in list(adaptation_set.findall("mpd:BaseURL", self.MPD_NAMESPACE)):
adaptation_set.remove(bu)
# Process ContentProtection
cp_elements = list(adaptation_set.findall("mpd:ContentProtection", self.MPD_NAMESPACE))
# Map and remove ContentProtection
if parent.tag.endswith("AdaptationSet"):
cp_elements = parent.findall(
"mpd:ContentProtection", self.MPD_NAMESPACE
)
if cp_elements:
as_id = parent.get("id", str(id(parent)))
encrypted_ids.add(str(as_id))
encrypted_ids.add(unique_id)
# Extract KID only in multi-key mode
if not self.key_config.single_key_mode and self.key_config.keys:
kid = self._extract_kid_from_contentprotection(cp_elements, adaptation_set)
if kid:
normalized_kid = kid.replace("-", "").lower()
as_id_to_kid[unique_id] = normalized_kid
logger.debug(f"AdaptationSet {unique_id} KID: {normalized_kid[:8]}...")
else:
logger.debug(f"AdaptationSet {unique_id} encrypted but no KID found")
# Remove ContentProtection elements
for cp in cp_elements:
parent.remove(cp)
return encrypted_ids
adaptation_set.remove(cp)
return encrypted_ids, as_id_to_kid
def _extract_kid_from_contentprotection(
self,
cp_elements: list,
adaptation_set: ET.Element
) -> Optional[str]:
"""
Extract KID from ContentProtection elements.
Tries multiple methods per DASH specification.
"""
# Method 1: default_KID attribute (most common)
for cp in cp_elements:
default_kid = (
cp.get("default_KID") or
cp.get("{urn:mpeg:cenc:2013}default_KID") or
cp.get("cenc:default_KID")
)
if default_kid:
return default_kid
# Method 2: Parse PSSH box
for cp in cp_elements:
# Try standard cenc:pssh
pssh_elem = cp.find("cenc:pssh", self.CENC_NAMESPACE)
if pssh_elem is None:
# Try without namespace
pssh_elem = cp.find("pssh")
if pssh_elem is not None and pssh_elem.text:
try:
kid = self._extract_kid_from_pssh(pssh_elem.text.strip())
if kid:
logger.debug("Extracted KID from PSSH box")
return kid
except Exception as e:
logger.debug(f"Failed to parse PSSH: {e}")
# Method 3: Check Representation-level (fallback)
rep = adaptation_set.find("mpd:Representation", self.MPD_NAMESPACE)
if rep is not None:
rep_cp = rep.findall("mpd:ContentProtection", self.MPD_NAMESPACE)
if rep_cp:
for cp in rep_cp:
default_kid = (
cp.get("default_KID") or
cp.get("{urn:mpeg:cenc:2013}default_KID")
)
if default_kid:
logger.debug("Found KID at Representation level")
return default_kid
return None
@staticmethod
def _extract_kid_from_pssh(self, pssh_b64: str) -> Optional[str]:
"""
Extract first KID from PSSH box (CENC specification).
PSSH structure (version 1):
- box_size: 4 bytes
- box_type: 4 bytes ('pssh')
- version: 1 byte (0 or 1)
- flags: 3 bytes
- system_id: 16 bytes
- [version 1 only] kid_count: 4 bytes
- [version 1 only] kids: 16 bytes each
- data_size: 4 bytes
- data: variable
"""
try:
pssh_data = base64.b64decode(pssh_b64)
if len(pssh_data) < 32:
return None
# Check version (byte 8)
version = pssh_data[8]
if version == 1:
# Version 1 includes KID list
if len(pssh_data) < 36:
return None
# KID count at bytes 28-31 (big-endian)
kid_count = int.from_bytes(pssh_data[28:32], 'big')
if kid_count > 0 and len(pssh_data) >= 48:
# First KID starts at byte 32 (16 bytes)
kid_bytes = pssh_data[32:48]
# Format as UUID string with hyphens
kid_hex = kid_bytes.hex()
kid_uuid = f"{kid_hex[0:8]}-{kid_hex[8:12]}-{kid_hex[12:16]}-{kid_hex[16:20]}-{kid_hex[20:32]}"
return kid_uuid
return None
except Exception as e:
logger.debug(f"Error extracting KID from PSSH: {e}")
return None
def _remove_adaptationsets_without_keys(
self,
root: ET.Element,
as_id_to_kid: Dict[str, str]
):
"""
Remove encrypted AdaptationSets for which we don't have decryption keys.
Optimized to avoid repeated getparent() calls.
"""
if self.key_config.single_key_mode:
# In single key mode, we can decrypt everything
return
removal_count = 0
# Process each period separately to avoid expensive getparent() calls
for period in root.findall(".//mpd:Period", self.MPD_NAMESPACE):
period_id = period.get("id", "")
adaptationsets_to_remove = []
for adaptation_set in period.findall("mpd:AdaptationSet", self.MPD_NAMESPACE):
as_id = adaptation_set.get("id")
if not as_id:
as_id = str(id(adaptation_set))
unique_id = f"{period_id}_{as_id}" if period_id else as_id
# Check if this AdaptationSet requires a key we don't have
if unique_id in as_id_to_kid:
required_kid = as_id_to_kid[unique_id]
if required_kid not in self.key_config.keys:
logger.warning(
f"Removing AdaptationSet {unique_id} - "
f"missing key for KID: {required_kid[:8]}..."
)
adaptationsets_to_remove.append(adaptation_set)
# Remove all marked AdaptationSets from this period
for adaptation_set in adaptationsets_to_remove:
period.remove(adaptation_set)
removal_count += 1
if removal_count > 0:
logger.info(f"Removed {removal_count} AdaptationSet(s) due to missing keys")
def _remove_all_encrypted_adaptationsets(self, root: ET.Element):
"""Remove all encrypted AdaptationSets when we have no keys."""
removal_count = 0
for period in root.findall(".//mpd:Period", self.MPD_NAMESPACE):
adaptationsets_to_remove = []
for adaptation_set in period.findall("mpd:AdaptationSet", self.MPD_NAMESPACE):
# Check if AdaptationSet has ContentProtection
cp_elements = adaptation_set.findall("mpd:ContentProtection", self.MPD_NAMESPACE)
if cp_elements:
adaptationsets_to_remove.append(adaptation_set)
# Remove all encrypted AdaptationSets from this period
for adaptation_set in adaptationsets_to_remove:
period.remove(adaptation_set)
removal_count += 1
if removal_count > 0:
logger.info(f"Removed {removal_count} encrypted AdaptationSet(s) (no keys available)")
def _rewrite_node(
self,
element: ET.Element,
base_url: str,
encrypted_ids: Set[str],
current_encrypted: bool,
self,
element: ET.Element,
base_url: str,
encrypted_ids: Set[str],
as_id_to_kid: Dict[str, str],
current_encrypted: bool,
current_kid: Optional[str] = None,
current_period_id: str = "",
):
"""Recursive node rewriter using state-passing for encryption context."""
# Update state: if we enter an AdaptationSet, check its encryption status
"""Recursive node rewriter with KID-aware key selection."""
# Track period ID as we traverse
if element.tag.endswith("Period"):
current_period_id = element.get("id", "")
# Update state when entering an AdaptationSet
if element.tag.endswith("AdaptationSet"):
as_id = element.get("id", str(id(element)))
current_encrypted = str(as_id) in encrypted_ids
# Use same unique ID logic as _prepare_tree_and_extract_kids
unique_id = f"{current_period_id}_{as_id}" if current_period_id else as_id
current_encrypted = unique_id in encrypted_ids
# Mapping of attributes to segment types for build_proxy_url
# Get specific KID for this AdaptationSet (multi-key mode only)
if current_encrypted and not self.key_config.single_key_mode:
current_kid = as_id_to_kid.get(unique_id)
# Rewrite URL attributes
attr_map = {
"media": "media",
"initialization": "initialization",
@@ -153,14 +434,14 @@ class MPDRewriter:
if "$" in resolved:
path, pattern = self.split_template_url(resolved)
element.attrib[attr] = self.build_proxy_url(
path, pattern, seg_type, current_encrypted
path, pattern, seg_type, current_encrypted, current_kid
)
else:
element.attrib[attr] = self.build_proxy_url(
resolved, None, seg_type, current_encrypted
resolved, None, seg_type, current_encrypted, current_kid
)
# Handle SegmentURL specifically (always 'media' type)
# Handle SegmentURL (always 'media' type)
if element.tag.endswith("SegmentURL") and "media" in element.attrib:
resolved = urljoin(base_url, element.attrib["media"])
path, pattern = (
@@ -169,12 +450,15 @@ class MPDRewriter:
else (resolved, None)
)
element.attrib["media"] = self.build_proxy_url(
path, pattern, "media", current_encrypted
path, pattern, "media", current_encrypted, current_kid
)
# Recurse to children, passing the current encryption state down
# Recurse to children
for child in element:
self._rewrite_node(child, base_url, encrypted_ids, current_encrypted)
self._rewrite_node(
child, base_url, encrypted_ids, as_id_to_kid,
current_encrypted, current_kid, current_period_id
)
def _extract_base_url(self, root: ET.Element, manifest_url: str) -> str:
base_url_elem = root.find(".//mpd:BaseURL", self.MPD_NAMESPACE)
@@ -236,4 +520,4 @@ class MPDRewriter:
int(d["hours"] or 0) * 3600
+ int(d["minutes"] or 0) * 60
+ float(d["seconds"] or 0)
)
)