extract kids from tenc if necessary

This commit is contained in:
Nirvana
2026-10-01 12:42:30 +02:00
parent 74de08d403
commit 952b3874de
5 changed files with 305 additions and 34 deletions
@@ -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
@@ -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
@@ -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]:
"""
@@ -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
+10
View File
@@ -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)