More validator refactoring

This commit is contained in:
jbannon 2022-04-03 18:02:22 +00:00
parent 775342f0cf
commit 3eca67101b
9 changed files with 101 additions and 103 deletions

View file

@ -1,77 +0,0 @@
from typing import Any
from typing import Dict
from typing import Optional
from typing import Set
from typing import Type
from typing import TypeVar
from ytdl_subscribe.validators.base.validators import Validator
from ytdl_subscribe.validators.exceptions import ValidationException
T = TypeVar("T", bound=Validator)
class DictValidator(Validator):
expected_value_type = dict
expected_value_type_name = "object" # for non-python users
required_fields: Set[str] = set()
optional_fields: Set[str] = set()
_allow_extra_fields = False
def __init__(self, name, value):
super().__init__(name, value)
# Ensure all required keys are present
for required_key in self.required_fields:
if required_key not in self.value:
error_msg = (
f"'{self.name}' is missing the required field '{required_key}'"
)
raise ValidationException(error_msg)
# If no extra fields are allowed, ensure all fields are either
# required or optional fields
if not self._allow_extra_fields:
for object_key in self.keys:
if object_key not in self.allowed_fields:
error_msg = (
f"'{self.name}' contains the field '{object_key}' which is not allowed. "
f"Allowed fields: {', '.join(self.allowed_fields)}"
)
raise ValidationException(error_msg)
def validate_key(
self, key: str, validator: Type[T], default: Optional[Any] = None
) -> T:
value = self.get(key=key, default=default)
if value is None:
raise self._validation_exception(
f"{key} is missing when it should be present."
)
return validator(
name=f"{self.name}.{key}",
value=self.get(key=key, default=default),
)
@property
def dict(self) -> dict:
return self._value
@property
def allowed_fields(self):
return sorted(self.required_fields.union(self.optional_fields))
@property
def keys(self):
return sorted(list(self.dict.keys()))
def get(self, key: str, default: Optional[Any] = None) -> Any:
return self.dict.get(key, default)
class DictWithExtraFieldsValidator(DictValidator):
_allow_extra_fields = True

View file

@ -0,0 +1,41 @@
from typing import Set
from ytdl_subscribe.validators.base.validators import DictValidator
from ytdl_subscribe.validators.exceptions import ValidationException
class StrictDictValidator(DictValidator):
required_keys: Set[str] = set()
optional_keys: Set[str] = set()
allow_extra_fields = False
def __init__(self, name, value):
super().__init__(name, value)
if len(self.required_keys) == 0:
raise ValueError(
"No required fields when using a StrictDictValidator. Should be using DictValidator instead."
)
# Ensure all required keys are present
for required_key in self.required_keys:
if required_key not in self.value:
error_msg = (
f"'{self.name}' is missing the required field '{required_key}'"
)
raise ValidationException(error_msg)
# Ensure all keys are either required or optional keys if no extra field are allowed
if not self.allow_extra_fields:
for object_key in self.keys:
if object_key not in self.allowed_keys:
error_msg = (
f"'{self.name}' contains the field '{object_key}' which is not allowed. "
f"Allowed fields: {', '.join(self.allowed_keys)}"
)
raise ValidationException(error_msg)
@property
def allowed_keys(self):
return sorted(self.required_keys.union(self.optional_keys))

View file

@ -1,6 +1,7 @@
from typing import Any from typing import Any
from typing import Optional from typing import Optional
from typing import Type from typing import Type
from typing import TypeVar
from ytdl_subscribe.validators.exceptions import ValidationException from ytdl_subscribe.validators.exceptions import ValidationException
@ -57,3 +58,36 @@ class StringValidator(Validator):
@property @property
def value(self) -> str: def value(self) -> str:
return self._value return self._value
T = TypeVar("T", bound=Validator)
class DictValidator(Validator):
expected_value_type = dict
expected_value_type_name = "object" # for non-python users
@property
def dict(self) -> dict:
return self._value
@property
def keys(self):
return sorted(list(self.dict.keys()))
def get(self, key: str, default: Optional[Any] = None) -> Any:
return self.dict.get(key, default)
def validate_key(
self, key: str, validator: Type[T], default: Optional[Any] = None
) -> T:
value = self.get(key=key, default=default)
if value is None:
raise self._validation_exception(
f"{key} is missing when it should be present."
)
return validator(
name=f"{self.name}.{key}",
value=self.get(key=key, default=default),
)

View file

@ -2,18 +2,18 @@ from typing import Any
import yaml import yaml
from ytdl_subscribe.validators.base.dict_validator import DictValidator from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValidator
from ytdl_subscribe.validators.base.dict_validator import DictWithExtraFieldsValidator from ytdl_subscribe.validators.base.validators import DictValidator
from ytdl_subscribe.validators.base.validators import StringValidator from ytdl_subscribe.validators.base.validators import StringValidator
class ConfigValidator(DictValidator): class ConfigValidator(StrictDictValidator):
required_fields = {"working_directory", "presets"} required_keys = {"working_directory", "presets"}
def __init__(self, name: str, value: Any): def __init__(self, name: str, value: Any):
super().__init__(name, value) super().__init__(name, value)
self.working_directory = self.validate_key("working_directory", StringValidator) self.working_directory = self.validate_key("working_directory", StringValidator)
self.presets = self.validate_key("presets", DictWithExtraFieldsValidator) self.presets = self.validate_key("presets", DictValidator)
@classmethod @classmethod
def from_file_path(cls, config_path) -> "ConfigValidator": def from_file_path(cls, config_path) -> "ConfigValidator":

