refactor subscriptions to use validators

This commit is contained in:
jbannon 2022-04-04 22:30:48 +00:00
parent 54d35dce0e
commit 3a22a9bf10
7 changed files with 249 additions and 145 deletions

View file

@ -78,10 +78,19 @@ class Entry:
"""Returns the entry's file name when downloaded locally""" """Returns the entry's file name when downloaded locally"""
return f"{self.uid}.{self.ext}" return f"{self.uid}.{self.ext}"
@property
def download_thumbnail_name(self) -> str:
"""Returns the thumbnail's file name when downloaded locally TODO: unit test this"""
return f"{self.uid}.{self.thumbnail_ext}"
def file_path(self, relative_directory: str): def file_path(self, relative_directory: str):
"""Returns the entry's file path with respect to the relative directory""" """Returns the entry's file path with respect to the relative directory"""
return str(Path(relative_directory) / self.download_file_name) return str(Path(relative_directory) / self.download_file_name)
def thumbnail_path(self, relative_directory: str):
"""Returns the entry's thumbnail path with respect to the relative directory"""
return str(Path(relative_directory) / self.download_thumbnail_name)
def to_dict(self) -> Dict: def to_dict(self) -> Dict:
"""Returns the entry's values as a dictionary""" """Returns the entry's values as a dictionary"""
return { return {

View file

@ -4,40 +4,57 @@ from ytdl_subscribe.downloaders.soundcloud_downloader import SoundcloudDownloade
from ytdl_subscribe.entries.soundcloud import SoundcloudAlbum from ytdl_subscribe.entries.soundcloud import SoundcloudAlbum
from ytdl_subscribe.entries.soundcloud import SoundcloudTrack from ytdl_subscribe.entries.soundcloud import SoundcloudTrack
from ytdl_subscribe.subscriptions.subscription import Subscription from ytdl_subscribe.subscriptions.subscription import Subscription
from ytdl_subscribe.validators.config.sources.soundcloud_validators import (
SoundcloudAlbumsAndSinglesDownloadValidator,
)
from ytdl_subscribe.validators.config.sources.soundcloud_validators import (
SoundcloudSourceValidator,
)
class SoundcloudSubscription(Subscription): class SoundcloudAlbumsAndSinglesSubscription(Subscription):
def extract_info(self): source_validator_type = SoundcloudSourceValidator
download_strategy_type = SoundcloudAlbumsAndSinglesDownloadValidator
@property
def source_options(self) -> SoundcloudSourceValidator:
return super().source_options
@property
def download_strategy_options(self) -> SoundcloudAlbumsAndSinglesDownloadValidator:
return super().download_strategy_options
def extract_info(
self,
):
tracks: List[SoundcloudTrack] = [] tracks: List[SoundcloudTrack] = []
soundcloud_downloader = SoundcloudDownloader( soundcloud_downloader = SoundcloudDownloader(
output_directory=self.output_path, output_directory=self.output_options.output_directory.value,
working_directory=self.WORKING_DIRECTORY, working_directory=self.config_options.working_directory.value,
ytdl_options=self.ytdl_opts, ytdl_options=self.ytdl_options.dict,
) )
if self.options.get("download_strategy") == "albums_and_singles": # Get the album info first. This tells us which track ids belong
# Get the album info first. This tells us which track ids belong # to an album. Unfortunately we cannot use download_archive or info.json for this
# to an album. Unfortunately we cannot use download_archive or info.json for this albums: List[SoundcloudAlbum] = soundcloud_downloader.download_albums(
albums: List[SoundcloudAlbum] = soundcloud_downloader.download_albums( artist_name=self.download_strategy_options.username.value
artist_name=self.options["username"] )
for album in albums:
tracks += album.album_tracks(
skip_premiere_tracks=self.source_options.skip_premiere_tracks.value
) )
for album in albums: # only add tracks that are not part of an album
tracks += album.album_tracks( single_tracks = soundcloud_downloader.download_tracks(
skip_premiere_tracks=self.options["skip_premiere_tracks"] artist_name=self.download_strategy_options.username.value
) )
tracks += [
track
for track in single_tracks
if not any(album.contains(track) for album in albums)
]
# only add tracks that are not part of an album for e in tracks:
single_tracks = soundcloud_downloader.download_tracks( self.post_process_entry(e)
artist_name=self.options["username"]
)
tracks += [
track
for track in single_tracks
if not any(album.contains(track) for album in albums)
]
for e in tracks:
self.post_process_entry(e)
else:
raise ValueError("Invalid download_strategy field for Soundcloud")

View file

@ -1,99 +1,115 @@
import os import os
from copy import deepcopy
from shutil import copyfile from shutil import copyfile
from typing import Type
from typing import TypeVar
import dicttoxml import dicttoxml
import music_tag import music_tag
from PIL import Image from PIL import Image
from sanitize_filename import sanitize
from ytdl_subscribe.entries.entry import Entry from ytdl_subscribe.entries.entry import Entry
from ytdl_subscribe.validators.config.config_validator import ConfigValidator
from ytdl_subscribe.validators.config.metadata_options.metadata_options_validator import (
MetadataOptionsValidator,
)
from ytdl_subscribe.validators.config.output_options.output_options_validator import (
OutputOptionsValidator,
)
from ytdl_subscribe.validators.config.preset_validator import OverridesValidator
from ytdl_subscribe.validators.config.preset_validator import YTDLOptionsValidator
from ytdl_subscribe.validators.config.sources.source_validator import (
DownloadStrategyValidator,
)
from ytdl_subscribe.validators.config.sources.source_validator import SourceValidator
SOURCE_T = TypeVar("SOURCE_T", bound=SourceValidator)
DOWNLOAD_STRATEGY_T = TypeVar("DOWNLOAD_STRATEGY_T", bound=DownloadStrategyValidator)
class Subscription(object): class Subscription(object):
WORKING_DIRECTORY = "" source_validator_type: Type[SOURCE_T]
download_strategy_type: Type[DOWNLOAD_STRATEGY_T]
def __init__(self, name, options, ytdl_opts, post_process, overrides, output_path): def __init__(
self,
name: str,
config_options: ConfigValidator,
source_options: SourceValidator,
output_options: OutputOptionsValidator,
metadata_options: MetadataOptionsValidator,
ytdl_options: YTDLOptionsValidator,
overrides: OverridesValidator,
):
""" """
Parameters Parameters
---------- ----------
name: str name: str
Name of the subscription Name of the subscription
options: dict config_options: ConfigValidator
Dictionary of ytdl options, specific to the source type source_options: SourceValidator
ytdl_opts: dict output_options: OutputOptionsValidator
Dictionary of options passed directly to ytdl. metadata_options: MetadataOptionsValidator
See `youtube_dl.YoutubeDL.YoutubeDL` for options. ytdl_options: YTDLOptionsValidator
post_process: dict overrides: OverridesValidator
Dictionary of ytdl-subscribe post processing options
overrides: dict
Dictionary that overrides every ytdl entry. Be careful what you override!
output_path: str
Base path to save files to
""" """
self.name = name self.name = name
self.options = options self.config_options = config_options
self.ytdl_opts = ytdl_opts or dict()
self.post_process = post_process self.__source_options = source_options
self.__download_strategy_options = source_options.download_strategy
if not isinstance(self.__source_options, self.source_validator_type):
raise ValueError("Source options does not match the expected type")
if not isinstance(
self.__download_strategy_options, self.download_strategy_type
):
raise ValueError("Download strategy does not match the expected type")
self.output_options = output_options
self.metadata_options = metadata_options
self.ytdl_options = ytdl_options
self.overrides = overrides self.overrides = overrides
for k, v in deepcopy(overrides).items():
self.overrides[f"sanitized_{k}"] = sanitize(v)
self.output_path = output_path
# Separate each subscription's working directory @property
self.WORKING_DIRECTORY += f"{'/' if self.WORKING_DIRECTORY else ''}{self.name}" def source_options(self) -> SOURCE_T:
return self.__source_options
# Always set outtmpl to the id and extension. Will be renamed using the subscription's @property
# output_path value def download_strategy_options(self) -> DOWNLOAD_STRATEGY_T:
self.ytdl_opts["outtmpl"] = self.WORKING_DIRECTORY + "/%(id)s.%(ext)s" return self.__download_strategy_options
self.ytdl_opts["writethumbnail"] = True
def format_filepath(self, filepath_formatter: str, entry: Entry, makedirs=False):
"""
Convert a filepath value in the config to an actual filepath.
Parameters
----------
filepath_formatter: str
File path relative to the specified output path
entry: Entry
Entry used to populate any format variables in the filepath
makedirs: bool
Whether to create all directories in the final filepath.
Returns
-------
str
Full filepath.
"""
file_name = entry.apply_formatter(filepath_formatter, overrides=self.overrides)
output_file_path = f"{self.output_path}/{file_name}"
if makedirs:
os.makedirs(os.path.dirname(output_file_path), exist_ok=True)
return output_file_path
def _post_process_tagging(self, entry: Entry): def _post_process_tagging(self, entry: Entry):
id3_options = self.metadata_options.id3
t = music_tag.load_file( t = music_tag.load_file(
entry.file_path(relative_directory=self.WORKING_DIRECTORY) entry.file_path(
relative_directory=self.config_options.working_directory.value
)
) )
for tag, tag_formatter in self.post_process["tagging"].items(): for tag, tag_formatter in id3_options.tags.dict.items():
t[tag] = entry.apply_formatter(tag_formatter, overrides=self.overrides) t[tag] = entry.apply_formatter(
format_string=tag_formatter, overrides=self.overrides.dict
)
t.save() t.save()
def _post_process_nfo(self, entry): def _post_process_nfo(self, entry):
nfo = {} nfo = {}
for tag, tag_formatter in self.post_process["nfo"].items(): nfo_options = self.metadata_options.nfo
nfo[tag] = entry.apply_formatter(tag_formatter, overrides=self.overrides)
for tag, tag_formatter in nfo_options.tags.dict.items():
nfo[tag] = entry.apply_formatter(
format_string=tag_formatter, overrides=self.overrides
)
xml = dicttoxml.dicttoxml( xml = dicttoxml.dicttoxml(
obj=nfo, obj=nfo,
root="nfo_root" in self.post_process, root=True, # We assume all NFOs have a root. Maybe we should not?
custom_root=self.post_process.get("nfo_root"), custom_root=nfo_options.nfo_root.value,
attr_type=False, attr_type=False,
) )
nfo_file_path = self.format_filepath( nfo_file_path = entry.apply_formatter(
self.post_process["nfo_name"], entry, makedirs=True format_string=nfo_options.nfo_name.format_string,
overrides=self.overrides.dict,
) )
with open(nfo_file_path, "wb") as f: with open(nfo_file_path, "wb") as f:
f.write(xml) f.write(xml)
@ -105,37 +121,46 @@ class Subscription(object):
raise NotImplemented("Each source needs to implement how it extracts info") raise NotImplemented("Each source needs to implement how it extracts info")
def post_process_entry(self, entry: Entry): def post_process_entry(self, entry: Entry):
if "tagging" in self.post_process: if self.metadata_options.id3:
self._post_process_tagging(entry) self._post_process_tagging(entry)
# Move the file after all direct file modifications are complete # Move the file after all direct file modifications are complete
output_file_path = self.format_filepath( entry_source_file_path = entry.file_path(
self.post_process["file_name"], entry, makedirs=True relative_directory=self.config_options.working_directory.value
) )
copyfile(
entry.file_path(relative_directory=self.WORKING_DIRECTORY), output_file_path entry_destination_file_path = entry.apply_formatter(
format_string=self.output_options.file_name.format_string,
overrides=self.overrides.dict,
) )
os.makedirs(os.path.dirname(entry_destination_file_path), exist_ok=True)
copyfile(entry_source_file_path, entry_destination_file_path)
# Download the thumbnail if its present # Download the thumbnail if its present
if "thumbnail_name" in self.post_process: if self.output_options.thumbnail_name:
thumbnail_dest_path = self.format_filepath( thumbnail_source_path = entry.thumbnail_path(
self.post_process["thumbnail_name"], relative_directory=self.config_options.working_directory.value
entry,
makedirs=True,
) )
if not os.path.isfile(thumbnail_dest_path):
thumbnail_file_path = ( thumbnail_destination_path = entry.apply_formatter(
f"{self.WORKING_DIRECTORY}/{entry.uid}.{entry.thumbnail_ext}" format_string=self.output_options.thumbnail_name.format_string,
overrides=self.overrides.dict,
)
os.makedirs(os.path.dirname(entry_destination_file_path), exist_ok=True)
# If the thumbnail is to be converted, then save the converted thumbnail to the
# output filepath
if self.output_options.convert_thumbnail:
im = Image.open(thumbnail_source_path).convert("RGB")
im.save(
fp=thumbnail_destination_path,
format=self.output_options.convert_thumbnail.value,
) )
if os.path.isfile(thumbnail_file_path): # Otherwise, just copy the downloaded thumbnail
copyfile(thumbnail_file_path, thumbnail_dest_path) else:
copyfile(thumbnail_source_path, thumbnail_destination_path)
if "convert_thumbnail" in self.post_process: if self.metadata_options.nfo:
# TODO: Clean with yaml definitions
if self.post_process["convert_thumbnail"] == "jpg":
self.post_process["convert_thumbnail"] = "jpeg"
im = Image.open(thumbnail_dest_path).convert("RGB")
im.save(thumbnail_dest_path, self.post_process["convert_thumbnail"])
if "nfo" in self.post_process:
self._post_process_nfo(entry) self._post_process_nfo(entry)

View file

@ -1,17 +1,34 @@
from ytdl_subscribe.downloaders.youtube_downloader import YoutubeDownloader from ytdl_subscribe.downloaders.youtube_downloader import YoutubeDownloader
from ytdl_subscribe.subscriptions.subscription import Subscription from ytdl_subscribe.subscriptions.subscription import Subscription
from ytdl_subscribe.validators.config.sources.youtube_validators import (
YoutubePlaylistDownloadValidator,
)
from ytdl_subscribe.validators.config.sources.youtube_validators import (
YoutubeSourceValidator,
)
class YoutubeSubscription(Subscription): class YoutubeSubscription(Subscription):
source_validator_type = YoutubeSourceValidator
download_strategy_type = YoutubePlaylistDownloadValidator
@property
def source_options(self) -> YoutubeSourceValidator:
return super().source_options
@property
def download_strategy_options(self) -> YoutubePlaylistDownloadValidator:
return super().download_strategy_options
def extract_info(self): def extract_info(self):
youtube_downloader = YoutubeDownloader( youtube_downloader = YoutubeDownloader(
output_directory=self.output_path, output_directory=self.output_options.output_directory.value,
working_directory=self.WORKING_DIRECTORY, working_directory=self.config_options.working_directory.value,
ytdl_options=self.ytdl_opts, ytdl_options=self.ytdl_options.dict,
) )
entries = youtube_downloader.download_playlist( entries = youtube_downloader.download_playlist(
playlist_id=self.options["playlist_id"] playlist_id=self.download_strategy_options.playlist_id.value
) )
for e in entries: for e in entries:

View file

@ -1,6 +1,11 @@
import copy
from typing import Any from typing import Any
from typing import Dict
from typing import Optional from typing import Optional
from typing import Tuple from typing import Tuple
from typing import Type
import sanitize_filename
from ytdl_subscribe.utils.enums import SubscriptionSourceName from ytdl_subscribe.utils.enums import SubscriptionSourceName
from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValidator from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValidator
@ -24,6 +29,23 @@ from ytdl_subscribe.validators.config.sources.youtube_validators import (
from ytdl_subscribe.validators.exceptions import ValidationException from ytdl_subscribe.validators.exceptions import ValidationException
class YTDLOptionsValidator(DictValidator):
pass
class OverridesValidator(DictFormatterValidator):
@property
def dict(self) -> dict:
"""For overrides, create sanitized versions of each entry for convenience"""
output_dict = copy.deepcopy(super().dict)
for key in self.keys:
output_dict[f"sanitized_{key}"] = sanitize_filename.sanitize(
output_dict[key]
)
return output_dict
class PresetValidator(StrictDictValidator): class PresetValidator(StrictDictValidator):
required_keys = {"output_options, metadata_options"} required_keys = {"output_options, metadata_options"}
optional_keys = { optional_keys = {
@ -32,6 +54,11 @@ class PresetValidator(StrictDictValidator):
*SubscriptionSourceName.all(), *SubscriptionSourceName.all(),
} }
subscription_source_validator_mapping: Dict[str, Type[SourceValidator]] = {
SubscriptionSourceName.SOUNDCLOUD: SoundcloudSourceValidator,
SubscriptionSourceName.YOUTUBE: YoutubeSourceValidator,
}
def __validate_and_get_subscription_source(self) -> Tuple[str, SourceValidator]: def __validate_and_get_subscription_source(self) -> Tuple[str, SourceValidator]:
subscription_source: Optional[SourceValidator] = None subscription_source: Optional[SourceValidator] = None
subscription_source_name: Optional[str] = None subscription_source_name: Optional[str] = None
@ -42,17 +69,10 @@ class PresetValidator(StrictDictValidator):
f"'{self.name}' can only have one of the following sources: {SubscriptionSourceName.pretty_all()}" f"'{self.name}' can only have one of the following sources: {SubscriptionSourceName.pretty_all()}"
) )
if key == SubscriptionSourceName.SOUNDCLOUD: if key in self.subscription_source_validator_mapping:
subscription_source_name = SubscriptionSourceName.SOUNDCLOUD subscription_source_name = key
subscription_source = self.validate_key( subscription_source = self.validate_key(
key=key, key=key, validator=self.subscription_source_validator_mapping[key]
validator=SoundcloudSourceValidator,
)
elif key == SubscriptionSourceName.YOUTUBE:
subscription_source_name = SubscriptionSourceName.YOUTUBE
subscription_source = self.validate_key(
key=key,
validator=YoutubeSourceValidator,
) )
# If subscription source was not set, error # If subscription source was not set, error
@ -78,10 +98,10 @@ class PresetValidator(StrictDictValidator):
key="metadata_options", validator=MetadataOptionsValidator key="metadata_options", validator=MetadataOptionsValidator
) )
self.ytdl_options = self.validate_key_if_present( self.ytdl_options = self.validate_key(
key="ytdl_options", validator=DictValidator key="ytdl_options", validator=YTDLOptionsValidator, default={}
) )
self.overrides = self.validate_key_if_present( self.overrides = self.validate_key(
key="overrides", validator=DictFormatterValidator key="overrides", validator=OverridesValidator, default={}
) )

