diff --git a/lib/streaming_providers/base/drm_operations.py b/lib/streaming_providers/base/drm_operations.py index 9da7035..7d4f538 100644 --- a/lib/streaming_providers/base/drm_operations.py +++ b/lib/streaming_providers/base/drm_operations.py @@ -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: diff --git a/lib/streaming_providers/base/utils/mpd_rewriter.py b/lib/streaming_providers/base/utils/mpd_rewriter.py index 1fe1a2d..98822c4 100644 --- a/lib/streaming_providers/base/utils/mpd_rewriter.py +++ b/lib/streaming_providers/base/utils/mpd_rewriter.py @@ -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(" 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) - ) + ) \ No newline at end of file