View file

@ -2,7 +2,7 @@ from typing import Any
from typing import Optional from typing import Optional
from ytdl_subscribe.utils.enums import SubscriptionSourceName from ytdl_subscribe.utils.enums import SubscriptionSourceName
from ytdl_subscribe.validators.base.dict_validator import DictValidator from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValidator
from ytdl_subscribe.validators.config.sources.soundcloud_validators import ( from ytdl_subscribe.validators.config.sources.soundcloud_validators import (
SoundcloudSourceValidator, SoundcloudSourceValidator,
) )
@ -13,9 +13,9 @@ from ytdl_subscribe.validators.config.sources.youtube_validators import (
from ytdl_subscribe.validators.exceptions import ValidationException from ytdl_subscribe.validators.exceptions import ValidationException
class PresetValidator(DictValidator): class PresetValidator(StrictDictValidator):
required_fields = {"post_process"} required_keys = {"post_process"}
optional_fields = { optional_keys = {
"ytdl_options", "ytdl_options",
"output_path", "output_path",
"overrides", "overrides",

View file

@ -7,7 +7,7 @@ from ytdl_subscribe.validators.config.sources.source_validator import SourceVali
class SoundcloudAlbumsAndSinglesDownloadValidator(DownloadStrategyValidator): class SoundcloudAlbumsAndSinglesDownloadValidator(DownloadStrategyValidator):
required_fields = {"username"} required_keys = {"username"}
def __init__(self, name, value): def __init__(self, name, value):
super().__init__(name, value) super().__init__(name, value)
@ -17,7 +17,7 @@ class SoundcloudAlbumsAndSinglesDownloadValidator(DownloadStrategyValidator):
class SoundcloudSourceValidator(SourceValidator): class SoundcloudSourceValidator(SourceValidator):
optional_fields = {"skip_premiere_tracks"} optional_keys = {"skip_premiere_tracks"}
download_strategy_validator_mapping = { download_strategy_validator_mapping = {
"albums_and_singles": SoundcloudAlbumsAndSinglesDownloadValidator "albums_and_singles": SoundcloudAlbumsAndSinglesDownloadValidator

View file

@ -3,18 +3,20 @@ from typing import Any
from typing import Dict from typing import Dict
from typing import Type from typing import Type
from ytdl_subscribe.validators.base.dict_validator import DictValidator from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValidator
from ytdl_subscribe.validators.base.dict_validator import DictWithExtraFieldsValidator
from ytdl_subscribe.validators.base.validators import StringValidator from ytdl_subscribe.validators.base.validators import StringValidator
class DownloadStrategyValidator(DictValidator): class DownloadStrategyValidator(StrictDictValidator):
pass pass
class SourceValidator(DictWithExtraFieldsValidator): class SourceValidator(StrictDictValidator):
# All media sources must define a download strategy # All media sources must define a download strategy
required_fields = {"download_strategy"} required_keys = {"download_strategy"}
# Extra fields will be strict-validated using other StictDictValidators
allow_extra_fields = True
download_strategy_validator_mapping: Dict[str, Type[DownloadStrategyValidator]] = {} download_strategy_validator_mapping: Dict[str, Type[DownloadStrategyValidator]] = {}
@ -32,7 +34,7 @@ class SourceValidator(DictWithExtraFieldsValidator):
# Remove all non-download strategy keys before passing the dict to the validator # Remove all non-download strategy keys before passing the dict to the validator
download_strategy_dict = copy.deepcopy(self.dict) download_strategy_dict = copy.deepcopy(self.dict)
for key_to_delete in self.allowed_fields: for key_to_delete in self.allowed_keys:
del download_strategy_dict[key_to_delete] del download_strategy_dict[key_to_delete]
download_strategy_class = self.download_strategy_validator_mapping[ download_strategy_class = self.download_strategy_validator_mapping[

View file

@ -8,7 +8,7 @@ from ytdl_subscribe.validators.config.sources.source_validator import SourceVali
class YoutubePlaylistDownloadValidator(DownloadStrategyValidator): class YoutubePlaylistDownloadValidator(DownloadStrategyValidator):
required_fields = {"playlist_id"} required_keys = {"playlist_id"}
def __init__(self, name, value): def __init__(self, name, value):
super().__init__(name, value) super().__init__(name, value)

View file

@ -8,29 +8,27 @@ from ytdl_subscribe.subscriptions.soundcloud import SoundcloudSubscription
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
from ytdl_subscribe.validators.base.dict_validator import DictValidator from ytdl_subscribe.validators.base.strict_dict_validator import StrictDictValidator
from ytdl_subscribe.validators.base.dict_validator import DictWithExtraFieldsValidator 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 PresetValidator from ytdl_subscribe.validators.config.preset_validator import PresetValidator
class SubscriptionValidator(DictValidator): class SubscriptionValidator(StrictDictValidator):
""" """
A Subscription is a preset but overrides it with specific values A Subscription is a preset but overrides it with specific values
""" """
required_fields = {"preset"} required_keys = {"preset"}
optional_fields = PresetValidator.required_fields.union( optional_keys = PresetValidator.required_keys.union(PresetValidator.optional_keys)
PresetValidator.optional_fields
)
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( self.overrides = self.validate_key(
key="overrides", key="overrides",
validator=DictWithExtraFieldsValidator, validator=DictValidator,
default={}, default={},
) )