From b6922c132219ed6a8c726d91f7a3d9f6ea52e715 Mon Sep 17 00:00:00 2001 From: Jesse Bannon Date: Sat, 11 Mar 2023 22:57:14 -0800 Subject: [PATCH] [FEATURE ] Reformat existing downloads --- .../subscriptions/subscription_download.py | 85 +++++++++++++++++++ src/ytdl_sub/utils/file_handler.py | 17 ++++ .../validators/file_path_validators.py | 14 +-- 3 files changed, 104 insertions(+), 12 deletions(-) diff --git a/src/ytdl_sub/subscriptions/subscription_download.py b/src/ytdl_sub/subscriptions/subscription_download.py index 05460ba1..c493ed72 100644 --- a/src/ytdl_sub/subscriptions/subscription_download.py +++ b/src/ytdl_sub/subscriptions/subscription_download.py @@ -1,9 +1,14 @@ import contextlib +import json import os import shutil from abc import ABC +from pathlib import Path +from typing import Iterable from typing import List from typing import Optional +from typing import Set +from typing import Tuple from ytdl_sub.entries.entry import Entry from ytdl_sub.plugins.plugin import Plugin @@ -14,7 +19,9 @@ from ytdl_sub.utils.exceptions import ValidationException from ytdl_sub.utils.file_handler import FileHandler from ytdl_sub.utils.file_handler import FileHandlerTransactionLog from ytdl_sub.utils.file_handler import FileMetadata +from ytdl_sub.utils.file_handler import get_file_extension from ytdl_sub.utils.thumbnail import convert_download_thumbnail +from ytdl_sub.ytdl_additions.enhanced_download_archive import EnhancedDownloadArchive def _get_split_plugin(plugins: List[Plugin]) -> Optional[Plugin]: @@ -303,3 +310,81 @@ class SubscriptionDownload(BaseSubscription, ABC): plugin.post_process_subscription() return self._enhanced_download_archive.get_file_handler_transaction_log() + + def _get_entries_for_reformat( + self, original_enhanced_download_archive: EnhancedDownloadArchive + ) -> Iterable[Entry]: + # ytdl-sub reformat asf.yaml --output xzy/ + entry_mapping: List[Tuple[Entry, Set[str]]] = [] + for download_mapping in original_enhanced_download_archive.mapping._entry_mappings.values(): + maybe_entry: Optional[Entry] = None + for file_name in download_mapping.file_names: + if file_name.endswith(".info.json"): + try: + with open( + Path(self.output_directory) / file_name, "r", encoding="utf-8" + ) as maybe_info_json: + maybe_entry = Entry( + entry_dict=json.load(maybe_info_json), + working_directory=self.working_directory, + ) + except Exception: + raise ValidationException( + "info.json file cannot be loaded - subscription cannot be reformatted" + ) + + if not maybe_entry: + raise ValidationException( + ".info.json file could not be found - subscription cannot be reformatted" + ) + + entry_mapping.append((maybe_entry, download_mapping.file_names)) + + for entry, file_names in entry_mapping: + for file_name in file_names: + # Remove any subdirectory part of the file name + _, file_name_no_dirs = os.path.split(Path(file_name)) + ext = get_file_extension(file_name_no_dirs) + + FileHandler.copy( + src_file_path=Path(self.output_directory) / file_name, + dst_file_path=Path(self.working_directory) / f"{entry.uid}.{ext}", + ) + + yield entry + + def reformat( + self, reformat_output_directory: str, dry_run: bool = False + ) -> FileHandlerTransactionLog: + original_enhanced_download_archive = self._enhanced_download_archive + self._enhanced_download_archive = EnhancedDownloadArchive( + subscription_name=self.name, + working_directory=self.working_directory, + output_directory=reformat_output_directory, + dry_run=dry_run, + ) + + self._enhanced_download_archive.reinitialize(dry_run=dry_run) + plugins = self._initialize_plugins() + + with self._subscription_download_context_managers(): + for entry in self._get_entries_for_reformat( + original_enhanced_download_archive=original_enhanced_download_archive + ): + entry_metadata = FileMetadata() + if isinstance(entry, tuple): + entry, entry_metadata = entry + + if split_plugin := _get_split_plugin(plugins): + self._process_split_entry( + split_plugin=split_plugin, plugins=plugins, dry_run=dry_run, entry=entry + ) + else: + self._process_entry( + plugins=plugins, dry_run=dry_run, entry=entry, entry_metadata=entry_metadata + ) + + for plugin in plugins: + plugin.post_process_subscription() + + return self._enhanced_download_archive.get_file_handler_transaction_log() diff --git a/src/ytdl_sub/utils/file_handler.py b/src/ytdl_sub/utils/file_handler.py index 654e87cb..64691095 100644 --- a/src/ytdl_sub/utils/file_handler.py +++ b/src/ytdl_sub/utils/file_handler.py @@ -11,6 +11,23 @@ from typing import Optional from typing import Set from typing import Union +from ytdl_sub.utils.subtitles import SUBTITLE_EXTENSIONS + + +def get_file_extension(file_name: Path | str) -> str: + if file_name.endswith(".info.json"): + return "info.json" + if any(file_name.endswith(f".{subtitle_ext}") for subtitle_ext in SUBTITLE_EXTENSIONS): + file_name_split = file_name.split(".") + ext = file_name_split[-1] + + # Try to capture .lang.ext + if len(file_name_split) > 2 and len(file_name_split[-2]) < 6: + ext = f"{file_name_split[-2]}.{file_name_split[-1]}" + + return ext + return file_name.rsplit(".", maxsplit=1)[-1] + def get_file_md5_hash(full_file_path: Path | str) -> str: """ diff --git a/src/ytdl_sub/validators/file_path_validators.py b/src/ytdl_sub/validators/file_path_validators.py index f0850c54..2c6b65cd 100644 --- a/src/ytdl_sub/validators/file_path_validators.py +++ b/src/ytdl_sub/validators/file_path_validators.py @@ -4,6 +4,7 @@ from typing import Any from typing import Dict from typing import Tuple +from ytdl_sub.utils.file_handler import get_file_extension from ytdl_sub.utils.subtitles import SUBTITLE_EXTENSIONS from ytdl_sub.utils.system import IS_WINDOWS from ytdl_sub.validators.string_formatter_validators import OverridesStringFormatterValidator @@ -46,18 +47,7 @@ class FilePathValidatorMixin: @classmethod def _get_extension_split(cls, file_name: str) -> Tuple[str, str]: - if file_name.endswith(".info.json"): - ext = "info.json" - elif any(file_name.endswith(f".{subtitle_ext}") for subtitle_ext in SUBTITLE_EXTENSIONS): - file_name_split = file_name.split(".") - ext = file_name_split[-1] - - # Try to capture .lang.ext - if len(file_name_split) > 2 and len(file_name_split[-2]) < 6: - ext = f"{file_name_split[-2]}.{file_name_split[-1]}" - else: - ext = file_name.rsplit(".", maxsplit=1)[-1] - + ext = get_file_extension(file_name) return file_name[: -len(ext)], ext @classmethod