View file

@ -11,9 +11,7 @@ class SoundcloudAlbumsAndSinglesDownloadValidator(DownloadStrategyValidator):
def __init__(self, name, value): def __init__(self, name, value):
super().__init__(name, value) super().__init__(name, value)
self.username = self.validate_key( self.username = self.validate_key(key="username", validator=StringValidator)
key="username", validator=StringValidator
).value
class SoundcloudSourceValidator(SourceValidator): class SoundcloudSourceValidator(SourceValidator):
@ -27,4 +25,4 @@ class SoundcloudSourceValidator(SourceValidator):
super().__init__(name=name, value=value) super().__init__(name=name, value=value)
self.skip_premiere_tracks = self.validate_key( self.skip_premiere_tracks = self.validate_key(
"skip_premiere_tracks", BoolValidator, default=True "skip_premiere_tracks", BoolValidator, default=True
).value )

View file

@ -4,7 +4,9 @@ from typing import List
import yaml import yaml
from mergedeep import mergedeep from mergedeep import mergedeep
from ytdl_subscribe.subscriptions.soundcloud import SoundcloudSubscription from ytdl_subscribe.subscriptions.soundcloud import (
SoundcloudAlbumsAndSinglesSubscription,
)
from ytdl_subscribe.subscriptions.subscription import Subscription from ytdl_subscribe.subscriptions.subscription import Subscription
from ytdl_subscribe.subscriptions.youtube import YoutubeSubscription from ytdl_subscribe.subscriptions.youtube import YoutubeSubscription
from ytdl_subscribe.utils.enums import SubscriptionSourceName from ytdl_subscribe.utils.enums import SubscriptionSourceName
@ -12,7 +14,14 @@ from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValid
from ytdl_subscribe.validators.base.validators import DictValidator from ytdl_subscribe.validators.base.validators import DictValidator
from ytdl_subscribe.validators.base.validators import StringValidator from ytdl_subscribe.validators.base.validators import StringValidator
from ytdl_subscribe.validators.config.config_validator import ConfigValidator from ytdl_subscribe.validators.config.config_validator import ConfigValidator
from ytdl_subscribe.validators.config.preset_validator import OverridesValidator
from ytdl_subscribe.validators.config.preset_validator import PresetValidator from ytdl_subscribe.validators.config.preset_validator import PresetValidator
from ytdl_subscribe.validators.config.sources.soundcloud_validators import (
SoundcloudSourceValidator,
)
from ytdl_subscribe.validators.config.sources.youtube_validators import (
YoutubePlaylistDownloadValidator,
)
class SubscriptionValidator(StrictDictValidator): class SubscriptionValidator(StrictDictValidator):
@ -26,9 +35,11 @@ class SubscriptionValidator(StrictDictValidator):
def __init__(self, config: ConfigValidator, name: str, value: Any): def __init__(self, config: ConfigValidator, name: str, value: Any):
super().__init__(name, value) super().__init__(name, value)
self.config = config self.config = config
self.overrides = self.validate_key(
# Ensure the overrides defined here are valid
_ = self.validate_key(
key="overrides", key="overrides",
validator=DictValidator, validator=OverridesValidator,
default={}, default={},
) )
@ -58,20 +69,27 @@ class SubscriptionValidator(StrictDictValidator):
) )
def to_subscription(self) -> Subscription: def to_subscription(self) -> Subscription:
if self.preset.subscription_source_name == SubscriptionSourceName.SOUNDCLOUD: if isinstance(
subscription_class = SoundcloudSubscription self.preset.subscription_source.download_strategy,
elif self.preset.subscription_source_name == SubscriptionSourceName.YOUTUBE: SoundcloudAlbumsAndSinglesSubscription,
):
subscription_class = SoundcloudAlbumsAndSinglesSubscription
elif isinstance(
self.preset.subscription_source.download_strategy,
YoutubePlaylistDownloadValidator,
):
subscription_class = YoutubeSubscription subscription_class = YoutubeSubscription
else: else:
raise ValueError("subscription source class not found") raise ValueError("subscription source class not found")
return subscription_class( return subscription_class(
name=self.name, name=self.name,
options=self.preset.dict.get(self.preset.subscription_source_name), config_options=self.config,
ytdl_opts=self.preset.dict.get("ytdl_options"), source_options=self.preset.subscription_source,
post_process=self.preset.dict.get("post_process"), output_options=self.preset.output_options,
overrides=self.preset.dict.get("overrides"), metadata_options=self.preset.metadata_options,
output_path=self.preset.dict.get("output_path"), ytdl_options=self.preset.ytdl_options,
overrides=self.preset.overrides,
) )
@classmethod @classmethod