From 9bef9770c8c3aec3c01066cc9e844cc39c861ef5 Mon Sep 17 00:00:00 2001 From: NTFSvolume <172021377+NTFSvolume@users.noreply.github.com> Date: Mon, 13 Jul 2026 16:37:43 -0500 Subject: [PATCH 1/2] refactor: add type hints --- m3u8/__init__.py | 39 ++- m3u8/httpclient.py | 32 +- m3u8/mixins.py | 48 ++- m3u8/model.py | 749 ++++++++++++++++++++++++++++---------------- runtests | 6 +- tests/m3u8server.py | 30 +- tests/playlists.py | 19 +- 7 files changed, 592 insertions(+), 331 deletions(-) diff --git a/m3u8/__init__.py b/m3u8/__init__.py index f6fc5174..5433ddf8 100644 --- a/m3u8/__init__.py +++ b/m3u8/__init__.py @@ -3,9 +3,11 @@ # license that can be found in the LICENSE file. import os +from collections.abc import Callable, Mapping +from typing import Any from urllib.parse import urljoin, urlsplit -from m3u8.httpclient import DefaultHTTPClient +from m3u8.httpclient import DefaultHTTPClient, _HTTPClientProtocol from m3u8.model import ( M3U8, ContentSteering, @@ -64,7 +66,14 @@ ) -def loads(content, uri=None, custom_tags_parser=None): +type _CustomTagsParser = Callable[[str, int, dict[str, Any], dict[str, Any]], object] + + +def loads( + content: str, + uri: str | None = None, + custom_tags_parser: _CustomTagsParser | None = None, +) -> M3U8: """ Given a string with a m3u8 content, returns a M3U8 object. Optionally parses a uri to set a correct base_uri on the M3U8 object. @@ -79,26 +88,30 @@ def loads(content, uri=None, custom_tags_parser=None): def load( - uri, - timeout=None, - headers={}, - custom_tags_parser=None, - http_client=DefaultHTTPClient(), - verify_ssl=True, -): + uri: str, + timeout: float | None = None, + headers: Mapping[str, Any] | None = None, + custom_tags_parser: _CustomTagsParser | None = None, + http_client: _HTTPClientProtocol = DefaultHTTPClient(), + verify_ssl: bool = True, +) -> M3U8: """ Retrieves the content from a given URI and returns a M3U8 object. Raises ValueError if invalid content or IOError if request fails. """ base_uri_parts = urlsplit(uri) if base_uri_parts.scheme and base_uri_parts.netloc: - content, base_uri = http_client.download(uri, timeout, headers, verify_ssl) + content, base_uri = http_client.download( + uri, timeout, headers or {}, verify_ssl + ) return M3U8(content, base_uri=base_uri, custom_tags_parser=custom_tags_parser) - else: - return _load_from_file(uri, custom_tags_parser) + + return _load_from_file(uri, custom_tags_parser) -def _load_from_file(uri, custom_tags_parser=None): +def _load_from_file( + uri: os.PathLike[str] | str, custom_tags_parser: _CustomTagsParser | None = None +) -> M3U8: with open(uri, encoding="utf8") as fileobj: raw_content = fileobj.read().strip() base_uri = os.path.dirname(uri) diff --git a/m3u8/httpclient.py b/m3u8/httpclient.py index a6babad3..b2ed834d 100644 --- a/m3u8/httpclient.py +++ b/m3u8/httpclient.py @@ -1,18 +1,38 @@ import gzip import ssl import urllib.request +from collections.abc import Mapping +from typing import Any, Protocol from urllib.parse import urljoin -class DefaultHTTPClient: - def __init__(self, proxies=None): - self.proxies = proxies +class _HTTPClientProtocol(Protocol): # noqa: Y046 + def download( + self, + uri: str, + timeout: float | None = None, + headers: Mapping[str, Any] | None = None, + verify_ssl: bool = True, + ) -> tuple[str, str]: ... + + +class DefaultHTTPClient(_HTTPClientProtocol): + def __init__(self, proxies: dict[str, str] | None = None) -> None: + self.proxies: dict[str, str] | None = proxies + + def download( + self, + uri: str, + timeout: float | None = None, + headers: Mapping[str, Any] | None = None, + verify_ssl: bool = True, + ) -> tuple[str, str]: - def download(self, uri, timeout=None, headers={}, verify_ssl=True): proxy_handler = urllib.request.ProxyHandler(self.proxies) https_handler = HTTPSHandler(verify_ssl=verify_ssl) opener = urllib.request.build_opener(proxy_handler, https_handler) - opener.addheaders = headers.items() + if headers: + opener.addheaders = list(headers.items()) resource = opener.open(uri, timeout=timeout) base_uri = urljoin(resource.geturl(), ".") @@ -28,7 +48,7 @@ def download(self, uri, timeout=None, headers={}, verify_ssl=True): class HTTPSHandler: - def __new__(self, verify_ssl=True): + def __new__(cls, verify_ssl: bool = True) -> urllib.request.HTTPSHandler: context = ssl.create_default_context() if not verify_ssl: context.check_hostname = False diff --git a/m3u8/mixins.py b/m3u8/mixins.py index 40ffb2fc..2d8cd9b7 100644 --- a/m3u8/mixins.py +++ b/m3u8/mixins.py @@ -1,10 +1,19 @@ +from __future__ import annotations + +from collections.abc import Iterable from os.path import dirname +from typing import Protocol from urllib.parse import urljoin, urlsplit -class BasePathMixin: +class HasUri(Protocol): + uri: str | None + base_uri: str | None + + +class BasePathMixin(HasUri): @property - def absolute_uri(self): + def absolute_uri(self) -> str | None: if self.uri is None: return None @@ -20,33 +29,40 @@ def absolute_uri(self): return ret @property - def base_path(self): + def base_path(self) -> str | None: if self.uri is None: return None return dirname(self.get_path_from_uri()) - def get_path_from_uri(self): + def get_path_from_uri(self) -> str: """Some URIs have a slash in the query string.""" + assert self.uri return self.uri.split("?")[0] @base_path.setter - def base_path(self, newbase_path): - if self.uri is not None: - if not self.base_path: - self.uri = f"{newbase_path}/{self.uri}" - else: - self.uri = self.uri.replace(self.base_path, newbase_path) + def base_path(self, newbase_path: str) -> None: + if self.uri is None: + return + + if not self.base_path: + self.uri = f"{newbase_path}/{self.uri}" + else: + self.uri = self.uri.replace(self.base_path, newbase_path) -class GroupedBasePathMixin: - def _set_base_uri(self, new_base_uri): +class GroupedBasePathMixin[T: BasePathMixin](Iterable[T]): + @property + def base_uri(self) -> str: ... + + @base_uri.setter + def base_uri(self, new_base_uri: str) -> None: for item in self: item.base_uri = new_base_uri - base_uri = property(None, _set_base_uri) + @property + def base_path(self) -> str: ... - def _set_base_path(self, newbase_path): + @base_path.setter + def base_path(self, newbase_path: str) -> None: for item in self: item.base_path = newbase_path - - base_path = property(None, _set_base_path) diff --git a/m3u8/model.py b/m3u8/model.py index df1e7530..3fffafd8 100644 --- a/m3u8/model.py +++ b/m3u8/model.py @@ -1,8 +1,11 @@ -# Copyright 2014 Globo.com Player authors. All rights reserved. -# Use of this source code is governed by a MIT License -# license that can be found in the LICENSE file. +from __future__ import annotations + +import dataclasses +import datetime as dt import decimal import os +from collections.abc import Callable, Iterable, Mapping, Sequence +from typing import Any, ClassVar, Literal from m3u8.mixins import BasePathMixin, GroupedBasePathMixin from m3u8.parser import format_date_time, parse @@ -15,9 +18,10 @@ ext_x_start, ) +type _CustomTagsParser = Callable[[str, int, dict[str, Any], dict[str, Any]], object] + -class MalformedPlaylistError(Exception): - pass +class MalformedPlaylistError(Exception): ... class M3U8: @@ -138,7 +142,7 @@ class M3U8: https://github.com/image-media-playlist/spec/blob/master/image_media_playlist_v0_4.pdf """ - simple_attributes = ( + simple_attributes: tuple[tuple[str, str], ...] = ( # obj attribute # parser attribute ("is_variant", "is_variant"), ("is_endlist", "is_endlist"), @@ -154,36 +158,71 @@ class M3U8: ("is_images_only", "is_images_only"), ) + simple_attributes: tuple[tuple[str, str], ...] + data: dict[str, Any] + keys: list[Key] + segment_map: list[InitializationSection] + segments: SegmentList + files: list[str | None] + media: MediaList + playlists: PlaylistList[Playlist] + iframe_playlists: PlaylistList[IFramePlaylist] + image_playlists: PlaylistList[ImagePlaylist] + start: Start + server_control: ServerControl + part_inf: PartInformation + skip: Skip | None + rendition_reports: RenditionReportList + session_data: SessionDataList + session_keys: list[SessionKey | None] + preload_hint: PreloadHint + content_steering: ContentSteering + + is_variant: bool | None + is_endlist: bool | None + is_i_frames_only: bool | None + target_duration: float | None + media_sequence: int | None + program_date_time: dt.datetime | None + is_independent_segments: bool | None + version: str | None + allow_cache: str | None + playlist_type: str | None + discontinuity_sequence: Any | None + is_images_only: bool | None + def __init__( - self, - content=None, - base_path=None, - base_uri=None, - strict=False, - custom_tags_parser=None, + self, + content: str | None = None, + base_path: str | None = None, + base_uri: str | None = None, + strict: bool = False, + custom_tags_parser: _CustomTagsParser | None = None, ): if content is not None: self.data = parse(content, strict, custom_tags_parser) else: self.data = {} - self._base_uri = base_uri + self._base_uri: str | None = base_uri if self._base_uri: if not self._base_uri.endswith("/"): self._base_uri += "/" self._initialize_attributes() + self._base_path: str | None = base_path self.base_path = base_path - def _initialize_attributes(self): - self.keys = [] + def _initialize_attributes(self) -> None: + self.keys: list[tuple[Key, ...]] = [] for keys in self.data.get("keys", []): - self.keys.append(tuple( - [Key(base_uri=self.base_uri, **params) for params in keys] - )) + self.keys.append( + tuple(Key(base_uri=self.base_uri, **params) for params in keys) + ) self.segment_map = [ - InitializationSection(base_uri=self.base_uri, **params) if params else None + InitializationSection(base_uri=self.base_uri, **params) for params in self.data.get("segment_map", []) + if params ] self.segments = SegmentList( [ @@ -290,11 +329,11 @@ def __unicode__(self): return self.dumps() @property - def base_uri(self): + def base_uri(self) -> str | None: return self._base_uri @base_uri.setter - def base_uri(self, new_base_uri): + def base_uri(self, new_base_uri: str) -> None: self._base_uri = new_base_uri self.media.base_uri = new_base_uri self.playlists.base_uri = new_base_uri @@ -314,15 +353,15 @@ def base_uri(self, new_base_uri): self.content_steering.base_uri = new_base_uri @property - def base_path(self): + def base_path(self) -> str | None: return self._base_path @base_path.setter - def base_path(self, newbase_path): + def base_path(self, newbase_path: str | None) -> None: self._base_path = newbase_path self._update_base_path() - def _update_base_path(self): + def _update_base_path(self) -> None: if self._base_path is None: return for keys in self.keys: @@ -331,6 +370,7 @@ def _update_base_path(self): for key in self.session_keys: if key: key.base_path = self._base_path + self.media.base_path = self._base_path self.segments.base_path = self._base_path self.playlists.base_path = self._base_path @@ -342,30 +382,28 @@ def _update_base_path(self): if self.content_steering: self.content_steering.base_path = self._base_path - def add_playlist(self, playlist): + def add_playlist(self, playlist: Playlist) -> None: self.is_variant = True self.playlists.append(playlist) - def add_iframe_playlist(self, iframe_playlist): - if iframe_playlist is not None: - self.is_variant = True - self.iframe_playlists.append(iframe_playlist) + def add_iframe_playlist(self, iframe_playlist: IFramePlaylist) -> None: + self.is_variant = True + self.iframe_playlists.append(iframe_playlist) - def add_image_playlist(self, image_playlist): - if image_playlist is not None: - self.is_variant = True - self.image_playlists.append(image_playlist) + def add_image_playlist(self, image_playlist: ImagePlaylist) -> None: + self.is_variant = True + self.image_playlists.append(image_playlist) - def add_media(self, media): + def add_media(self, media: Media) -> None: self.media.append(media) - def add_segment(self, segment): + def add_segment(self, segment: Segment) -> None: self.segments.append(segment) - def add_rendition_report(self, report): + def add_rendition_report(self, report: RenditionReport) -> None: self.rendition_reports.append(report) - def dumps(self, timespec="milliseconds", infspec="auto"): + def dumps(self, timespec: str = "milliseconds", infspec: str = "auto") -> str: """ Returns the current m3u8 as a string. You could also use unicode() or str() @@ -434,7 +472,7 @@ def dumps(self, timespec="milliseconds", infspec="auto"): return "\n".join(output) - def dump(self, filename): + def dump(self, filename: str) -> None: """ Saves the current m3u8 to ``filename`` """ @@ -443,7 +481,7 @@ def dump(self, filename): with open(filename, "w") as fileobj: fileobj.write(self.dumps()) - def _create_sub_directories(self, filename): + def _create_sub_directories(self, filename: str) -> None: if not os.path.isabs(filename): filename = os.path.join(os.getcwd(), filename) @@ -525,34 +563,59 @@ class Segment(BasePathMixin): Additional values which custom_tags_parser might store per segment """ + media_sequence: int | None + uri: str | None + duration: float | None + title: str + bitrate: int | None + byterange: str | None + program_date_time: dt.datetime | None + current_program_date_time: dt.datetime | None + discontinuity: bool + cue_out_start: bool + cue_out_explicitly_duration: bool + cue_out: bool + cue_in: bool + scte35: str | None + oatcls_scte35: str | None + scte35_duration: float | None + scte35_elapsedtime: Any | None + asset_metadata: dict[str, Any] | None + key: Key | None + parts: PartialSegmentList + init_section: InitializationSection | None + dateranges: DateRangeList + gap_tag: Any | None + custom_parser_values: dict[str, Any] + def __init__( - self, - uri=None, - base_uri=None, - program_date_time=None, - current_program_date_time=None, - duration=None, - title=None, - bitrate=None, - byterange=None, - cue_out=False, - cue_out_start=False, - cue_out_explicitly_duration=False, - cue_in=False, - discontinuity=False, - keys=None, - scte35=None, - oatcls_scte35=None, - scte35_duration=None, - scte35_elapsedtime=None, - asset_metadata=None, - keyobjects=(), - parts=None, - init_section=None, - dateranges=None, - gap_tag=None, - media_sequence=None, - custom_parser_values=None, + self, + uri: str | None = None, + base_uri: str | None = None, + program_date_time: dt.datetime | None = None, + current_program_date_time: dt.datetime | None = None, + duration: float | None = None, + title: str | None = None, + bitrate: int | None = None, + byterange: str | None = None, + cue_out: bool = False, + cue_out_start: bool = False, + cue_out_explicitly_duration: bool = False, + cue_in: bool = False, + discontinuity: bool = False, + key: object = None, + scte35: str | None = None, + oatcls_scte35: str | None = None, + scte35_duration: float | None = None, + scte35_elapsedtime=None, + asset_metadata: Mapping[str, str] | None = None, + keyobjects: tuple[Key, ...] = (), + parts: Iterable[Mapping[str, Any]] | None = None, + init_section: Mapping[str, Any] | None = None, + dateranges: Iterable[Mapping[str, Any]] | None = None, + gap_tag: list[Mapping[str, Any]] | None = None, + media_sequence: int | None = None, + custom_parser_values: dict[str, Any] | None = None, ): self.media_sequence = media_sequence self.uri = uri @@ -572,8 +635,8 @@ def __init__( self.oatcls_scte35 = oatcls_scte35 self.scte35_duration = scte35_duration self.scte35_elapsedtime = scte35_elapsedtime - self.asset_metadata = asset_metadata - self._keys = () + self.asset_metadata = dict(asset_metadata) if asset_metadata else None + self._keys: tuple[Key, ...] = () self.keys = keyobjects self.parts = PartialSegmentList( [PartialSegment(base_uri=self._base_uri, **partial) for partial in parts] @@ -591,19 +654,24 @@ def __init__( self.custom_parser_values = custom_parser_values or {} @property - def keys(self): + def keys(self) -> tuple[Key, ...]: return self._keys @keys.setter - def keys(self, value): + def keys(self, value: tuple[Key, ...]) -> None: if not isinstance(value, tuple): raise ValueError("keys must be a tuple") self._keys = value - def add_part(self, part): + def add_part(self, part: PartialSegment) -> None: self.parts.append(part) - def dumps(self, last_segment, timespec="milliseconds", infspec="auto"): + def dumps( + self, + last_segment: PartialSegment, + timespec: str = "milliseconds", + infspec: str = "auto", + ): output = [] if last_segment and self.keys != last_segment.keys: @@ -696,34 +764,34 @@ def dumps(self, last_segment, timespec="milliseconds", infspec="auto"): return "".join(output) - def __str__(self): + def __str__(self) -> str: return self.dumps(None) @property - def base_path(self): + def base_path(self) -> str | None: return super().base_path @base_path.setter - def base_path(self, newbase_path): - super(Segment, self.__class__).base_path.fset(self, newbase_path) + def base_path(self, newbase_path: str) -> None: + super().base_path = newbase_path self.parts.base_path = newbase_path if self.init_section is not None: self.init_section.base_path = newbase_path @property - def base_uri(self): + def base_uri(self) -> str | None: return self._base_uri @base_uri.setter - def base_uri(self, newbase_uri): + def base_uri(self, newbase_uri: str | None) -> None: self._base_uri = newbase_uri self.parts.base_uri = newbase_uri if self.init_section is not None: self.init_section.base_uri = newbase_uri -class SegmentList(list, GroupedBasePathMixin): - def dumps(self, timespec="milliseconds", infspec="auto"): +class SegmentList(list[Segment], GroupedBasePathMixin[Segment]): + def dumps(self, timespec: str = "milliseconds", infspec: str = "auto") -> str: output = [] last_segment = None for segment in self: @@ -731,14 +799,14 @@ def dumps(self, timespec="milliseconds", infspec="auto"): last_segment = segment return "\n".join(output) - def __str__(self): + def __str__(self) -> str: return self.dumps() @property - def uri(self): + def uri(self) -> list[str | None]: return [seg.uri for seg in self] - def by_keys(self, keys): + def by_keys(self, keys: tuple[Key, ...]) -> list[Segment]: return [segment for segment in self if segment.keys == keys] @@ -781,18 +849,29 @@ class PartialSegment(BasePathMixin): attribute. """ + base_uri: str + uri: str | None + duration: float | None + program_date_time: dt.datetime | None + current_program_date_time: dt.datetime | None + byterange: str | None + independent: bool + gap: str | None + dateranges: DateRangeList + gap_tag: str | None + def __init__( - self, - base_uri, - uri, - duration, - program_date_time=None, - current_program_date_time=None, - byterange=None, - independent=None, - gap=None, - dateranges=None, - gap_tag=None, + self, + base_uri: str, + uri: str | None, + duration: float | None, + program_date_time: dt.datetime | None = None, + current_program_date_time: dt.datetime | None = None, + byterange: str | None = None, + independent=None, + gap=None, + dateranges: Iterable[Mapping[str, Any]] | None = None, + gap_tag=None, ): self.base_uri = base_uri self.uri = uri @@ -807,7 +886,7 @@ def __init__( ) self.gap_tag = gap_tag - def dumps(self, last_segment): + def dumps(self, last_segment: object) -> str: output = [] if len(self.dateranges): @@ -833,12 +912,12 @@ def dumps(self, last_segment): return "".join(output) - def __str__(self): + def __str__(self) -> str: return self.dumps(None) -class PartialSegmentList(list, GroupedBasePathMixin): - def __str__(self): +class PartialSegmentList(list[PartialSegment], GroupedBasePathMixin[PartialSegment]): + def __str__(self) -> str: output = [str(part) for part in self] return "\n".join(output) @@ -861,17 +940,23 @@ class Key(BasePathMixin): """ - tag = ext_x_key + tag: ClassVar[str] = ext_x_key + method: str + base_uri: str + uri: str | None + iv: str | None + keyformat: str | None + keyformatversions: str | None def __init__( - self, - method, - base_uri, - uri=None, - iv=None, - keyformat=None, - keyformatversions=None, - **kwargs, + self, + method: str, + base_uri: str, + uri: str | None = None, + iv: str | None = None, + keyformat: str | None = None, + keyformatversions: str | None = None, + **kwargs, ): self.method = method self.uri = uri @@ -881,7 +966,7 @@ def __init__( self.base_uri = base_uri self._extra_params = kwargs - def __str__(self): + def __str__(self) -> str: output = [ "METHOD=%s" % self.method, ] @@ -896,21 +981,19 @@ def __str__(self): return self.tag + ":" + ",".join(output) - def __eq__(self, other): - if not other: - return False - return ( - self.method == other.method - and self.uri == other.uri - and self.iv == other.iv - and self.base_uri == other.base_uri - and self.keyformat == other.keyformat - and self.keyformatversions == other.keyformatversions + def __eq__(self, other: object) -> bool: + if not isinstance(other, type(self)): + return NotImplemented + + return bool( + self.method == other.method + and self.uri == other.uri + and self.iv == other.iv + and self.base_uri == other.base_uri + and self.keyformat == other.keyformat + and self.keyformatversions == other.keyformatversions ) - def __ne__(self, other): - return not self.__eq__(other) - class InitializationSection(BasePathMixin): """ @@ -927,14 +1010,19 @@ class InitializationSection(BasePathMixin): uri the segment comes from in URI hierarchy. ex.: http://example.com/path/to """ - tag = ext_x_map + tag: ClassVar[str] = ext_x_map + base_uri: str + uri: str | None + byterange: str | None - def __init__(self, base_uri, uri, byterange=None): + def __init__( + self, base_uri: str, uri: str | None, byterange: str | None = None + ) -> None: self.base_uri = base_uri self.uri = uri self.byterange = byterange - def __str__(self): + def __str__(self) -> str: output = [] if self.uri: output.append("URI=" + quoted(self.uri)) @@ -942,21 +1030,19 @@ def __str__(self): output.append("BYTERANGE=" + quoted(self.byterange)) return "{tag}:{attributes}".format(tag=self.tag, attributes=",".join(output)) - def __eq__(self, other): - if not other: - return False - return ( - self.uri == other.uri - and self.byterange == other.byterange - and self.base_uri == other.base_uri - ) + def __eq__(self, other: object) -> bool: + if not isinstance(other, type(self)): + return NotImplemented - def __ne__(self, other): - return not self.__eq__(other) + return bool( + self.uri == other.uri + and self.byterange == other.byterange + and self.base_uri == other.base_uri + ) class SessionKey(Key): - tag = ext_x_session_key + tag: ClassVar[str] = ext_x_session_key class Playlist(BasePathMixin): @@ -974,7 +1060,18 @@ class Playlist(BasePathMixin): More info: http://tools.ietf.org/html/draft-pantos-http-live-streaming-07#section-3.3.10 """ - def __init__(self, uri, stream_info, media, base_uri): + base_uri: str | None + uri: str | None + stream_info: StreamInfo + media: MediaList + + def __init__( + self, + uri: str | None, + stream_info: Mapping[str, Any], + media: MediaList, + base_uri: str, + ) -> None: self.uri = uri self.base_uri = base_uri @@ -1011,9 +1108,10 @@ def __init__(self, uri, stream_info, media, base_uri): self.media += filter(lambda m: m.group_id == group_id, media) - def __str__(self): + def __str__(self) -> str: media_types = [] stream_inf = [str(self.stream_info)] + assert self.uri for media in self.media: if media.type in media_types: continue @@ -1025,7 +1123,7 @@ def __str__(self): return "#EXT-X-STREAM-INF:" + ",".join(stream_inf) + "\n" + self.uri -class IFramePlaylist(BasePathMixin): +class IFramePlaylist(Playlist): """ IFramePlaylist object representing a link to a variant M3U8 i-frame playlist with a specific bitrate. @@ -1039,7 +1137,16 @@ class IFramePlaylist(BasePathMixin): More info: http://tools.ietf.org/html/draft-pantos-http-live-streaming-07#section-3.3.13 """ - def __init__(self, base_uri, uri, iframe_stream_info): + uri: str | None + base_uri: str | None + iframe_stream_info: StreamInfo + + def __init__( + self, + base_uri: str, + uri: str | None, + iframe_stream_info: Mapping[str, Any], + ) -> None: self.uri = uri self.base_uri = base_uri @@ -1070,7 +1177,7 @@ def __init__(self, base_uri, uri, iframe_stream_info): req_video_layout=None, ) - def __str__(self): + def __str__(self) -> str: iframe_stream_inf = [] if self.iframe_stream_info.program_id: iframe_stream_inf.append( @@ -1084,9 +1191,9 @@ def __str__(self): ) if self.iframe_stream_info.resolution: res = ( - str(self.iframe_stream_info.resolution[0]) - + "x" - + str(self.iframe_stream_info.resolution[1]) + str(self.iframe_stream_info.resolution[0]) + + "x" + + str(self.iframe_stream_info.resolution[1]) ) iframe_stream_inf.append("RESOLUTION=" + res) if self.iframe_stream_info.codecs: @@ -1113,41 +1220,25 @@ def __str__(self): return "#EXT-X-I-FRAME-STREAM-INF:" + ",".join(iframe_stream_inf) +@dataclasses.dataclass(slots=True) class StreamInfo: - bandwidth = None - closed_captions = None - average_bandwidth = None - program_id = None - resolution = None - codecs = None - audio = None - video = None - subtitles = None - frame_rate = None - video_range = None - hdcp_level = None - pathway_id = None - stable_variant_id = None - req_video_layout = None - - def __init__(self, **kwargs): - self.bandwidth = kwargs.get("bandwidth") - self.closed_captions = kwargs.get("closed_captions") - self.average_bandwidth = kwargs.get("average_bandwidth") - self.program_id = kwargs.get("program_id") - self.resolution = kwargs.get("resolution") - self.codecs = kwargs.get("codecs") - self.audio = kwargs.get("audio") - self.video = kwargs.get("video") - self.subtitles = kwargs.get("subtitles") - self.frame_rate = kwargs.get("frame_rate") - self.video_range = kwargs.get("video_range") - self.hdcp_level = kwargs.get("hdcp_level") - self.pathway_id = kwargs.get("pathway_id") - self.stable_variant_id = kwargs.get("stable_variant_id") - self.req_video_layout = kwargs.get("req_video_layout") - - def __str__(self): + bandwidth: int | None + closed_captions: Any | None + average_bandwidth: int | None + program_id: int | None + resolution: tuple[int, int] | None + codecs: str | None + audio: str | None + video: str | None + subtitles: str | None + frame_rate: float | None + video_range: str | None + hdcp_level: str | None + pathway_id: str | None + stable_variant_id: str | None + req_video_layout: str | None + + def __str__(self) -> str: stream_inf = [] if self.program_id is not None: stream_inf.append("PROGRAM-ID=%d" % self.program_id) @@ -1206,23 +1297,39 @@ class Media(BasePathMixin): uri the media comes from in URI hierarchy. ex.: http://example.com/path/to """ + base_uri: str | None + uri: str | None + type: str | None + group_id: str | None + language: str | None + name: str | None + default: str | None + autoselect: str | None + forced: str | None + assoc_language: str | None + instream_id: str | None + characteristics: str | None + channels: str | None + stable_rendition_id: str | None + extras: dict[str, Any] + def __init__( - self, - uri=None, - type=None, - group_id=None, - language=None, - name=None, - default=None, - autoselect=None, - forced=None, - characteristics=None, - channels=None, - stable_rendition_id=None, - assoc_language=None, - instream_id=None, - base_uri=None, - **extras, + self, + uri: str | None = None, + type: str | None = None, + group_id: str | None = None, + language: str | None = None, + name: str | None = None, + default: str | None = None, + autoselect: str | None = None, + forced: str | None = None, + characteristics: str | None = None, + channels: str | None = None, + stable_rendition_id: str | None = None, + assoc_language: str | None = None, + instream_id: str | None = None, + base_uri: str | None = None, + **extras, ): self.base_uri = base_uri self.uri = uri @@ -1240,7 +1347,7 @@ def __init__( self.stable_rendition_id = stable_rendition_id self.extras = extras - def dumps(self): + def dumps(self) -> str: media_out = [] if self.uri: @@ -1272,36 +1379,40 @@ def dumps(self): return "#EXT-X-MEDIA:" + ",".join(media_out) - def __str__(self): + def __str__(self) -> str: return self.dumps() -class TagList(list): - def __str__(self): +class TagList[T](list[T]): + def __str__(self) -> str: output = [str(tag) for tag in self] return "\n".join(output) -class MediaList(TagList, GroupedBasePathMixin): +class MediaList(TagList[Media], GroupedBasePathMixin[Media]): @property - def uri(self): + def uri(self) -> list[str | None]: return [media.uri for media in self] -class PlaylistList(TagList, GroupedBasePathMixin): - pass +class PlaylistList(TagList[Playlist], GroupedBasePathMixin[Playlist]): ... -class SessionDataList(TagList): - pass +class SessionDataList(TagList[SessionData]): ... class Start: - def __init__(self, time_offset, precise=None): + time_offset: float + precise: Literal["YES", "NO"] | None + + def __init__( + self, time_offset: float, precise: Literal["YES", "NO"] | None = None + ) -> None: + self.time_offset = float(time_offset) self.precise = precise - def __str__(self): + def __str__(self) -> str: output = ["TIME-OFFSET=" + str(self.time_offset)] if self.precise and self.precise in ["YES", "NO"]: output.append("PRECISE=" + str(self.precise)) @@ -1310,13 +1421,24 @@ def __str__(self): class RenditionReport(BasePathMixin): - def __init__(self, base_uri, uri, last_msn=None, last_part=None): + base_uri: str | None + uri: str | None + last_msn: int + last_part: int | None + + def __init__( + self, + base_uri: str | None, + uri: str | None, + last_msn: int, + last_part: int | None = None, + ) -> None: self.base_uri = base_uri self.uri = uri self.last_msn = last_msn self.last_part = last_part - def dumps(self): + def dumps(self) -> str: report = [] report.append("URI=" + quoted(self.uri)) if self.last_msn is not None: @@ -1326,35 +1448,41 @@ def dumps(self): return "#EXT-X-RENDITION-REPORT:" + ",".join(report) - def __str__(self): + def __str__(self) -> str: return self.dumps() -class RenditionReportList(list, GroupedBasePathMixin): - def __str__(self): +class RenditionReportList(list[RenditionReport], GroupedBasePathMixin[RenditionReport]): + def __str__(self) -> str: output = [str(report) for report in self] return "\n".join(output) class ServerControl: + can_skip_until: float | None + can_block_reload: str | None + hold_back: float | None + part_hold_back: float | None + can_skip_dateranges: str | None + def __init__( - self, - can_skip_until=None, - can_block_reload=None, - hold_back=None, - part_hold_back=None, - can_skip_dateranges=None, - ): + self, + can_skip_until: float | None = None, + can_block_reload: str | None = None, + hold_back: float | None = None, + part_hold_back: float | None = None, + can_skip_dateranges: str | None = None, + ) -> None: self.can_skip_until = can_skip_until self.can_block_reload = can_block_reload self.hold_back = hold_back self.part_hold_back = part_hold_back self.can_skip_dateranges = can_skip_dateranges - def __getitem__(self, item): + def __getitem__(self, item: str) -> str | float | None: return getattr(self, item) - def dumps(self): + def dumps(self) -> str: ctrl = [] if self.can_block_reload: ctrl.append("CAN-BLOCK-RELOAD=%s" % self.can_block_reload) @@ -1373,16 +1501,21 @@ def dumps(self): return "#EXT-X-SERVER-CONTROL:" + ",".join(ctrl) - def __str__(self): + def __str__(self) -> str: return self.dumps() class Skip: - def __init__(self, skipped_segments, recently_removed_dateranges=None): + skipped_segments: int | None + recently_removed_dateranges: str | None + + def __init__( + self, skipped_segments: int, recently_removed_dateranges: str | None = None + ) -> None: self.skipped_segments = skipped_segments self.recently_removed_dateranges = recently_removed_dateranges - def dumps(self): + def dumps(self) -> str: skip = [] skip.append("SKIPPED-SEGMENTS=%s" % self.skipped_segments) if self.recently_removed_dateranges is not None: @@ -1393,35 +1526,48 @@ def dumps(self): return "#EXT-X-SKIP:" + ",".join(skip) - def __str__(self): + def __str__(self) -> str: return self.dumps() class PartInformation: - def __init__(self, part_target=None): + part_target: float | None + + def __init__(self, part_target: float | None = None) -> None: self.part_target = part_target - def dumps(self): + def dumps(self) -> str: return "#EXT-X-PART-INF:PART-TARGET=%s" % number_to_string(self.part_target) - def __str__(self): + def __str__(self) -> str: return self.dumps() class PreloadHint(BasePathMixin): + hint_type: str | None + base_uri: str | None + uri: str | None + byterange_start: int | None + byterange_length: int | None + def __init__( - self, type, base_uri, uri, byterange_start=None, byterange_length=None - ): + self, + type: str | None, + base_uri: str | None, + uri: str | None, + byterange_start: int | None = None, + byterange_length: int | None = None, + ) -> None: self.hint_type = type self.base_uri = base_uri self.uri = uri self.byterange_start = byterange_start self.byterange_length = byterange_length - def __getitem__(self, item): + def __getitem__(self, item: str) -> str | int | None: return getattr(self, item) - def dumps(self): + def dumps(self) -> str: hint = [] hint.append("TYPE=" + self.hint_type) hint.append("URI=" + quoted(self.uri)) @@ -1432,18 +1578,29 @@ def dumps(self): return "#EXT-X-PRELOAD-HINT:" + ",".join(hint) - def __str__(self): + def __str__(self) -> str: return self.dumps() class SessionData: - def __init__(self, data_id, value=None, uri=None, language=None): + data_id: str + value: str | None + uri: str | None + language: str | None + + def __init__( + self, + data_id: str, + value: str | None = None, + uri: str | None = None, + language: str | None = None, + ) -> None: self.data_id = data_id self.value = value self.uri = uri self.language = language - def dumps(self): + def dumps(self) -> str: session_data_out = ["DATA-ID=" + quoted(self.data_id)] if self.value: @@ -1455,31 +1612,56 @@ def dumps(self): return "#EXT-X-SESSION-DATA:" + ",".join(session_data_out) - def __str__(self): + def __str__(self) -> str: return self.dumps() -class DateRangeList(TagList): - pass +class DateRangeList(TagList[DateRange]): ... class DateRange: - def __init__(self, **kwargs): - self.id = kwargs["id"] - self.start_date = kwargs.get("start_date") - self.class_ = kwargs.get("class") - self.end_date = kwargs.get("end_date") - self.duration = kwargs.get("duration") - self.planned_duration = kwargs.get("planned_duration") - self.scte35_cmd = kwargs.get("scte35_cmd") - self.scte35_out = kwargs.get("scte35_out") - self.scte35_in = kwargs.get("scte35_in") - self.end_on_next = kwargs.get("end_on_next") + id: str + start_date: str | None + class_: str | None + end_date: str | None + duration: float | None + planned_duration: float | None + scte35_cmd: str | None + scte35_out: str | None + scte35_in: str | None + end_on_next: Any + x_client_attrs: list[tuple[str, str]] + + def __init__( + self, + *, + id: str, + start_date: str | None = None, + class_: str | None = None, # actually passing as `class` argument + end_date: str | None = None, + duration: float | None = None, + planned_duration: float | None = None, + scte35_cmd: str | None = None, + scte35_out: str | None = None, + scte35_in: str | None = None, + end_on_next=None, + **kwargs: str, # for arguments with `x_` prefix + ) -> None: + self.id = id + self.start_date = start_date + self.class_ = class_ + self.end_date = end_date + self.duration = duration + self.planned_duration = planned_duration + self.scte35_cmd = scte35_cmd + self.scte35_out = scte35_out + self.scte35_in = scte35_in + self.end_on_next = end_on_next self.x_client_attrs = [ - (attr, kwargs.get(attr)) for attr in kwargs if attr.startswith("x_") + (attr, kwargs[attr]) for attr in kwargs if attr.startswith("x_") ] - def dumps(self): + def dumps(self) -> str: daterange = [] daterange.append("ID=" + quoted(self.id)) @@ -1514,17 +1696,26 @@ def dumps(self): return "#EXT-X-DATERANGE:" + ",".join(daterange) - def __str__(self): + def __str__(self) -> str: return self.dumps() class ContentSteering(BasePathMixin): - def __init__(self, base_uri, server_uri, pathway_id=None): + base_uri: str | None + uri: str | None + pathway_id: str | None + + def __init__( + self, + base_uri: str | None, + server_uri: str | None, + pathway_id: str | None = None, + ) -> None: self.base_uri = base_uri self.uri = server_uri self.pathway_id = pathway_id - def dumps(self): + def dumps(self) -> str: steering = [] steering.append("SERVER-URI=" + quoted(self.uri)) @@ -1533,11 +1724,11 @@ def dumps(self): return "#EXT-X-CONTENT-STEERING:" + ",".join(steering) - def __str__(self): + def __str__(self) -> str: return self.dumps() -class ImagePlaylist(BasePathMixin): +class ImagePlaylist(Playlist): """ ImagePlaylist object representing a link to a variant M3U8 image playlist with a specific bitrate. @@ -1550,7 +1741,17 @@ class ImagePlaylist(BasePathMixin): More info: https://github.com/image-media-playlist/spec/blob/master/image_media_playlist_v0_4.pdf """ - def __init__(self, base_uri, uri, image_stream_info): + uri: str | None + base_uri: str | None + image_stream_info: StreamInfo + + def __init__( + self, + base_uri: str | None, + uri: str | None, + image_stream_info: Mapping[str, Any], + ) -> None: + self.uri = uri self.base_uri = base_uri @@ -1580,7 +1781,7 @@ def __init__(self, base_uri, uri, image_stream_info): stable_variant_id=image_stream_info.get("stable_variant_id"), ) - def __str__(self): + def __str__(self) -> str: image_stream_inf = [] if self.image_stream_info.program_id: image_stream_inf.append("PROGRAM-ID=%d" % self.image_stream_info.program_id) @@ -1592,9 +1793,9 @@ def __str__(self): ) if self.image_stream_info.resolution: res = ( - str(self.image_stream_info.resolution[0]) - + "x" - + str(self.image_stream_info.resolution[1]) + str(self.image_stream_info.resolution[0]) + + "x" + + str(self.image_stream_info.resolution[1]) ) image_stream_inf.append("RESOLUTION=" + res) if self.image_stream_info.codecs: @@ -1627,12 +1828,17 @@ class Tiles(BasePathMixin): duration attribute from EXT-X-TILES tag """ - def __init__(self, resolution, layout, duration): + uri: str | None + resolution: Any + layout: Any + duration: Any + + def __init__(self, resolution, layout, duration) -> None: self.resolution = resolution self.layout = layout self.duration = duration - def dumps(self): + def dumps(self) -> str: tiles = [] tiles.append("RESOLUTION=" + self.resolution) tiles.append("LAYOUT=" + self.layout) @@ -1644,36 +1850,39 @@ def __str__(self): return self.dumps() -def find_keys(keys, keylist) -> tuple["Key"]: +def find_keys( + keys: Mapping[int, Any] | None, keylist: Sequence[tuple[Key, ...]] +) -> tuple["Key", ...]: if not keys: - return tuple() + return () for key_objects in keylist: if len(key_objects) != len(keys): continue - match_keys = [] + match_keys: list[Key] = [] for i in range(len(keys)): if ( - key_objects[i].uri == keys[i].get("uri", None) - and key_objects[i].method == keys[i].get("method", "NONE") - and key_objects[i].iv == keys[i].get("iv", None) + key_objects[i].uri == keys[i].get("uri", None) + and key_objects[i].method == keys[i].get("method", "NONE") + and key_objects[i].iv == keys[i].get("iv", None) ): match_keys.append(key_objects[i]) + if len(match_keys) == len(keys): return tuple(match_keys) raise KeyError("No matching keys found") -def denormalize_attribute(attribute): +def denormalize_attribute(attribute: str) -> str: return attribute.replace("_", "-").upper() -def quoted(string): +def quoted(string: str) -> str: return '"%s"' % string -def number_to_string(number): +def number_to_string(number: str | float | decimal.Decimal) -> str: with decimal.localcontext() as ctx: ctx.prec = 20 # set floating point precision d = decimal.Decimal(str(number)) diff --git a/runtests b/runtests index 7d55bd76..72ed6d0d 100755 --- a/runtests +++ b/runtests @@ -3,12 +3,12 @@ test_server_stdout=tests/server.stdout function install_deps { - pip install -r requirements-dev.txt + uv sync --locked } function start_server { rm -f ${test_server_stdout} - python tests/m3u8server.py >${test_server_stdout} 2>&1 & + uv run tests/m3u8server.py >${test_server_stdout} 2>&1 & } function stop_server { @@ -17,7 +17,7 @@ function stop_server { } function run { - PYTHONPATH=. py.test -vv --cov-report term-missing --cov m3u8 tests/ + uv run pytest -vv --cov-report term-missing --cov m3u8 tests/ } function main { diff --git a/tests/m3u8server.py b/tests/m3u8server.py index 1d135e7f..ecba6e14 100644 --- a/tests/m3u8server.py +++ b/tests/m3u8server.py @@ -4,43 +4,47 @@ # Test server to deliver stubed M3U8s -from os.path import dirname, abspath, join +import time +from pathlib import Path -from bottle import route, run, response, redirect import bottle -import time +from bottle import os, redirect, response, route, run -playlists = abspath(join(dirname(__file__), "playlists")) +playlists = Path(__file__).parent / "playlists" @route("/path/to/redirect_me") -def redirect_route(): +def redirect_route() -> None: redirect("/simple.m3u8") @route("/simple.m3u8") -def simple(): +def simple() -> str: response.set_header("Content-Type", "application/vnd.apple.mpegurl") return m3u8_file("simple-playlist.m3u8") @route("/timeout_simple.m3u8") -def timeout_simple(): +def timeout_simple() -> str: time.sleep(5) response.set_header("Content-Type", "application/vnd.apple.mpegurl") return m3u8_file("simple-playlist.m3u8") @route("/path/to/relative-playlist.m3u8") -def relative_playlist(): +def relative_playlist() -> str: response.set_header("Content-Type", "application/vnd.apple.mpegurl") return m3u8_file("relative-playlist.m3u8") -def m3u8_file(filename): - with open(join(playlists, filename)) as fileobj: - return fileobj.read().strip() +def m3u8_file(filename: os.PathLike[str] | str) -> str: + return (playlists / filename).read_text().strip() + + +def run_server() -> None: + bottle.debug = True + run(host="localhost", port=8112) -bottle.debug = True -run(host="localhost", port=8112) +if __name__ == "__main__": + run_server() diff --git a/tests/playlists.py b/tests/playlists.py index b2435d2e..7d27951b 100755 --- a/tests/playlists.py +++ b/tests/playlists.py @@ -2,7 +2,15 @@ # Use of this source code is governed by a MIT License # license that can be found in the LICENSE file. -from os.path import abspath, dirname, join + +from pathlib import Path + +_PLAYLISTS = Path(__file__).parent / "playlists" +SIMPLE_PLAYLIST_FILENAME = _PLAYLISTS / "simple-playlist.m3u8" +RELATIVE_PLAYLIST_FILENAME = _PLAYLISTS / "relative-playlist.m3u8" +CUE_OUT_PLAYLIST_FILENAME = _PLAYLISTS / "cue_out.m3u8" + +del Path TEST_HOST = "http://localhost:8112" @@ -54,9 +62,6 @@ #EXT-X-ENDLIST """ -SIMPLE_PLAYLIST_FILENAME = abspath( - join(dirname(__file__), "playlists/simple-playlist.m3u8") -) SIMPLE_PLAYLIST_URI = TEST_HOST + "/simple.m3u8" TIMEOUT_SIMPLE_PLAYLIST_URI = TEST_HOST + "/timeout_simple.m3u8" @@ -1173,13 +1178,9 @@ #EXT-X-PRELOAD-HINT:TYPE=PART,URI="fs271.mp4",BYTERANGE-START=61000,BYTERANGE-LENGTH=20000 """ -RELATIVE_PLAYLIST_FILENAME = abspath( - join(dirname(__file__), "playlists/relative-playlist.m3u8") -) RELATIVE_PLAYLIST_URI = TEST_HOST + "/path/to/relative-playlist.m3u8" -CUE_OUT_PLAYLIST_FILENAME = abspath(join(dirname(__file__), "playlists/cue_out.m3u8")) CUE_OUT_PLAYLIST_URI = TEST_HOST + "/path/to/cue_out.m3u8" @@ -1586,5 +1587,3 @@ #EXTINF:9.600, C:\HLS Video\test1.ts """ - -del abspath, dirname, join From f231cb8d7df080b21ff265ad8a86451c2077954f Mon Sep 17 00:00:00 2001 From: NTFSvolume <172021377+NTFSvolume@users.noreply.github.com> Date: Tue, 14 Jul 2026 00:36:52 -0500 Subject: [PATCH 2/2] refactor: more type hints --- m3u8/model.py | 4 +- m3u8/parser.py | 236 +++++++++++++++++++++------------ m3u8/version_matching.py | 51 +++---- m3u8/version_matching_rules.py | 85 ++++++------ 4 files changed, 223 insertions(+), 153 deletions(-) diff --git a/m3u8/model.py b/m3u8/model.py index 3fffafd8..f4d8a001 100644 --- a/m3u8/model.py +++ b/m3u8/model.py @@ -8,7 +8,7 @@ from typing import Any, ClassVar, Literal from m3u8.mixins import BasePathMixin, GroupedBasePathMixin -from m3u8.parser import format_date_time, parse +from m3u8.parser import parse from m3u8.protocol import ( ext_oatcls_scte35, ext_x_asset, @@ -695,7 +695,7 @@ def dumps( if self.program_date_time: output.append( "#EXT-X-PROGRAM-DATE-TIME:%s\n" - % format_date_time(self.program_date_time, timespec=timespec) + % self.program_date_time.isoformat(timespec=timespec) ) if len(self.dateranges): diff --git a/m3u8/parser.py b/m3u8/parser.py index 4399eaba..e1f7cfa4 100644 --- a/m3u8/parser.py +++ b/m3u8/parser.py @@ -2,16 +2,12 @@ # Use of this source code is governed by a MIT License # license that can be found in the LICENSE file. +import dataclasses import itertools import re +from collections.abc import Callable from datetime import datetime, timedelta - -try: - from backports.datetime_fromisoformat import MonkeyPatch - - MonkeyPatch.patch_fromisoformat() -except ImportError: - pass +from typing import Any, TypedDict, Unpack from m3u8 import protocol, version_matching @@ -19,31 +15,73 @@ http://tools.ietf.org/html/draft-pantos-http-live-streaming-08#section-3.2 http://stackoverflow.com/questions/2785755/how-to-split-but-ignore-separators-in-quoted-strings-in-python """ -ATTRIBUTELISTPATTERN = re.compile(r"""((?:[^,"']|"[^"]*"|'[^']*')+)""") +ATTRIBUTE_LIST_PATTERN = re.compile(r"""((?:[^,"']|"[^"]*"|'[^']*')+)""") -def cast_date_time(value): +def cast_date_time(value: str) -> datetime: return datetime.fromisoformat(value) -def format_date_time(value, **kwargs): - return value.isoformat(**kwargs) - - +@dataclasses.dataclass(slots=True) class ParseError(Exception): - def __init__(self, lineno, line): - self.lineno = lineno - self.line = line + lineno: int + line: str - def __str__(self): + def __str__(self) -> str: return "Syntax error in manifest on line %d: %s" % (self.lineno, self.line) -def parse(content, strict=False, custom_tags_parser=None): +class M3U8Data(TypedDict): + media_sequence: int + is_variant: bool + is_endlist: bool + is_i_frames_only: bool + is_independent_segments: bool + is_images_only: bool + playlist_type: str | None + playlists: list[Any] + segments: list[Any] + iframe_playlists: list[Any] + image_playlists: list[Any] + tiles: list[Any] + media: list[Any] + keys: list[Any] + rendition_reports: list[Any] + skip: dict[str, Any] + part_inf: dict[str, Any] + session_data: list[Any] + session_keys: list[Any] + segment_map: list[Any] + + +class State(TypedDict): + expect_segment: bool + expect_playlist: bool + current_keys: list[Any] + current_segment_map: None + continuous_key: bool + + +class _ParseKwargs(TypedDict): + line: str + lineno: int + data: M3U8Data + state: State + strict: bool + + +type _CustomTagsParser = Callable[[str, int, M3U8Data, State], object] + + +def parse( + content: str, + strict: bool = False, + custom_tags_parser: _CustomTagsParser | None = None, +) -> M3U8Data: """ Given a M3U8 playlist content returns a dictionary with all data found """ - data = { + data: M3U8Data = { "media_sequence": 0, "is_variant": False, "is_endlist": False, @@ -66,7 +104,7 @@ def parse(content, strict=False, custom_tags_parser=None): "segment_map": [], } - state = { + state: State = { "expect_segment": False, "expect_playlist": False, "current_keys": [], @@ -83,7 +121,7 @@ def parse(content, strict=False, custom_tags_parser=None): for lineno, line in enumerate(lines, 1): line = line.strip() - parse_kwargs = { + parse_kwargs: _ParseKwargs = { "line": line, "lineno": lineno, "data": data, @@ -247,8 +285,8 @@ def parse(content, strict=False, custom_tags_parser=None): return data -def _parse_key(line, data, state, **kwargs): - params = ATTRIBUTELISTPATTERN.split(line.replace(protocol.ext_x_key + ":", ""))[ +def _parse_key(line: str, data: M3U8Data, state: State, **kwargs: object) -> None: + params = ATTRIBUTE_LIST_PATTERN.split(line.replace(protocol.ext_x_key + ":", ""))[ 1::2 ] key = {} @@ -265,7 +303,9 @@ def _parse_key(line, data, state, **kwargs): # data["keys"].append(key) -def _parse_extinf(line, state, lineno, strict, **kwargs): +def _parse_extinf( + line: str, state: State, lineno: int, strict: bool, **kwargs: object +) -> None: chunks = line.replace(protocol.extinf + ":", "").split(",", 1) if len(chunks) == 2: duration, title = chunks @@ -282,7 +322,7 @@ def _parse_extinf(line, state, lineno, strict, **kwargs): state["expect_segment"] = True -def _parse_ts_chunk(line, data, state, **kwargs): +def _parse_ts_chunk(line: str, data: M3U8Data, state: State, **kwargs: object) -> None: segment = state.pop("segment") if state.get("program_date_time"): segment["program_date_time"] = state.pop("program_date_time") @@ -317,10 +357,17 @@ def _parse_ts_chunk(line, data, state, **kwargs): state["expect_segment"] = False -def _parse_attribute_list(prefix, line, attribute_parser, default_parser=None): - params = ATTRIBUTELISTPATTERN.split(line.replace(prefix + ":", ""))[1::2] +def _parse_attribute_list[T: str | float | int]( + prefix: str, + line: str, + attribute_parser: dict[str, Callable[[str], T]], + default_parser: Callable[[str], T] | None = None, +) -> dict[str, str | T]: + params: list[str] = ATTRIBUTE_LIST_PATTERN.split(line.replace(prefix + ":", ""))[ + 1::2 + ] - attributes = {} + attributes: dict[str, T | str] = {} if not line.startswith(prefix + ":"): return attributes @@ -343,7 +390,9 @@ def _parse_attribute_list(prefix, line, attribute_parser, default_parser=None): return attributes -def _parse_stream_inf(line, data, state, **kwargs): +def _parse_stream_inf( + line: str, data: M3U8Data, state: State, **kwargs: object +) -> None: state["expect_playlist"] = True data["is_variant"] = True data["media_sequence"] = None @@ -366,7 +415,7 @@ def _parse_stream_inf(line, data, state, **kwargs): ) -def _parse_i_frame_stream_inf(line, data, **kwargs): +def _parse_i_frame_stream_inf(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = remove_quotes_parser( "codecs", "uri", "pathway_id", "stable_variant_id" ) @@ -385,7 +434,7 @@ def _parse_i_frame_stream_inf(line, data, **kwargs): data["iframe_playlists"].append(iframe_playlist) -def _parse_image_stream_inf(line, data, **kwargs): +def _parse_image_stream_inf(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = remove_quotes_parser( "codecs", "uri", "pathway_id", "stable_variant_id" ) @@ -404,11 +453,11 @@ def _parse_image_stream_inf(line, data, **kwargs): data["image_playlists"].append(image_playlist) -def _parse_is_images_only(line, data, **kwargs): +def _parse_is_images_only(line: str, data: M3U8Data, **kwargs: object) -> None: data["is_images_only"] = True -def _parse_tiles(line, data, state, **kwargs): +def _parse_tiles(line: str, data: M3U8Data, state: State, **kwargs: object) -> None: attribute_parser = remove_quotes_parser("uri") attribute_parser["resolution"] = str attribute_parser["layout"] = str @@ -417,7 +466,7 @@ def _parse_tiles(line, data, state, **kwargs): data["tiles"].append(tiles_info) -def _parse_media(line, data, **kwargs): +def _parse_media(line: str, data: M3U8Data, **kwargs: object) -> None: quoted = remove_quotes_parser( "uri", "group_id", @@ -435,38 +484,42 @@ def _parse_media(line, data, **kwargs): data["media"].append(media) -def _parse_variant_playlist(line, data, state, **kwargs): +def _parse_variant_playlist( + line: str, data: M3U8Data, state: State, **kwargs: object +) -> None: playlist = {"uri": line, "stream_info": state.pop("stream_info")} data["playlists"].append(playlist) state["expect_playlist"] = False -def _parse_bitrate(state, **kwargs): +def _parse_bitrate(state: State, **kwargs: object) -> None: if "segment" not in state: state["segment"] = {} state["segment"]["bitrate"] = _parse_simple_parameter(cast_to=int, **kwargs) -def _parse_byterange(line, state, **kwargs): +def _parse_byterange(line: str, state: State, **kwargs: object) -> None: if "segment" not in state: state["segment"] = {} state["segment"]["byterange"] = line.replace(protocol.ext_x_byterange + ":", "") state["expect_segment"] = True -def _parse_targetduration(**parse_kwargs): +def _parse_targetduration(**parse_kwargs: Unpack[_ParseKwargs]) -> int: return _parse_simple_parameter(cast_to=int, **parse_kwargs) -def _parse_media_sequence(**parse_kwargs): +def _parse_media_sequence(**parse_kwargs: Unpack[_ParseKwargs]) -> int: return _parse_simple_parameter(cast_to=int, **parse_kwargs) -def _parse_discontinuity_sequence(**parse_kwargs): +def _parse_discontinuity_sequence(**parse_kwargs: Unpack[_ParseKwargs]) -> int: return _parse_simple_parameter(cast_to=int, **parse_kwargs) -def _parse_program_date_time(line, state, data, **parse_kwargs): +def _parse_program_date_time( + line: str, state: State, data: M3U8Data, **parse_kwargs +) -> None: _, program_date_time = _parse_simple_parameter_raw_value( line, cast_to=cast_date_time, **parse_kwargs ) @@ -476,80 +529,94 @@ def _parse_program_date_time(line, state, data, **parse_kwargs): state["program_date_time"] = program_date_time -def _parse_discontinuity(state, **parse_kwargs): +def _parse_discontinuity(state: State, **parse_kwargs: object) -> None: state["discontinuity"] = True -def _parse_cue_in(state, **parse_kwargs): +def _parse_cue_in(state: State, **parse_kwargs: object) -> None: state["cue_in"] = True -def _parse_cue_span(state, **parse_kwargs): +def _parse_cue_span(state: State, **parse_kwargs: object) -> None: state["cue_out"] = True -def _parse_version(**parse_kwargs): +def _parse_version(**parse_kwargs: Unpack[_ParseKwargs]) -> int: return _parse_simple_parameter(cast_to=int, **parse_kwargs) -def _parse_allow_cache(**parse_kwargs): +def _parse_allow_cache(**parse_kwargs: Unpack[_ParseKwargs]) -> str: return _parse_simple_parameter(cast_to=str, **parse_kwargs) -def _parse_playlist_type(line, data, **kwargs): +def _parse_playlist_type(line: str, data: M3U8Data, **kwargs: object) -> str: return _parse_simple_parameter(line, data) -def _parse_x_map(line, data, state, **kwargs): +def _parse_x_map(line: str, data: M3U8Data, state: State, **kwargs: object) -> None: quoted_parser = remove_quotes_parser("uri", "byterange") segment_map_info = _parse_attribute_list(protocol.ext_x_map, line, quoted_parser) state["current_segment_map"] = segment_map_info data["segment_map"].append(segment_map_info) -def _parse_start(line, data, **kwargs): +def _parse_start(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = {"time_offset": lambda x: float(x)} start_info = _parse_attribute_list(protocol.ext_x_start, line, attribute_parser) data["start"] = start_info -def _parse_gap(state, **kwargs): +def _parse_gap(state: State, **kwargs: object) -> None: state["gap"] = True -def _parse_simple_parameter_raw_value(line, cast_to=str, normalize=False, **kwargs): - param, value = line.split(":", 1) +def _parse_simple_parameter_raw_value[T: str | float | int | datetime]( + line: str, + cast_to: Callable[[str], T] = str, + normalize: bool = False, + **kwargs: object, +) -> tuple[str, T]: + param, _, value = line.partition(":") param = normalize_attribute(param.replace("#EXT-X-", "")) if normalize: value = value.strip().lower() return param, cast_to(value) -def _parse_and_set_simple_parameter_raw_value( - line, data, cast_to=str, normalize=False, **kwargs -): +def _parse_and_set_simple_parameter_raw_value[T: str | float | int]( + line: str, + data: M3U8Data, + cast_to: Callable[[str], T] = str, + normalize: bool = False, + **kwargs: object, +) -> T: param, value = _parse_simple_parameter_raw_value(line, cast_to, normalize) data[param] = value - return data[param] + return value -def _parse_simple_parameter(line, data, cast_to=str, **kwargs): +def _parse_simple_parameter[T: str | float | int]( + line: str, + data: M3U8Data, + cast_to: Callable[[str], T] = str, + **kwargs: object, +) -> T: return _parse_and_set_simple_parameter_raw_value(line, data, cast_to, True) -def _parse_i_frames_only(data, **kwargs): +def _parse_i_frames_only(data: M3U8Data, **kwargs: object) -> None: data["is_i_frames_only"] = True -def _parse_is_independent_segments(data, **kwargs): +def _parse_is_independent_segments(data: M3U8Data, **kwargs: object) -> None: data["is_independent_segments"] = True -def _parse_endlist(data, **kwargs): +def _parse_endlist(data: M3U8Data, **kwargs: object) -> None: data["is_endlist"] = True -def _parse_cueout_cont(line, state, **kwargs): +def _parse_cueout_cont(line: str, state: State, **kwargs: object) -> None: state["cue_out"] = True elements = line.split(":", 1) @@ -564,8 +631,7 @@ def _parse_cueout_cont(line, state, **kwargs): ) # EXT-X-CUE-OUT-CONT:2.436/120 style - progress = cue_info.get("") - if progress: + if progress := cue_info.get(""): progress_parts = progress.split("/", 1) if len(progress_parts) == 1: state["current_cue_out_duration"] = progress_parts[0] @@ -586,7 +652,7 @@ def _parse_cueout_cont(line, state, **kwargs): state["current_cue_out_elapsedtime"] = elapsedtime -def _parse_cueout(line, state, **kwargs): +def _parse_cueout(line: str, state: State, **kwargs: object) -> None: state["cue_out_start"] = True state["cue_out"] = True if "DURATION" in line.upper(): @@ -609,7 +675,7 @@ def _parse_cueout(line, state, **kwargs): state["current_cue_out_duration"] = cue_out_duration -def _parse_server_control(line, data, **kwargs): +def _parse_server_control(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = { "can_block_reload": str, "hold_back": lambda x: float(x), @@ -623,7 +689,7 @@ def _parse_server_control(line, data, **kwargs): ) -def _parse_part_inf(line, data, **kwargs): +def _parse_part_inf(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = {"part_target": lambda x: float(x)} data["part_inf"] = _parse_attribute_list( @@ -631,7 +697,7 @@ def _parse_part_inf(line, data, **kwargs): ) -def _parse_rendition_report(line, data, **kwargs): +def _parse_rendition_report(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = remove_quotes_parser("uri") attribute_parser["last_msn"] = int attribute_parser["last_part"] = int @@ -643,7 +709,7 @@ def _parse_rendition_report(line, data, **kwargs): data["rendition_reports"].append(rendition_report) -def _parse_part(line, state, **kwargs): +def _parse_part(line: str, state: State, **kwargs) -> None: attribute_parser = remove_quotes_parser("uri") attribute_parser["duration"] = lambda x: float(x) attribute_parser["independent"] = str @@ -669,21 +735,21 @@ def _parse_part(line, state, **kwargs): segment["parts"].append(part) -def _parse_skip(line, data, **parse_kwargs): +def _parse_skip(line: str, data: M3U8Data, **parse_kwargs: object) -> None: attribute_parser = remove_quotes_parser("recently_removed_dateranges") attribute_parser["skipped_segments"] = int data["skip"] = _parse_attribute_list(protocol.ext_x_skip, line, attribute_parser) -def _parse_session_data(line, data, **kwargs): +def _parse_session_data(line: str, data: M3U8Data, **kwargs: object) -> None: quoted = remove_quotes_parser("data_id", "value", "uri", "language") session_data = _parse_attribute_list(protocol.ext_x_session_data, line, quoted) data["session_data"].append(session_data) -def _parse_session_key(line, data, **kwargs): - params = ATTRIBUTELISTPATTERN.split( +def _parse_session_key(line: str, data: M3U8Data, **kwargs: object) -> None: + params = ATTRIBUTE_LIST_PATTERN.split( line.replace(protocol.ext_x_session_key + ":", "") )[1::2] key = {} @@ -693,7 +759,7 @@ def _parse_session_key(line, data, **kwargs): data["session_keys"].append(key) -def _parse_preload_hint(line, data, **kwargs): +def _parse_preload_hint(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = remove_quotes_parser("uri") attribute_parser["type"] = str attribute_parser["byterange_start"] = int @@ -704,7 +770,7 @@ def _parse_preload_hint(line, data, **kwargs): ) -def _parse_daterange(line, state, **kwargs): +def _parse_daterange(line: str, state: State, **kwargs: object) -> None: attribute_parser = remove_quotes_parser("id", "class", "start_date", "end_date") attribute_parser["duration"] = float attribute_parser["planned_duration"] = float @@ -721,7 +787,7 @@ def _parse_daterange(line, state, **kwargs): state["dateranges"].append(parsed) -def _parse_content_steering(line, data, **kwargs): +def _parse_content_steering(line: str, data: M3U8Data, **kwargs: object) -> None: attribute_parser = remove_quotes_parser("server_uri", "pathway_id") data["content_steering"] = _parse_attribute_list( @@ -729,13 +795,13 @@ def _parse_content_steering(line, data, **kwargs): ) -def _parse_oatcls_scte35(line, state, **kwargs): +def _parse_oatcls_scte35(line: str, state: State, **kwargs: object) -> None: scte35_cue = line.split(":", 1)[1] state["current_cue_out_oatcls_scte35"] = scte35_cue state["current_cue_out_scte35"] = scte35_cue -def _parse_asset(line, state, **kwargs): +def _parse_asset(line: str, state: State, **kwargs: object) -> None: # EXT-X-ASSET attribute values may or may not be quoted, and need to be URL-encoded. # They are preserved as-is here to prevent loss of information. state["asset_metadata"] = _parse_attribute_list( @@ -743,15 +809,15 @@ def _parse_asset(line, state, **kwargs): ) -def string_to_lines(string): +def string_to_lines(string: str) -> list[str]: return string.strip().splitlines() -def remove_quotes_parser(*attrs): - return dict(zip(attrs, itertools.repeat(remove_quotes))) +def remove_quotes_parser(*attrs: str) -> dict[str, Callable[[str], str | int | float]]: + return dict(zip(attrs, itertools.repeat(remove_quotes), strict=False)) -def remove_quotes(string): +def remove_quotes(string: str) -> str: """ Remove quotes from string. @@ -767,11 +833,13 @@ def remove_quotes(string): return string -def normalize_attribute(attribute): +def normalize_attribute(attribute: str) -> str: return attribute.replace("-", "_").lower().strip() -def get_segment_custom_value(state, key, default=None): +def get_segment_custom_value[T]( + state: State, key: str, default: T | None = None +) -> T | None: """ Helper function for getting custom values for Segment Are useful with custom_tags_parser @@ -783,7 +851,7 @@ def get_segment_custom_value(state, key, default=None): return state["segment"]["custom_parser_values"].get(key, default) -def save_segment_custom_value(state, key, value): +def save_segment_custom_value(state: State, key: str, value: Any) -> None: """ Helper function for saving custom values for Segment Are useful with custom_tags_parser diff --git a/m3u8/version_matching.py b/m3u8/version_matching.py index 00846366..ae2a4678 100644 --- a/m3u8/version_matching.py +++ b/m3u8/version_matching.py @@ -1,43 +1,47 @@ +from collections.abc import Generator, Iterable + from m3u8 import protocol -from m3u8.version_matching_rules import VersionMatchingError, available_rules +from m3u8.version_matching_rules import RULES, VersionMatchingError -def get_version(file_lines: list[str]): +def get_version(file_lines: Iterable[str]) -> float | None: for line in file_lines: if line.startswith(protocol.ext_x_version): - version = line.split(":")[1] + version = line.partition(":")[1] return float(version) return None -def validate_multiple_version_decl(file_lines: list[str]): + +def validate_multiple_version_decl( + file_lines: Iterable[str], errors: list[VersionMatchingError] +) -> Generator[str]: version_number_line_ctr = 0 - errors = [] for number, line in enumerate(file_lines): if line.startswith(protocol.ext_x_version): - version_number_line_ctr = version_number_line_ctr + 1 + version_number_line_ctr += 1 if version_number_line_ctr > 1: - errors.append(VersionMatchingError( - line_number=number, - line=line, - description="There are multiple version declarations in the file.", - how_to_fix="Remove all extra version declarations." - )) - return errors + errors.append( + VersionMatchingError( + line_number=number, + line=line, + description="There are multiple version declarations in the file.", + how_to_fix="Remove all extra version declarations.", + ) + ) + yield line + def valid_in_all_rules( line_number: int, line: str, version: float -) -> list[VersionMatchingError]: - errors = [] - for rule in available_rules: +) -> Generator[VersionMatchingError]: + for rule in RULES: validator = rule(version, line_number, line) if not validator.validate(): - errors.append(validator.get_error()) - - return errors + yield validator.get_error() def validate(file_lines: list[str]) -> list[VersionMatchingError]: @@ -45,12 +49,9 @@ def validate(file_lines: list[str]) -> list[VersionMatchingError]: if found_version is None: return [] - errors = [] - - for number, line in enumerate(file_lines): - errors_in_line = valid_in_all_rules(number, line, found_version) - errors.extend(errors_in_line) + errors: list[VersionMatchingError] = [] - errors.extend(validate_multiple_version_decl(file_lines)) + for number, line in enumerate(validate_multiple_version_decl(file_lines, errors)): + errors.extend(valid_in_all_rules(number, line, found_version)) return errors diff --git a/m3u8/version_matching_rules.py b/m3u8/version_matching_rules.py index fae77dd0..07fbc1d9 100644 --- a/m3u8/version_matching_rules.py +++ b/m3u8/version_matching_rules.py @@ -1,16 +1,20 @@ -from dataclasses import dataclass +from __future__ import annotations + +import dataclasses +from abc import ABC, abstractmethod +from typing import ClassVar from m3u8 import protocol -@dataclass +@dataclasses.dataclass(slots=True) class VersionMatchingError(Exception): line_number: int line: str how_to_fix: str = "Please fix the version matching error." description: str = "There is a version matching error in the file." - def __str__(self): + def __str__(self) -> str: return ( "Version matching error found in the file when parsing in strict mode.\n" f"Line {self.line_number}: {self.description}\n" @@ -20,37 +24,45 @@ def __str__(self): ) -class VersionMatchRuleBase: - description: str = "" - how_to_fix: str = "" +RULES: list[type[VersionMatchRule]] = [] + + +@dataclasses.dataclass +class VersionMatchRule(ABC): version: float line_number: int line: str + DESCRIPTION: ClassVar[str] = "" + HOW_TO_FIX: ClassVar[str] = "" def __init__(self, version: float, line_number: int, line: str) -> None: self.version = version self.line_number = line_number self.line = line - def validate(self): + def __init_subclass__(cls) -> None: + RULES.append(cls) + + @abstractmethod + def validate(self) -> bool: raise NotImplementedError - def get_error(self): + def get_error(self) -> VersionMatchingError: return VersionMatchingError( line_number=self.line_number, line=self.line, - description=self.description, - how_to_fix=self.how_to_fix, + description=self.DESCRIPTION, + how_to_fix=self.HOW_TO_FIX, ) -class ValidIVInEXTXKEY(VersionMatchRuleBase): - description = ( +class ValidIVInEXTXKEY(VersionMatchRule): + DESCRIPTION: ClassVar[str] = ( "You must use at least protocol version 2 if you have IV in EXT-X-KEY." ) - how_to_fix = "Change the protocol version to 2 or higher." + HOW_TO_FIX: ClassVar[str] = "Change the protocol version to 2 or higher." - def validate(self): + def validate(self) -> bool: if protocol.ext_x_key not in self.line: return True @@ -60,38 +72,34 @@ def validate(self): return True -class ValidFloatingPointEXTINF(VersionMatchRuleBase): - description = "You must use at least protocol version 3 if you have floating point EXTINF duration values." - how_to_fix = "Change the protocol version to 3 or higher." +class ValidFloatingPointEXTINF(VersionMatchRule): + DESCRIPTION: ClassVar[str] = ( + "You must use at least protocol version 3 if you have floating point EXTINF duration values." + ) + HOW_TO_FIX: ClassVar[str] = "Change the protocol version to 3 or higher." - def validate(self): + def validate(self) -> bool: if protocol.extinf not in self.line: return True chunks = self.line.replace(protocol.extinf + ":", "").split(",", 1) duration = chunks[0] - def is_number(value: str): - try: - float(value) - return True - except ValueError: - return False + try: + float(duration) + except ValueError: + return False - def is_floating_number(value: str): - return is_number(value) and "." in value + return self.version >= 3 if "." in duration else True - if is_floating_number(duration): - return self.version >= 3 - return is_number(duration) - - -class ValidEXTXBYTERANGEOrEXTXIFRAMESONLY(VersionMatchRuleBase): - description = "You must use at least protocol version 4 if you have EXT-X-BYTERANGE or EXT-X-IFRAME-ONLY." - how_to_fix = "Change the protocol version to 4 or higher." +class ValidEXTXBYTERANGEOrEXTXIFRAMESONLY(VersionMatchRule): + DESCRIPTION: ClassVar[str] = ( + "You must use at least protocol version 4 if you have EXT-X-BYTERANGE or EXT-X-IFRAME-ONLY." + ) + HOW_TO_FIX: ClassVar[str] = "Change the protocol version to 4 or higher." - def validate(self): + def validate(self) -> bool: if ( protocol.ext_x_byterange not in self.line and protocol.ext_i_frames_only not in self.line @@ -99,10 +107,3 @@ def validate(self): return True return self.version >= 4 - - -available_rules: list[type[VersionMatchRuleBase]] = [ - ValidIVInEXTXKEY, - ValidFloatingPointEXTINF, - ValidEXTXBYTERANGEOrEXTXIFRAMESONLY, -]