diff --git a/lib/streaming_providers/base/utils/__init__.py b/lib/streaming_providers/base/utils/__init__.py index bf75f18..1d616ee 100644 --- a/lib/streaming_providers/base/utils/__init__.py +++ b/lib/streaming_providers/base/utils/__init__.py @@ -7,6 +7,9 @@ from .mpd_rewriter import MPDRewriter from .timestamp_converter import TimestampConverter from .mp4_pssh_extractor import MP4PSSHExtractor from .vfs import VFS +from .drm_extractor import DRMExtractor +from .url_resolver import URLResolver +from .manifest_utils import ManifestUtils __all__ = [ "logger", @@ -17,4 +20,7 @@ __all__ = [ "MPDCacheManager", "MP4PSSHExtractor", "TimestampConverter", + "DRMExtractor", + "URLResolver", + "ManifestUtils", ] diff --git a/lib/streaming_providers/base/utils/drm_extractor.py b/lib/streaming_providers/base/utils/drm_extractor.py new file mode 100644 index 0000000..1ece0ba --- /dev/null +++ b/lib/streaming_providers/base/utils/drm_extractor.py @@ -0,0 +1,232 @@ +# streaming_providers/base/utils/drm_extractor.py +""" +DRM-specific extraction utilities for PSSH boxes and key IDs from manifests and segments. +Separated from general manifest parsing to maintain clear separation of concerns. +""" + +import base64 +import re +from typing import List + +from ..models.drm import PSSHData +from .logger import logger + + +class DRMExtractor: + """Extracts PSSH boxes and DRM information from manifests and segments.""" + + @staticmethod + def extract_pssh_from_manifest( + manifest_content: str, + manifest_url: str = "", + fallback_to_segments: bool = True, + segment_urls: List[str] = None, + ) -> List[PSSHData]: + """ + Extract PSSH data from DASH manifest content. + + DEPRECATED: Use extract_single_init_segment_url from ManifestParser + and then extract from that segment instead. + Kept for backwards compatibility. + + Args: + manifest_content: Full manifest XML content + manifest_url: URL of the manifest (unused, kept for compatibility) + fallback_to_segments: Whether to extract from segments if manifest incomplete + segment_urls: List of segment URLs to try if fallback enabled + + Returns: + List of PSSHData objects found + """ + pssh_list = DRMExtractor._extract_from_manifest_content(manifest_content) + + if fallback_to_segments and segment_urls: + incomplete_pssh = [p for p in pssh_list if not p.pssh_box or not p.key_ids] + if incomplete_pssh: + segment_pssh = DRMExtractor._extract_from_single_segment( + segment_urls[0], [p.system_id for p in incomplete_pssh] + ) + return DRMExtractor._merge_pssh_data(pssh_list, segment_pssh) + + return pssh_list + + @staticmethod + def _extract_from_manifest_content(manifest_content: str) -> List[PSSHData]: + """Extract PSSH and DRM systems from manifest content.""" + # Try regex extraction first + pssh_list = DRMExtractor._extract_with_regex(manifest_content) + if pssh_list: + return pssh_list + + # Fallback: extract DRM systems from schemeIdUri only + drm_systems_found = set() + result = [] + + cp_pattern = re.compile( + r']*schemeIdUri="urn:uuid:([^"]+)"[^>]*>', + re.IGNORECASE, + ) + + for match in cp_pattern.finditer(manifest_content): + system_id = match.group(1).lower() + + # Skip mp4protection + if "mp4protection" in manifest_content[max(0, match.start() - 100):match.start()]: + continue + + # Let PSSHData handle normalization + pssh_data = PSSHData( + system_id=system_id, + pssh_box="", + key_ids=[], + source="manifest_scheme_only", + ) + + if pssh_data.drm_system and system_id not in drm_systems_found: + drm_systems_found.add(pssh_data.system_id) # Use normalized ID + result.append(pssh_data) + logger.debug(f"Found DRM system: {pssh_data.drm_system.value}") + + return result + + @staticmethod + def _extract_from_single_segment( + segment_url: str, + expected_system_ids: List[str] = None + ) -> List[PSSHData]: + """Extract PSSH from a single segment URL.""" + from .mp4_pssh_extractor import MP4PSSHExtractor + + try: + pssh_from_segment = MP4PSSHExtractor.extract_from_url(segment_url) + + if expected_system_ids: + # Normalize expected IDs using the model + normalized_expected = [] + for sys_id in expected_system_ids: + # Create temporary PSSHData to leverage its normalization + temp = PSSHData(system_id=sys_id, source="filter") + normalized_expected.append(temp.system_id) + + filtered_pssh = [ + p for p in pssh_from_segment + if p.system_id in normalized_expected + ] + + if filtered_pssh: + return filtered_pssh + else: + # If filtering produced no matches, return all segment PSSH + logger.debug( + f"Filtering by expected_system_ids produced no matches, " + f"returning all {len(pssh_from_segment)} PSSH from segment" + ) + return pssh_from_segment + + # No filtering requested, return all segment PSSH + return pssh_from_segment + + except Exception as e: + logger.warning(f"Failed to extract PSSH from segment: {e}") + + return [] + + @staticmethod + def _merge_pssh_data( + manifest_pssh: List[PSSHData], + segment_pssh: List[PSSHData] + ) -> List[PSSHData]: + """ + Merge manifest and segment PSSH data. + Prefer segment data as it's typically more complete. + """ + if not manifest_pssh: + return segment_pssh + if not segment_pssh: + return manifest_pssh + + merged = [] + segment_by_system = {p.system_id: p for p in segment_pssh} + + for manifest_p in manifest_pssh: + if manifest_p.system_id in segment_by_system: + # Use segment data (complete) + merged.append(segment_by_system[manifest_p.system_id]) + else: + # Keep manifest data (incomplete) + merged.append(manifest_p) + + return merged + + @staticmethod + def _extract_with_regex(mpd_content: str) -> List[PSSHData]: + """Extract PSSH boxes using regex.""" + pssh_dict = {} + global_key_ids = [] + + # Compile patterns + pssh_pattern = re.compile(r"<(?:cenc:)?pssh[^>]*>([^<]+)") + default_kid_pattern = re.compile( + r'(?:cenc:)?default_KID="([^"]+)"', re.IGNORECASE + ) + system_id_pattern = re.compile(r'schemeIdUri="urn:uuid:([^"]+)"', re.IGNORECASE) + + # Find ContentProtection blocks + cp_blocks = re.findall( + r"]*>.*?", + mpd_content, + re.DOTALL + ) + + # First pass: collect all default KIDs + for block in cp_blocks: + kid_match = default_kid_pattern.search(block) + if kid_match: + clean_kid = kid_match.group(1).replace("-", "").lower() + if clean_kid not in global_key_ids: + global_key_ids.append(clean_kid) + + # Second pass: extract PSSH data + for block in cp_blocks: + try: + system_id = None + + # Extract system ID from schemeIdUri + scheme_match = system_id_pattern.search(block) + if scheme_match: + system_id = scheme_match.group(1).lower() + + # Extract PSSH data + for pssh_match in pssh_pattern.finditer(block): + pssh_b64 = pssh_match.group(1) + try: + pssh_data = base64.b64decode(pssh_b64) + + if len(pssh_data) >= 28: + # Extract system ID from PSSH if not found + if not system_id: + system_id_bytes = pssh_data[12:28] + system_id = "-".join([ + system_id_bytes[0:4].hex(), + system_id_bytes[4:6].hex(), + system_id_bytes[6:8].hex(), + system_id_bytes[8:10].hex(), + system_id_bytes[10:16].hex(), + ]) + + # Deduplicate by PSSH box content + if pssh_b64 not in pssh_dict: + pssh_dict[pssh_b64] = PSSHData( + system_id=system_id, + pssh_box=pssh_b64, + key_ids=global_key_ids.copy(), + source="manifest_pssh", + ) + + except Exception as e: + logger.debug(f"Error decoding PSSH: {e}") + + except Exception as e: + logger.debug(f"Error processing ContentProtection block: {e}") + + return list(pssh_dict.values()) \ No newline at end of file diff --git a/lib/streaming_providers/base/utils/manifest_parser.py b/lib/streaming_providers/base/utils/manifest_parser.py index f6f4000..aed0cc2 100644 --- a/lib/streaming_providers/base/utils/manifest_parser.py +++ b/lib/streaming_providers/base/utils/manifest_parser.py @@ -1,300 +1,83 @@ -import base64 -import re -from typing import List, Optional -from urllib.parse import quote, urljoin, urlparse +# streaming_providers/base/utils/manifest_parser.py +""" +DASH manifest parser for extracting init segment URLs. +For PSSH/DRM extraction, use drm_extractor module. +""" + +from typing import Optional -from ..models.drm import PSSHData from .logger import logger +from .url_resolver import URLResolver +from .manifest_utils import ManifestUtils class ManifestParser: - @staticmethod - def extract_pssh_from_manifest( - manifest_content: str, - manifest_url: str = "", - fallback_to_segments: bool = True, - segment_urls: List[str] = None, - ) -> List[PSSHData]: - """ - DEPRECATED: Use extract_single_init_segment_url instead. - Kept for backwards compatibility. - """ - pssh_list = ManifestParser._extract_from_manifest_content(manifest_content) - - if fallback_to_segments and segment_urls: - incomplete_pssh = [p for p in pssh_list if not p.pssh_box or not p.key_ids] - if incomplete_pssh: - segment_pssh = ManifestParser._extract_from_single_segment( - segment_urls[0], [p.system_id for p in incomplete_pssh] - ) - return ManifestParser._merge_pssh_data(pssh_list, segment_pssh) - - return pssh_list - - @staticmethod - def _extract_from_manifest_content(manifest_content: str) -> List[PSSHData]: - """Extract PSSH and DRM systems from manifest content""" - # Try regex extraction first - pssh_list = ManifestParser._extract_with_regex(manifest_content) - if pssh_list: - return pssh_list - - # Fallback: extract DRM systems from schemeIdUri only - drm_systems_found = set() - result = [] - - cp_pattern = re.compile( - r']*schemeIdUri="urn:uuid:([^"]+)"[^>]*>', - re.IGNORECASE, - ) - - for match in cp_pattern.finditer(manifest_content): - system_id = match.group(1).lower() - - # Skip mp4protection - if "mp4protection" in manifest_content[max(0, match.start() - 100):match.start()]: - continue - - # Let PSSHData handle normalization! - pssh_data = PSSHData( - system_id=system_id, - pssh_box="", - key_ids=[], - source="manifest_scheme_only", - ) - - if pssh_data.drm_system and system_id not in drm_systems_found: - drm_systems_found.add(pssh_data.system_id) # Use normalized ID - result.append(pssh_data) - logger.debug(f"Found DRM system: {pssh_data.drm_system.value}") - - return result - - @staticmethod - def _extract_from_single_segment( - segment_url: str, expected_system_ids: List[str] = None - ) -> List[PSSHData]: - from .mp4_pssh_extractor import MP4PSSHExtractor - - try: - pssh_from_segment = MP4PSSHExtractor.extract_from_url(segment_url) - - if expected_system_ids: - # Normalize expected IDs using the model! - normalized_expected = [] - for sys_id in expected_system_ids: - # Create temporary PSSHData to leverage its normalization - temp = PSSHData(system_id=sys_id, source="filter") - normalized_expected.append(temp.system_id) - - filtered_pssh = [ - p for p in pssh_from_segment - if p.system_id in normalized_expected - ] - if filtered_pssh: - return filtered_pssh - else: - # If filtering produced no matches, return all segment PSSH - # This can happen if the manifest had incomplete system IDs - logger.debug( - f"Filtering by expected_system_ids produced no matches, " - f"returning all {len(pssh_from_segment)} PSSH from segment" - ) - return pssh_from_segment - - # No filtering requested, return all segment PSSH - return pssh_from_segment - - except Exception as e: - logger.warning(f"Failed to extract PSSH from segment: {e}") - - return [] - - @staticmethod - def _merge_pssh_data( - manifest_pssh: List[PSSHData], segment_pssh: List[PSSHData] - ) -> List[PSSHData]: - """Merge manifest and segment PSSH data""" - if not manifest_pssh: - return segment_pssh - if not segment_pssh: - return manifest_pssh - - merged = [] - segment_by_system = {p.system_id: p for p in segment_pssh} - - for manifest_p in manifest_pssh: - if manifest_p.system_id in segment_by_system: - # Use segment data (complete) - merged.append(segment_by_system[manifest_p.system_id]) - else: - # Keep manifest data (incomplete) - merged.append(manifest_p) - - return merged - - @staticmethod - def _extract_with_regex(mpd_content: str) -> List[PSSHData]: - """Extract PSSH boxes using regex""" - pssh_dict = {} - global_key_ids = [] - - # Compile patterns once - pssh_pattern = re.compile(r"<(?:cenc:)?pssh[^>]*>([^<]+)") - default_kid_pattern = re.compile( - r'(?:cenc:)?default_KID="([^"]+)"', re.IGNORECASE - ) - system_id_pattern = re.compile(r'schemeIdUri="urn:uuid:([^"]+)"', re.IGNORECASE) - - # Find ContentProtection blocks efficiently - cp_blocks = re.findall( - r"]*>.*?", mpd_content, re.DOTALL - ) - - # First pass: collect all default KIDs - for block in cp_blocks: - kid_match = default_kid_pattern.search(block) - if kid_match: - clean_kid = kid_match.group(1).replace("-", "").lower() - if clean_kid not in global_key_ids: - global_key_ids.append(clean_kid) - - # Second pass: extract PSSH data - for block in cp_blocks: - try: - system_id = None - - # Extract system ID from schemeIdUri - scheme_match = system_id_pattern.search(block) - if scheme_match: - system_id = scheme_match.group(1).lower() - - # Extract PSSH data - for pssh_match in pssh_pattern.finditer(block): - pssh_b64 = pssh_match.group(1) - try: - pssh_data = base64.b64decode(pssh_b64) - - if len(pssh_data) >= 28: - # Extract system ID from PSSH if not found - if not system_id: - system_id_bytes = pssh_data[12:28] - system_id = "-".join( - [ - system_id_bytes[0:4].hex(), - system_id_bytes[4:6].hex(), - system_id_bytes[6:8].hex(), - system_id_bytes[8:10].hex(), - system_id_bytes[10:16].hex(), - ] - ) - - # Deduplicate by PSSH box content - if pssh_b64 not in pssh_dict: - pssh_dict[pssh_b64] = PSSHData( - system_id=system_id, - pssh_box=pssh_b64, - key_ids=global_key_ids.copy(), - source="manifest_pssh", - ) - - except Exception as e: - logger.debug(f"Error decoding PSSH: {e}") - - except Exception as e: - logger.debug(f"Error processing ContentProtection block: {e}") - - return list(pssh_dict.values()) + """Parser for DASH manifests focused on segment URL extraction.""" @staticmethod def extract_single_init_segment_url( - manifest_content: str, manifest_url: str + manifest_content: str, + manifest_url: str ) -> Optional[str]: """ Extract ONE init segment URL from DASH manifest. Prioritizes video representations as they typically have the same DRM as audio. + + Args: + manifest_content: Full manifest XML content + manifest_url: URL where the manifest was fetched from + + Returns: + Full URL to an initialization segment, or None if not found """ - # Parse manifest base URL - parsed = urlparse(manifest_url) - manifest_base = ( - f"{parsed.scheme}://{parsed.netloc}{'/'.join(parsed.path.split('/')[:-1])}" - ) - if not manifest_base.endswith("/"): - manifest_base += "/" - - # Extract BaseURL elements (can appear at multiple levels) - base_urls = re.findall(r"]*>([^<]+)", manifest_content) - - # Build effective base URL - effective_base = manifest_base - for base_url in base_urls: - if base_url.startswith("http"): - effective_base = base_url - else: - effective_base = urljoin(effective_base, base_url) - - if not effective_base.endswith("/"): - effective_base += "/" + # Build effective base URL from manifest URL and BaseURL elements + base_urls = ManifestUtils.extract_base_urls(manifest_content) + effective_base = URLResolver.build_effective_base_url(manifest_url, base_urls) logger.debug(f"Effective base URL: {effective_base}") - # Find SegmentTemplate with initialization attribute - # Prioritize video AdaptationSets - adaptation_sets = re.findall( - r"]*>.*?", manifest_content, re.DOTALL - ) - - video_sets = [] - audio_sets = [] - - for ad_set in adaptation_sets: - if 'contentType="video"' in ad_set or 'mimeType="video/' in ad_set: - video_sets.append(ad_set) - elif 'contentType="audio"' in ad_set or 'mimeType="audio/' in ad_set: - audio_sets.append(ad_set) + # Parse all adaptation sets + adaptation_sets = ManifestUtils.parse_adaptation_sets(manifest_content) + video_sets, audio_sets = ManifestUtils.separate_video_audio_sets(adaptation_sets) # Try video first, then audio target_sets = video_sets + audio_sets - for ad_set in target_sets: - # Find SegmentTemplate initialization - seg_template_match = re.search( - r']*initialization="([^"]+)"', ad_set, re.IGNORECASE + for ad_set_info in target_sets: + # Extract SegmentTemplate initialization attribute + init_template = ManifestUtils.extract_segment_template_initialization( + ad_set_info.content ) - if not seg_template_match: + if not init_template: continue - init_template = seg_template_match.group(1) logger.debug(f"Found init template: {init_template}") - # Find first Representation in this AdaptationSet - rep_match = re.search(r']*id="([^"]+)"', ad_set) - if not rep_match: + # Get first Representation ID from this AdaptationSet + rep_id = ManifestUtils.extract_first_representation_id(ad_set_info.content) + + if not rep_id: + logger.debug("No Representation ID found in AdaptationSet") continue - rep_id = rep_match.group(1) logger.debug(f"Using Representation ID: {rep_id}") - # Substitute template variables - init_url = init_template.replace("$RepresentationID$", rep_id) - - # Handle other common template variables - init_url = init_url.replace("$Bandwidth$", "0") - init_url = init_url.replace("$Time$", "0") - init_url = init_url.replace("$Number$", "1") + # Substitute template variables with defaults + init_url = URLResolver.substitute_template_variables( + init_template, + representation_id=rep_id, + bandwidth="0", + time="0", + number="1" + ) # Construct full URL - if init_url.startswith("http"): - full_url = init_url - else: - # URL encode special characters in representation ID - # Split path and encode only the filename part - path_parts = init_url.split("/") - path_parts[-1] = quote(path_parts[-1], safe=".-_") - init_url = "/".join(path_parts) - - full_url = urljoin(effective_base, init_url) + full_url = URLResolver.construct_full_url( + effective_base, + init_url, + url_encode_filename=True + ) logger.info(f"Constructed init segment URL: {full_url}") return full_url @@ -303,10 +86,12 @@ class ManifestParser: return None @staticmethod - def extract_segment_urls(manifest_content: str, manifest_url: str) -> List[str]: + def extract_segment_urls(manifest_content: str, manifest_url: str) -> list[str]: """ DEPRECATED: Use extract_single_init_segment_url instead. This extracts ALL segments which is inefficient. + + This method is kept for backwards compatibility only. """ logger.warning( "extract_segment_urls is deprecated, use extract_single_init_segment_url" @@ -317,11 +102,11 @@ class ManifestParser: return [init_url] if init_url else [] @staticmethod - def extract_init_segment_urls( - manifest_content: str, manifest_url: str - ) -> List[str]: + def extract_init_segment_urls(manifest_content: str, manifest_url: str) -> list[str]: """ DEPRECATED: Use extract_single_init_segment_url instead. + + This method is kept for backwards compatibility only. """ logger.warning( "extract_init_segment_urls is deprecated, use extract_single_init_segment_url" diff --git a/lib/streaming_providers/base/utils/manifest_utils.py b/lib/streaming_providers/base/utils/manifest_utils.py new file mode 100644 index 0000000..b13a09a --- /dev/null +++ b/lib/streaming_providers/base/utils/manifest_utils.py @@ -0,0 +1,142 @@ +# streaming_providers/base/utils/manifest_utils.py +""" +Utilities for parsing DASH manifest structure. +Extracts AdaptationSets, Representations, and other manifest elements. +""" + +import re +from typing import List, Tuple, Optional +from dataclasses import dataclass + +@dataclass +class AdaptationSetInfo: + """Information about a parsed AdaptationSet.""" + content: str # Raw XML content + content_type: str # "video", "audio", or "unknown" + mime_type: str + is_video: bool + is_audio: bool + + +class ManifestUtils: + """Utilities for parsing DASH manifest structure.""" + + @staticmethod + def parse_adaptation_sets(manifest_content: str) -> List[AdaptationSetInfo]: + """ + Parse all AdaptationSets from manifest content. + + Args: + manifest_content: Full manifest XML content + + Returns: + List of AdaptationSetInfo objects + """ + adaptation_sets = [] + + # Find all AdaptationSet blocks + ad_set_pattern = re.compile( + r"]*>.*?", + re.DOTALL + ) + + for match in ad_set_pattern.finditer(manifest_content): + ad_set_content = match.group(0) + + # Extract content type and mime type + content_type = ManifestUtils._extract_content_type(ad_set_content) + mime_type = ManifestUtils._extract_mime_type(ad_set_content) + + is_video = content_type == "video" or mime_type.startswith("video/") + is_audio = content_type == "audio" or mime_type.startswith("audio/") + + adaptation_sets.append(AdaptationSetInfo( + content=ad_set_content, + content_type=content_type, + mime_type=mime_type, + is_video=is_video, + is_audio=is_audio + )) + + return adaptation_sets + + @staticmethod + def separate_video_audio_sets( + adaptation_sets: List[AdaptationSetInfo] + ) -> Tuple[List[AdaptationSetInfo], List[AdaptationSetInfo]]: + """ + Separate adaptation sets into video and audio lists. + + Args: + adaptation_sets: List of parsed AdaptationSets + + Returns: + Tuple of (video_sets, audio_sets) + """ + video_sets = [ad_set for ad_set in adaptation_sets if ad_set.is_video] + audio_sets = [ad_set for ad_set in adaptation_sets if ad_set.is_audio] + + return video_sets, audio_sets + + @staticmethod + def _extract_content_type(ad_set_content: str) -> str: + """Extract contentType attribute from AdaptationSet.""" + match = re.search(r'contentType="([^"]+)"', ad_set_content) + return match.group(1) if match else "unknown" + + @staticmethod + def _extract_mime_type(ad_set_content: str) -> str: + """Extract mimeType attribute from AdaptationSet or Representation.""" + # Try AdaptationSet level first + match = re.search(r']*mimeType="([^"]+)"', ad_set_content) + if match: + return match.group(1) + + # Try Representation level + match = re.search(r']*mimeType="([^"]+)"', ad_set_content) + return match.group(1) if match else "" + + @staticmethod + def extract_first_representation_id(ad_set_content: str) -> Optional[str]: + """ + Extract the ID of the first Representation in an AdaptationSet. + + Args: + ad_set_content: AdaptationSet XML content + + Returns: + Representation ID or None if not found + """ + match = re.search(r']*id="([^"]+)"', ad_set_content) + return match.group(1) if match else None + + @staticmethod + def extract_segment_template_initialization(ad_set_content: str) -> Optional[str]: + """ + Extract initialization attribute from SegmentTemplate. + + Args: + ad_set_content: AdaptationSet XML content + + Returns: + Initialization template string or None if not found + """ + match = re.search( + r']*initialization="([^"]+)"', + ad_set_content, + re.IGNORECASE + ) + return match.group(1) if match else None + + @staticmethod + def extract_base_urls(manifest_content: str) -> List[str]: + """ + Extract all BaseURL elements from manifest. + + Args: + manifest_content: Full manifest XML content + + Returns: + List of BaseURL text contents + """ + return re.findall(r"]*>([^<]+)", manifest_content) \ No newline at end of file diff --git a/lib/streaming_providers/base/utils/mpd_rewriter.py b/lib/streaming_providers/base/utils/mpd_rewriter.py index e58838f..b8e9c0c 100644 --- a/lib/streaming_providers/base/utils/mpd_rewriter.py +++ b/lib/streaming_providers/base/utils/mpd_rewriter.py @@ -1,16 +1,22 @@ # streaming_providers/base/utils/mpd_rewriter.py +""" +MPD rewriter for DASH manifests. +Handles URL proxying, DRM key injection, quality filtering, and representation blocklisting. +""" + import base64 import struct import xml.etree.ElementTree as ET import re from typing import Optional, Tuple, Set, Dict, List -from urllib.parse import urljoin, urlparse, quote, urlencode +from urllib.parse import urljoin, quote, urlencode from datetime import datetime, timezone from email.utils import parsedate_to_datetime from dataclasses import dataclass, field from .logger import logger from .vfs import get_vfs +from .url_resolver import URLResolver # Pre-compile regex for ISO duration parsing at module level ISO_8601_PERIOD_RE = re.compile( @@ -174,17 +180,17 @@ class MPDRewriter: media_proxy_url: str, provider_proxy_url: Optional[str] = None, clearkey_keyids: Optional[dict] = None, - highest_quality_video_only: bool = False, # NEW: Enable highest quality filtering - provider: Optional[str] = None, # NEW: Provider name for blocklist filtering - channel: Optional[str] = None, # NEW: Channel name for blocklist filtering - blocklist_path: str = "representation_blocklist.json", # NEW: Path to blocklist file + highest_quality_video_only: bool = False, + provider: Optional[str] = None, + channel: Optional[str] = None, + blocklist_path: str = "representation_blocklist.json", ): self.media_proxy_url = media_proxy_url.rstrip("/") self.provider_proxy_url = provider_proxy_url self.key_config = KeyConfiguration(clearkey_keyids or {}) self.highest_quality_video_only = highest_quality_video_only - # NEW: Blocklist configuration + # Blocklist configuration self.provider = provider self.channel = channel self.blocklist = RepresentationBlocklist(blocklist_path) @@ -211,8 +217,8 @@ class MPDRewriter: template_pattern: Optional[str] = None, segment_type: Optional[str] = None, is_encrypted: bool = False, - kid: Optional[str] = None, # Specific KID for this AdaptationSet - representation_id: Optional[str] = None, # NEW: For template substitution + kid: Optional[str] = None, + representation_id: Optional[str] = None, ) -> str: params = {"url": original_url, **self._static_params} @@ -238,7 +244,7 @@ class MPDRewriter: params["kid"] = kid params["key"] = key else: - # Fallback to first key (should only happen if we couldn't extract KID) + # Fallback to first key if segment_type == "initialization": params["kid"] = self.key_config.default_kid elif segment_type == "media": @@ -253,7 +259,7 @@ class MPDRewriter: proxy_url = f"{self.media_proxy_url}/api/{endpoint}/{encoded}" if template_pattern: - # NEW: Substitute $RepresentationID$ if we have it and highest_quality_video_only is enabled + # Substitute $RepresentationID$ if we have it and highest_quality_video_only is enabled if self.highest_quality_video_only and representation_id and "$RepresentationID$" in template_pattern: template_pattern = template_pattern.replace("$RepresentationID$", representation_id) @@ -261,29 +267,18 @@ class MPDRewriter: return proxy_url - @staticmethod - def split_template_url(url: str) -> Tuple[str, Optional[str]]: - if "$" not in url: - return url, None - first_template_pos = url.find("$") - 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:] - def rewrite_mpd(self, mpd_content: str, manifest_url: str) -> str: try: root = ET.fromstring(mpd_content) ET.register_namespace("", self.MPD_NAMESPACE["mpd"]) - # Extract MPD-level base URL before any modifications + # Extract MPD-level base URL using shared utility mpd_base_url = self._extract_mpd_base_url(root, manifest_url) # Single-pass tree preparation with BaseURL extraction encrypted_ids, as_id_to_kid, base_url_map = self._prepare_tree_and_extract_kids(root, mpd_base_url) # FIRST: Filter out encrypted AdaptationSets without available keys - # This ensures we only consider decryptable content for quality selection if self.key_config.keys: self._remove_adaptationsets_without_keys(root, as_id_to_kid) else: @@ -548,131 +543,102 @@ class MPDRewriter: return encrypted_ids, as_id_to_kid, base_url_map def _extract_kid_from_adaptationset(self, adaptation_set: ET.Element) -> Optional[str]: - """Extract KID from ContentProtection elements.""" - # Check AdaptationSet-level ContentProtection + """Extract KID from ContentProtection elements in an AdaptationSet.""" + # Try cenc:default_KID first for cp in adaptation_set.findall("mpd:ContentProtection", self.MPD_NAMESPACE): - kid = self._extract_kid_from_cp_element(cp) - if kid: - return kid + default_kid = cp.get("{urn:mpeg:cenc:2013}default_KID") + if default_kid: + return default_kid.replace("-", "").lower() - # Check Representation-level ContentProtection - for representation in adaptation_set.findall("mpd:Representation", self.MPD_NAMESPACE): - for cp in representation.findall("mpd:ContentProtection", self.MPD_NAMESPACE): - kid = self._extract_kid_from_cp_element(cp) - if kid: - return kid - - return None - - def _extract_kid_from_cp_element(self, cp: ET.Element) -> Optional[str]: - """Extract KID from a single ContentProtection element.""" - # Check default_KID attribute - kid_attr = cp.get("{urn:mpeg:cenc:2013}default_KID") - if kid_attr: - # Normalize immediately! Remove hyphens and lowercase - normalized = kid_attr.replace("-", "").lower() - logger.debug(f"Extracted KID from default_KID attribute: {normalized}") - return normalized - - # Check cenc:pssh - for pssh in cp.findall("cenc:pssh", self.CENC_NAMESPACE): - if pssh.text: + # Try PSSH box + for cp in adaptation_set.findall("mpd:ContentProtection", self.MPD_NAMESPACE): + pssh_elem = cp.find("cenc:pssh", self.CENC_NAMESPACE) + if pssh_elem is not None and pssh_elem.text: try: - pssh_data = base64.b64decode(pssh.text) - logger.debug(f"PSSH box size: {len(pssh_data)} bytes") - - if len(pssh_data) >= 36: - version = pssh_data[8] if len(pssh_data) > 8 else 0 - system_id = pssh_data[12:28].hex() if len(pssh_data) >= 28 else "unknown" - logger.debug(f"PSSH version: {version}, System ID: {system_id}") - - if version > 0: + pssh_data = base64.b64decode(pssh_elem.text) + if len(pssh_data) >= 32: + version = pssh_data[8] + if version == 1: kid_count = struct.unpack(">I", pssh_data[28:32])[0] - logger.debug(f"KID count from header: {kid_count}") - - if kid_count > 0 and len(pssh_data) >= 48: - kid_bytes = pssh_data[32:48] - kid_hex = kid_bytes.hex().lower() - logger.debug(f"Extracted KID from PSSH header: {kid_hex}") - return kid_hex - - # Try to extract from payload if header extraction failed - # This is where we'd use the new Widevine payload parsing + if kid_count > 0 and len(pssh_data) >= 32 + 16: + kid = pssh_data[32:48].hex() + return kid except Exception as e: - logger.debug(f"Error extracting KID from PSSH: {e}") + logger.debug(f"Failed to extract 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 that require keys we don't have.""" - removal_count = 0 - - # Debug: Log what keys we have - logger.debug(f"Available keys: {list(self.key_config.keys.keys())}") + """Remove encrypted AdaptationSets for which we don't have keys.""" + removed_count = 0 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)) - + as_id = adaptation_set.get("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]}..." + kid = as_id_to_kid[unique_id] + if kid not in self.key_config.keys: + logger.info( + f"Removing AdaptationSet {unique_id} - no key available for KID {kid[:8]}..." ) adaptationsets_to_remove.append(adaptation_set) + removed_count += 1 - # 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") + if removed_count > 0: + logger.info(f"Removed {removed_count} AdaptationSet(s) without available keys") def _remove_all_encrypted_adaptationsets(self, root: ET.Element): - """Remove all encrypted AdaptationSets when we have no keys.""" - removal_count = 0 + """Remove all encrypted AdaptationSets when no keys are available.""" + removed_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 + # Check for ContentProtection cp_elements = adaptation_set.findall("mpd:ContentProtection", self.MPD_NAMESPACE) - if cp_elements: - adaptationsets_to_remove.append(adaptation_set) + has_cp = len(cp_elements) > 0 + + # Also check Representation level + if not has_cp: + for representation in adaptation_set.findall("mpd:Representation", self.MPD_NAMESPACE): + rep_cp = representation.findall("mpd:ContentProtection", self.MPD_NAMESPACE) + if rep_cp: + has_cp = True + break + + if has_cp: + as_id = adaptation_set.get("id", "unknown") + logger.info(f"Removing encrypted AdaptationSet {as_id} - no keys available") + adaptationsets_to_remove.append(adaptation_set) + removed_count += 1 - # 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)") + if removed_count > 0: + logger.info(f"Removed {removed_count} encrypted AdaptationSet(s)") def _remove_blocked_representations(self, root: ET.Element): - """ - Remove Representation elements that are blocked for this provider/channel. - If an AdaptationSet has only one Representation and it's blocked, remove the entire AdaptationSet. - """ - if not self.provider or not self.channel: - return - + """Remove representations that are blocklisted for this provider/channel.""" blocked_ids = self.blocklist.get_blocked_ids(self.provider, self.channel) + if not blocked_ids: return + logger.info( + f"Applying representation blocklist for {self.provider}/{self.channel}: " + f"{len(blocked_ids)} representation(s) blocked" + ) + total_reps_removed = 0 total_as_removed = 0 @@ -729,7 +695,7 @@ class MPDRewriter: current_encrypted: bool, current_kid: Optional[str] = None, current_period_id: str = "", - best_video_info: Optional[VideoRepresentation] = None, # NEW + best_video_info: Optional[VideoRepresentation] = None, ): """Recursive node rewriter with KID-aware key selection and context-aware base URLs.""" # Track period ID as we traverse @@ -738,7 +704,7 @@ class MPDRewriter: # Update state when entering an AdaptationSet current_as_id = None - current_rep_id = None # NEW: Track current representation ID + current_rep_id = None if element.tag.endswith("AdaptationSet"): as_id = element.get("id", str(id(element))) current_as_id = as_id @@ -754,7 +720,7 @@ class MPDRewriter: if current_encrypted and not self.key_config.single_key_mode: current_kid = as_id_to_kid.get(unique_id) - # NEW: Track representation ID for template substitution + # Track representation ID for template substitution if element.tag.endswith("Representation"): current_rep_id = element.get("id", "") @@ -773,64 +739,51 @@ class MPDRewriter: resolved = urljoin(base_url, val) if "$" in resolved: - path, pattern = self.split_template_url(resolved) + # Use shared utility for splitting template URLs + path, pattern = URLResolver.split_template_url(resolved) element.attrib[attr] = self.build_proxy_url( path, pattern, seg_type, current_encrypted, current_kid, - representation_id=current_rep_id # NEW + representation_id=current_rep_id ) else: element.attrib[attr] = self.build_proxy_url( resolved, None, seg_type, current_encrypted, current_kid, - representation_id=current_rep_id # NEW + representation_id=current_rep_id ) # Handle SegmentURL (always 'media' type) if element.tag.endswith("SegmentURL") and "media" in element.attrib: resolved = urljoin(base_url, element.attrib["media"]) path, pattern = ( - self.split_template_url(resolved) + URLResolver.split_template_url(resolved) if "$" in resolved else (resolved, None) ) element.attrib["media"] = self.build_proxy_url( path, pattern, "media", current_encrypted, current_kid, - representation_id=current_rep_id # NEW + representation_id=current_rep_id ) # Recurse to children for child in element: self._rewrite_node( child, base_url, encrypted_ids, as_id_to_kid, base_url_map, - current_encrypted, current_kid, current_period_id, best_video_info # NEW + current_encrypted, current_kid, current_period_id, best_video_info ) def _extract_mpd_base_url(self, root: ET.Element, manifest_url: str) -> str: - """Extract and resolve MPD-level BaseURL.""" + """ + Extract and resolve MPD-level BaseURL. + Uses shared URLResolver utility. + """ base_url_elem = root.find("mpd:BaseURL", self.MPD_NAMESPACE) - - # Check if this is one of the special services - SPECIAL_PREFIXES = [ - "https://bpcdnmanprod.nexttv.ht.hr/bpk-tv/", - "https://lineartv-cdn.t-mobile.pl/bpk-tv/" - ] - - # Determine manifest directory based on service type - if any(manifest_url.startswith(prefix) for prefix in SPECIAL_PREFIXES): - # Special service: KEEP index.mpd - manifest_dir = manifest_url if manifest_url.endswith('/') else f"{manifest_url}/" - else: - # Normal service: remove index.mpd - parsed_manifest = urlparse(manifest_url) - manifest_dir = f"{parsed_manifest.scheme}://{parsed_manifest.netloc}{parsed_manifest.path.rsplit('/', 1)[0]}/" + base_url_text = None if base_url_elem is not None and base_url_elem.text: base_url_text = base_url_elem.text.strip() - if not base_url_text.startswith(("http://", "https://")): - return urljoin(manifest_dir, base_url_text) - return base_url_text - # No BaseURL element - return manifest_dir + # Use shared utility - it handles special service prefixes + return URLResolver.resolve_base_url_with_element(manifest_url, base_url_text) @staticmethod def extract_cache_ttl(headers: dict) -> int: diff --git a/lib/streaming_providers/base/utils/url_resolver.py b/lib/streaming_providers/base/utils/url_resolver.py new file mode 100644 index 0000000..fdbfbed --- /dev/null +++ b/lib/streaming_providers/base/utils/url_resolver.py @@ -0,0 +1,204 @@ +# streaming_providers/base/utils/url_resolver.py +""" +Shared URL resolution utilities for manifest parsing and MPD rewriting. +Centralizes URL construction, template substitution, and base URL resolution. +""" + +from typing import Optional, Tuple +from urllib.parse import urljoin, urlparse, quote + + +class URLResolver: + """Handles URL resolution, template substitution, and base URL extraction for DASH manifests.""" + + # Special service prefixes that require different URL handling + SPECIAL_SERVICE_PREFIXES = [ + "https://bpcdnmanprod.nexttv.ht.hr/bpk-tv/", + "https://lineartv-cdn.t-mobile.pl/bpk-tv/" + ] + + @staticmethod + def extract_manifest_base_url(manifest_url: str) -> str: + """ + Extract the base directory URL from a manifest URL. + + Special services keep 'index.mpd' in the path, while normal services remove it. + + Args: + manifest_url: Full URL to the manifest file + + Returns: + Base URL with trailing slash + """ + # Check if this is a special service + is_special_service = any( + manifest_url.startswith(prefix) + for prefix in URLResolver.SPECIAL_SERVICE_PREFIXES + ) + + if is_special_service: + # Special service: KEEP index.mpd in path + manifest_dir = manifest_url if manifest_url.endswith('/') else f"{manifest_url}/" + else: + # Normal service: remove index.mpd from path + parsed = urlparse(manifest_url) + manifest_dir = f"{parsed.scheme}://{parsed.netloc}{parsed.path.rsplit('/', 1)[0]}/" + + return manifest_dir + + @staticmethod + def resolve_base_url_with_element( + manifest_url: str, + base_url_text: Optional[str] = None + ) -> str: + """ + Resolve base URL considering manifest URL and optional BaseURL element text. + + Args: + manifest_url: URL of the manifest + base_url_text: Text content from element, if present + + Returns: + Resolved base URL with trailing slash + """ + manifest_base = URLResolver.extract_manifest_base_url(manifest_url) + + if base_url_text: + base_url_text = base_url_text.strip() + if base_url_text.startswith(("http://", "https://")): + # Absolute URL + return base_url_text if base_url_text.endswith('/') else f"{base_url_text}/" + else: + # Relative URL + resolved = urljoin(manifest_base, base_url_text) + return resolved if resolved.endswith('/') else f"{resolved}/" + + return manifest_base + + @staticmethod + def build_effective_base_url( + manifest_url: str, + base_url_elements: list[str] + ) -> str: + """ + Build effective base URL by chaining BaseURL elements. + Used when multiple BaseURL elements exist at different levels. + + Args: + manifest_url: URL of the manifest + base_url_elements: List of BaseURL text contents in order + + Returns: + Final effective base URL with trailing slash + """ + effective_base = URLResolver.extract_manifest_base_url(manifest_url) + + for base_url in base_url_elements: + if base_url.startswith("http"): + effective_base = base_url + else: + effective_base = urljoin(effective_base, base_url) + + return effective_base if effective_base.endswith('/') else f"{effective_base}/" + + @staticmethod + def substitute_template_variables( + template: str, + representation_id: Optional[str] = None, + bandwidth: Optional[str] = None, + time: Optional[str] = None, + number: Optional[str] = None + ) -> str: + """ + Substitute DASH template variables with actual values. + + Common template variables: + - $RepresentationID$ - Representation identifier + - $Bandwidth$ - Representation bandwidth + - $Time$ - Segment time + - $Number$ - Segment number + + Args: + template: URL template with $Variable$ placeholders + representation_id: Value for $RepresentationID$ + bandwidth: Value for $Bandwidth$ + time: Value for $Time$ + number: Value for $Number$ + + Returns: + Template with variables substituted + """ + result = template + + if representation_id is not None: + result = result.replace("$RepresentationID$", str(representation_id)) + + if bandwidth is not None: + result = result.replace("$Bandwidth$", str(bandwidth)) + + if time is not None: + result = result.replace("$Time$", str(time)) + + if number is not None: + result = result.replace("$Number$", str(number)) + + return result + + @staticmethod + def construct_full_url( + base_url: str, + relative_path: str, + url_encode_filename: bool = False + ) -> str: + """ + Construct full URL from base URL and relative path. + + Args: + base_url: Base URL (should end with /) + relative_path: Relative path to append + url_encode_filename: If True, URL-encode the filename portion + + Returns: + Complete URL + """ + if relative_path.startswith("http"): + # Already absolute + return relative_path + + if url_encode_filename: + # URL encode special characters in filename only + path_parts = relative_path.split("/") + path_parts[-1] = quote(path_parts[-1], safe=".-_") + relative_path = "/".join(path_parts) + + return urljoin(base_url, relative_path) + + @staticmethod + def split_template_url(url: str) -> Tuple[str, Optional[str]]: + """ + Split a URL with template variables into base path and template pattern. + + Example: + "https://cdn.com/path/segment-$Number$.m4s" + -> ("https://cdn.com/path", "segment-$Number$.m4s") + + Args: + url: URL potentially containing template variables ($Variable$) + + Returns: + Tuple of (base_path, template_pattern) + If no template variables found, returns (url, None) + """ + if "$" not in url: + return url, None + + first_template_pos = url.find("$") + last_slash_before_template = url.rfind("/", 0, first_template_pos) + + if last_slash_before_template == -1: + return "", url + + base_path = url[:last_slash_before_template] + template_pattern = url[last_slash_before_template + 1:] + + return base_path, template_pattern \ No newline at end of file