diff --git a/lib/streaming_providers/base/models/drm/tenc_parser.py b/lib/streaming_providers/base/models/drm/tenc_parser.py index a15bb60..7fcd2c7 100644 --- a/lib/streaming_providers/base/models/drm/tenc_parser.py +++ b/lib/streaming_providers/base/models/drm/tenc_parser.py @@ -5,7 +5,7 @@ from typing import Optional class TencParser: """Parser for tenc (Track Encryption) boxes - simplified version.""" - + @staticmethod def extract_kid_from_tenc(tenc_data: bytes) -> Optional[bytes]: """ @@ -21,15 +21,15 @@ class TencParser: is_protected = tenc_data[7] if is_protected == 0: return None - + # Extract KID from bytes 9-24 if len(tenc_data) >= 25: kid_bytes = tenc_data[9:25] return kid_bytes - + except Exception: pass - + return None @staticmethod @@ -53,29 +53,14 @@ class TencParser: if box_type != b'tenc': return [] - # The KID appears to be at offset 16-31 in your data - # Let's check if that looks like a valid KID - if len(tenc_data) >= 32: - # Try offset 16 first (based on your data) - kid_bytes = tenc_data[16:32] - kid_hex = kid_bytes.hex().lower() + # Full box layout: header(8) + version/flags(4) + reserved(1) + # + crypt/skip(1) + default_isProtected(1) + ivSize(1) + KID(16) + if tenc_data[14] == 0: + return [] # unprotected track: the KID field is meaningless - # Validate it's not all zeros - if not all(c == '0' for c in kid_hex): - return [kid_hex] - - # If that didn't work, try scanning for valid-looking KID - for offset in range(0, len(tenc_data) - 16): - chunk = tenc_data[offset:offset + 16] - # Check if it looks like a valid KID (not all zeros, not repetitive) - if all(b == 0 for b in chunk): - continue - if all(b == chunk[0] for b in chunk): - continue - - # This could be a KID - kid_hex = chunk.hex().lower() - return [kid_hex] + kid_bytes = tenc_data[16:32] + if any(kid_bytes): + return [kid_bytes.hex().lower()] except Exception: pass diff --git a/lib/streaming_providers/base/utils/init_kid_resolver.py b/lib/streaming_providers/base/utils/init_kid_resolver.py new file mode 100644 index 0000000..56cc0d3 --- /dev/null +++ b/lib/streaming_providers/base/utils/init_kid_resolver.py @@ -0,0 +1,122 @@ +# streaming_providers/base/utils/init_kid_resolver.py +""" +Resolves the default KID of an AdaptationSet from its init segment (tenc box). + +Used by MPDRewriter when the MPD carries neither cenc:default_KID nor a +PSSH with KIDs and multiple keys are configured: without the KID the rewriter +cannot pick the right key for a segment. + +Design notes +------------ +* The rewriter stays free of network I/O: it receives this resolver as an + injected callable (init_url -> KID or None). +* Results are cached per init-segment path (scheme + host + path). The query + string is ignored on purpose: signed URLs change on every manifest refresh, + the KID of a given init segment does not. +* Failures (fetch error, no tenc) are cached too, but only briefly, so a live + MPD that refreshes every few seconds does not trigger a download per refresh + while a transient error still heals quickly. +""" + +import threading +import time +from collections import OrderedDict +from typing import Dict, Optional, Tuple +from urllib.parse import urlsplit + +from .logger import logger +from .mp4_pssh_extractor import MP4PSSHExtractor + +# The tenc box lives in moov at the start of the file; for single-file +# (SegmentBase) manifests this avoids downloading the whole MP4. +_PROBE_BYTES = 100 * 1024 + +_MISS = object() + + +class InitSegmentKidResolver: + """Thread-safe, cached init-segment -> KID lookup.""" + + def __init__( + self, + ttl_seconds: int = 3600, + failure_ttl_seconds: int = 60, + max_size: int = 1024, + ) -> None: + self._ttl = ttl_seconds + self._failure_ttl = failure_ttl_seconds + self._max_size = max_size + self._entries: "OrderedDict[str, Tuple[Optional[str], float]]" = OrderedDict() + self._lock = threading.Lock() + + @staticmethod + def _cache_key(init_url: str) -> str: + parts = urlsplit(init_url) + return f"{parts.scheme}://{parts.netloc}{parts.path}" + + def _get(self, key: str): + with self._lock: + entry = self._entries.get(key) + if entry is None: + return _MISS + kid, expires = entry + if expires <= time.monotonic(): + del self._entries[key] + return _MISS + self._entries.move_to_end(key) + return kid + + def _set(self, key: str, kid: Optional[str]) -> None: + ttl = self._ttl if kid else self._failure_ttl + with self._lock: + self._entries[key] = (kid, time.monotonic() + ttl) + self._entries.move_to_end(key) + while len(self._entries) > self._max_size: + self._entries.popitem(last=False) + + def resolve( + self, + init_url: str, + headers: Optional[Dict[str, str]] = None, + http_manager=None, + ) -> Optional[str]: + """ + Return the KID (32 lowercase hex chars) of the init segment, or None. + + Args: + init_url: Direct CDN URL of the init segment (not the proxied one). + headers: Segment auth headers (StreamingProvider.get_segment_headers()). + http_manager: Provider's HTTP manager, if any. + """ + key = self._cache_key(init_url) + cached = self._get(key) + if cached is not _MISS: + return cached + + request_headers = { + **(headers or {}), + "Range": f"bytes=0-{_PROBE_BYTES - 1}", + } + kids = MP4PSSHExtractor.extract_tenc_kids_from_url( + init_url, headers=request_headers, http_manager=http_manager + ) + kid = kids[0] if kids else None + + if len(kids) > 1: + logger.debug(f"Init segment has {len(kids)} tenc KIDs, using the first: {key}") + + self._set(key, kid) + return kid + + +_instance: Optional[InitSegmentKidResolver] = None +_instance_lock = threading.Lock() + + +def get_init_kid_resolver() -> InitSegmentKidResolver: + """Process-wide resolver so the cache survives across per-request rewriters.""" + global _instance + with _instance_lock: + if _instance is None: + _instance = InitSegmentKidResolver() + return _instance \ No newline at end of file diff --git a/lib/streaming_providers/base/utils/mp4_pssh_extractor.py b/lib/streaming_providers/base/utils/mp4_pssh_extractor.py index c570fb9..adbcbdf 100644 --- a/lib/streaming_providers/base/utils/mp4_pssh_extractor.py +++ b/lib/streaming_providers/base/utils/mp4_pssh_extractor.py @@ -42,19 +42,46 @@ class MP4PSSHExtractor: List of PSSHData objects with extracted information """ try: - if http_manager is not None: - response = http_manager.get(segment_url, headers=headers or {}, timeout=timeout, operation="api") - else: - import requests - response = requests.get(segment_url, timeout=timeout, headers=headers or {}) - response.raise_for_status() - data = response.content[:1024 * 100] + data = MP4PSSHExtractor._fetch_head(segment_url, timeout, headers, http_manager) return MP4PSSHExtractor.extract_from_bytes(data) except Exception as e: logger.error(f"Failed to extract PSSH from {segment_url}: {e}") return [] + @staticmethod + def extract_tenc_kids_from_url( + segment_url: str, + timeout: int = 10, + headers: Optional[Dict[str, str]] = None, + http_manager=None, + ) -> List[str]: + """ + Download an init segment and return the default KIDs of its protected + tracks (from tenc boxes), as normalized 32-char hex strings. + + Unlike _extract_all_tenc_kids (a byte-resync scan used as PSSH fallback), + this walks the box tree structurally, so it does not depend on box sizes + happening to line up. + """ + try: + data = MP4PSSHExtractor._fetch_head(segment_url, timeout, headers, http_manager) + return MP4PSSHExtractor.extract_tenc_kids_structured(data) + except Exception as e: + logger.error(f"Failed to extract tenc KIDs from {segment_url}: {e}") + return [] + + @staticmethod + def _fetch_head(segment_url: str, timeout: int, headers: Optional[Dict[str, str]], http_manager) -> bytes: + """Fetch a segment and return at most its first 100 KB.""" + if http_manager is not None: + response = http_manager.get(segment_url, headers=headers or {}, timeout=timeout, operation="api") + else: + import requests + response = requests.get(segment_url, timeout=timeout, headers=headers or {}) + response.raise_for_status() + return response.content[:1024 * 100] + @staticmethod def extract_from_bytes(data: bytes) -> List[PSSHData]: """ @@ -148,6 +175,55 @@ class MP4PSSHExtractor: return pssh_data_list + # Fixed-size sample entry fields that precede child boxes (ISO/IEC 14496-12) + _SAMPLE_ENTRY_FIELDS = {b"encv": 78, b"enca": 28} + _TENC_CONTAINERS = {b"moov", b"trak", b"mdia", b"minf", b"stbl", b"sinf", b"schi"} + + @staticmethod + def extract_tenc_kids_structured(data: bytes) -> List[str]: + """ + Collect KIDs from tenc boxes by walking the box tree + moov/trak/mdia/minf/stbl/stsd/{encv,enca}/sinf/schi/tenc. + + Handles the two places where children do not start right after the + 8-byte box header: stsd (version/flags + entry_count) and the + encv/enca sample entries (fixed-size fields before child boxes). + Unprotected tracks (default_isProtected == 0) are skipped. + """ + kids: List[str] = [] + MP4PSSHExtractor._walk_for_tenc(data, 0, len(data), kids) + return kids + + @staticmethod + def _walk_for_tenc(data: bytes, start: int, end: int, kids: List[str]) -> None: + offset = start + while offset + 8 <= end: + box_size = struct.unpack(">I", data[offset: offset + 4])[0] + box_type = data[offset + 4: offset + 8] + if box_size < 8 or offset + box_size > end: + break # malformed or truncated: stop, don't guess + + body = offset + 8 + box_end = offset + box_size + + if box_type == b"tenc": + # header(8) + version/flags(4) + reserved(1) + crypt/skip(1) + # + isProtected(1) + ivSize(1) + KID(16) + if box_size >= 32 and data[offset + 14] != 0: + kid = data[offset + 16: offset + 32].hex() + if any(data[offset + 16: offset + 32]) and kid not in kids: + kids.append(kid) + elif box_type in MP4PSSHExtractor._TENC_CONTAINERS: + MP4PSSHExtractor._walk_for_tenc(data, body, box_end, kids) + elif box_type == b"stsd": + # version/flags(4) + entry_count(4), then sample entries + MP4PSSHExtractor._walk_for_tenc(data, body + 8, box_end, kids) + elif box_type in MP4PSSHExtractor._SAMPLE_ENTRY_FIELDS: + skip = MP4PSSHExtractor._SAMPLE_ENTRY_FIELDS[box_type] + MP4PSSHExtractor._walk_for_tenc(data, body + skip, box_end, kids) + + offset = box_end + @staticmethod def _extract_all_tenc_kids(data: bytes) -> List[str]: """ diff --git a/lib/streaming_providers/base/utils/mpd_rewriter.py b/lib/streaming_providers/base/utils/mpd_rewriter.py index 847b05b..4b02ff2 100644 --- a/lib/streaming_providers/base/utils/mpd_rewriter.py +++ b/lib/streaming_providers/base/utils/mpd_rewriter.py @@ -8,7 +8,7 @@ import base64 import struct import xml.etree.ElementTree as ET from dataclasses import dataclass -from typing import Optional, Tuple, Set, Dict +from typing import Callable, Optional, Tuple, Set, Dict from urllib.parse import urljoin, quote, urlencode from datetime import datetime, timezone from email.utils import parsedate_to_datetime @@ -68,7 +68,13 @@ class MPDRewriter: blocklist_path: str = "representation_blocklist.json", clearkey_receiver_side: bool = False, segment_headers: Optional[Dict[str, str]] = None, + kid_resolver: Optional[Callable[[str], Optional[str]]] = None, ): + # kid_resolver: init-segment URL -> KID (32 hex chars) or None. Injected + # so the rewriter itself stays free of network I/O. Only consulted in + # multi-key server-side decrypt mode, for AdaptationSets whose MPD + # carries no KID (see _resolve_kid_from_init_segment). + self.kid_resolver = kid_resolver self.media_proxy_url = media_proxy_url.rstrip("/") self.provider_proxy_url = provider_proxy_url self.key_config = KeyConfiguration(clearkey_keyids or {}) @@ -378,6 +384,10 @@ class MPDRewriter: # Extract KID extracted_kid = self._extract_kid_from_adaptationset(adaptation_set) + if not extracted_kid and self._needs_kid_from_init_segment(): + extracted_kid = self._resolve_kid_from_init_segment( + adaptation_set, as_base_url, unique_id + ) if extracted_kid: as_id_to_kid[unique_id] = extracted_kid # logger.debug(f"AdaptationSet {unique_id} KID: {extracted_kid[:8]}...") @@ -410,6 +420,74 @@ class MPDRewriter: return encrypted_ids, as_id_to_kid, base_url_map + def _needs_kid_from_init_segment(self) -> bool: + """KIDs are only consumed for key selection in multi-key server-side decrypt.""" + return ( + self.kid_resolver is not None + and bool(self.key_config.keys) + and not self.clearkey_receiver_side + and not self.key_config.single_key_mode + ) + + def _find_init_segment_url(self, adaptation_set: ET.Element, base_url: str) -> Optional[str]: + """ + Resolve the init segment URL of the first Representation, using the same + base-URL logic that is later used to proxy the segments. + + Supports SegmentTemplate (Representation or AdaptationSet level) and + single-file Representation BaseURL (SegmentBase). Period-level + SegmentTemplate is not handled. + """ + representations = adaptation_set.findall("mpd:Representation", self.MPD_NAMESPACE) + first_rep = representations[0] if representations else None + rep_id = first_rep.get("id", "") if first_rep is not None else "" + bandwidth = first_rep.get("bandwidth", "0") if first_rep is not None else "0" + + containers = ([first_rep] if first_rep is not None else []) + [adaptation_set] + for container in containers: + template = container.find("mpd:SegmentTemplate", self.MPD_NAMESPACE) + if template is None or not template.get("initialization"): + continue + init = URLResolver.substitute_template_variables( + template.get("initialization"), + representation_id=rep_id, + bandwidth=bandwidth, + ) + if "$" in init: + return None # unsupported variable/format specifier in an init template + return self._urljoin_preserve_query(base_url, init) + + if first_rep is not None: + base_elem = first_rep.find("mpd:BaseURL", self.MPD_NAMESPACE) + if base_elem is not None and base_elem.text and base_elem.text.strip(): + return self._urljoin_preserve_query(base_url, base_elem.text.strip()) + + return None + + def _resolve_kid_from_init_segment( + self, adaptation_set: ET.Element, base_url: str, unique_id: str + ) -> Optional[str]: + """Fallback KID lookup (tenc in the init segment) via the injected resolver.""" + init_url = self._find_init_segment_url(adaptation_set, base_url) + if not init_url: + logger.warning( + f"AdaptationSet {unique_id}: no KID in MPD and no init segment URL found" + ) + return None + try: + kid = self.kid_resolver(init_url) + except Exception as e: + logger.debug(f"KID resolver failed for AdaptationSet {unique_id}: {e}") + kid = None + if kid: + logger.debug(f"AdaptationSet {unique_id} KID from init tenc: {kid[:8]}...") + else: + logger.warning( + f"AdaptationSet {unique_id}: no KID in MPD and none in init segment; " + f"fallback key will be used" + ) + return kid + def _extract_kid_from_adaptationset(self, adaptation_set: ET.Element) -> Optional[str]: """Extract KID from ContentProtection elements in an AdaptationSet.""" # Try cenc:default_KID first diff --git a/service.py b/service.py index 71ce159..2d8812e 100644 --- a/service.py +++ b/service.py @@ -27,6 +27,7 @@ try: ProviderEnableManager, ) from streaming_providers.base.utils import MPDCacheManager, MPDRewriter, logger + from streaming_providers.base.utils.init_kid_resolver import get_init_kid_resolver from streaming_providers.base.utils.environment import ( get_environment_manager, get_vfs_instance, @@ -396,6 +397,13 @@ class UltimateService: return manifest_response.text, ttl, provider_proxy_url, segment_headers, manifest_response.url + def _make_kid_resolver(self, provider: str, segment_headers: Optional[dict]): + resolver = get_init_kid_resolver() + http_manager = self.manager.get_provider_http_manager(provider) + return lambda init_url: resolver.resolve( + init_url, headers=segment_headers, http_manager=http_manager + ) + def _get_decrypted_cached(self, key: str, max_stale: int = 0) -> Optional[str]: entry = self._decrypted_cache.get(key) if not entry: @@ -581,6 +589,7 @@ class UltimateService: self.media_proxy_url, provider_proxy_url, keyids, highest_quality_only, provider=provider, channel=channel_id, clearkey_receiver_side=receiver_side, segment_headers=segment_headers, + id_resolver=self._make_kid_resolver(provider, segment_headers), ) rewritten_mpd = rewriter.rewrite_mpd(manifest_text, effective_url) return rewritten_mpd, min(ttl, 10) # holds key material — keep exposure window short @@ -636,6 +645,7 @@ class UltimateService: self.media_proxy_url, provider_proxy_url, keyids, highest_quality_only, provider=provider, channel=channel_id, clearkey_receiver_side=receiver_side, segment_headers=segment_headers, + id_resolver=self._make_kid_resolver(provider, segment_headers), ) rewritten_mpd = rewriter.rewrite_mpd(manifest_text, effective_url) return rewritten_mpd, min(ttl, 30)