Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 32 additions & 10 deletions yamale/command_line.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import os
import re
import multiprocessing
from typing import Callable, List, Optional
from .yamale_error import YamaleError
from .schema.validationresults import Result
from .version import __version__
Expand All @@ -22,7 +23,7 @@
schemas = {}


def _validate(schema_path, data_path, parser, strict, _raise_error):
def _validate(schema_path: str, data_path: str, parser: str, strict: bool, _raise_error: bool) -> list[Result]:
schema = schemas.get(schema_path)
try:
if not schema:
Expand All @@ -37,7 +38,7 @@ def _validate(schema_path, data_path, parser, strict, _raise_error):
return yamale.validate(schema, data, strict, _raise_error)


def _find_data_path_schema(data_path, schema_name):
def _find_data_path_schema(data_path: str, schema_name: str) -> Optional[str]:
"""Starts in the data file folder and recursively looks
in parents for `schema_name`"""
if not data_path or data_path == os.path.abspath(os.sep) or data_path == ".":
Expand All @@ -49,7 +50,7 @@ def _find_data_path_schema(data_path, schema_name):
return path[0]


def _find_schema(data_path, schema_name):
def _find_schema(data_path: str, schema_name: str) -> Optional[str]:
"""Checks if `schema_name` is a valid file, if not
searches in `data_path` for it."""

Expand All @@ -65,7 +66,13 @@ def _find_schema(data_path, schema_name):
return _find_data_path_schema(data_path, schema_name)


def _validate_file(yaml_path, schema_name, parser, strict, should_exclude):
def _validate_file(
yaml_path: str,
schema_name: str,
parser: str,
strict: bool,
should_exclude: Callable[[str], bool],
) -> None:
if should_exclude(yaml_path):
return
s = _find_schema(yaml_path, schema_name)
Expand All @@ -74,9 +81,16 @@ def _validate_file(yaml_path, schema_name, parser, strict, should_exclude):
_validate(s, yaml_path, parser, strict, True)


def _validate_dir(root, schema_name, cpus, parser, strict, should_exclude):
def _validate_dir(
root: str,
schema_name: str,
cpus: int,
parser: str,
strict: bool,
should_exclude: Callable[[str], bool],
) -> None:
pool = multiprocessing.Pool(processes=cpus)
res = []
res: List[multiprocessing.pool.AsyncResult] = []
error_messages = []
for root, _, files in os.walk(root):
for f in files:
Expand All @@ -100,10 +114,18 @@ def _validate_dir(root, schema_name, cpus, parser, strict, should_exclude):
raise ValueError("\n----\n".join(set(error_messages)))


def _router(paths, schema_name, cpus, parser, excludes=None, strict=True, verbose=False):
def _router(
paths: List[str],
schema_name: str,
cpus: int,
parser: str,
excludes: Optional[List[str]] = None,
strict: bool = True,
verbose: bool = False,
) -> None:
EXCLUDE_REGEXES = tuple(re.compile(e) for e in excludes) if excludes else tuple()

def should_exclude(yaml_path):
def should_exclude(yaml_path: str) -> bool:
has_match = any(pattern.search(yaml_path) for pattern in EXCLUDE_REGEXES)
if has_match and verbose:
print("Skipping validation for %s due to exclude pattern" % yaml_path)
Expand All @@ -122,8 +144,8 @@ def should_exclude(yaml_path):
_validate_file(abs_path, schema_name, parser, strict, should_exclude)


def main():
def int_or_auto(num_cpu):
def main() -> None:
def int_or_auto(num_cpu: str) -> int:
if num_cpu == "auto":
return multiprocessing.cpu_count()
return int(num_cpu)
Expand Down
9 changes: 5 additions & 4 deletions yamale/readers/yaml_reader.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
from __future__ import absolute_import
from io import StringIO
from typing import Any, Callable, Dict, List, Optional, TextIO


def _pyyaml(f):
def _pyyaml(f: TextIO) -> List[Any]:
import yaml

try:
Expand All @@ -12,17 +13,17 @@ def _pyyaml(f):
return list(yaml.load_all(f, Loader=Loader))


def _ruamel(f):
def _ruamel(f: TextIO) -> List[Any]:
from ruamel.yaml import YAML

yaml = YAML(typ="safe")
return list(yaml.load_all(f))


_parsers = {"pyyaml": _pyyaml, "ruamel": _ruamel}
_parsers: Dict[str, Callable[[TextIO], List[Any]]] = {"pyyaml": _pyyaml, "ruamel": _ruamel}


def parse_yaml(path=None, parser="pyyaml", content=None):
def parse_yaml(path: Optional[str] = None, parser: str = "pyyaml", content: Optional[str] = None) -> List[Any]:
try:
parse = _parsers[parser.lower()]
except KeyError:
Expand Down
8 changes: 4 additions & 4 deletions yamale/schema/datapath.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,14 @@
class DataPath(object):
def __init__(self, *path):
def __init__(self, *path: object) -> None:
self._path = path

def __add__(self, other):
def __add__(self, other: "DataPath") -> "DataPath":
dp = DataPath()
dp._path = self._path + other._path
return dp

def __str__(self):
def __str__(self) -> str:
return ".".join(map(str, (self._path)))

def __repr__(self):
def __repr__(self) -> str:
return "DataPath({})".format(repr(self._path))
43 changes: 28 additions & 15 deletions yamale/schema/schema.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,16 @@
from typing import Any, Dict, List, Optional, Type, Union

from .datapath import DataPath
from .validationresults import ValidationResult
from .. import syntax, util
from .. import validators as val

Container = Union[Dict[Any, Any], List[Any]]
SchemaNode = Union[val.Validator, Container]


class FatalValidationError(Exception):
def __init__(self, error):
def __init__(self, error: str) -> None:
super().__init__()
self.error = error

Expand All @@ -16,7 +21,13 @@ class Schema(object):
Still acts like a dict.
"""

def __init__(self, schema_dict, name="", validators=None, includes=None):
def __init__(
self,
schema_dict: Any,
name: str = "",
validators: Optional[Dict[str, Type[val.Validator]]] = None,
includes: Optional[Dict[str, "Schema"]] = None,
) -> None:
self.validators = validators or val.DefaultValidators
self.dict = schema_dict
self.name = name
Expand All @@ -25,12 +36,12 @@ def __init__(self, schema_dict, name="", validators=None, includes=None):
# schema
self.includes = {} if includes is None else includes

def add_include(self, type_dict):
def add_include(self, type_dict: Dict[str, Any]) -> None:
for include_name, custom_type in type_dict.items():
t = Schema(custom_type, name=include_name, validators=self.validators, includes=self.includes)
self.includes[include_name] = t

def _process_schema(self, path, schema_data, validators):
def _process_schema(self, path: DataPath, schema_data: Any, validators: Dict[str, Type[val.Validator]]) -> Any:
"""
Go through a schema and construct validators.
"""
Expand All @@ -41,23 +52,25 @@ def _process_schema(self, path, schema_data, validators):
schema_data = self._parse_schema_item(path, schema_data, validators)
return schema_data

def _parse_schema_item(self, path, expression, validators):
def _parse_schema_item(
self, path: DataPath, expression: str, validators: Dict[str, Type[val.Validator]]
) -> val.Validator:
try:
return syntax.parse(expression, validators)
except SyntaxError as e:
# Tack on some more context and rethrow.
error = str(e) + " at node '%s'" % str(path)
raise SyntaxError(error)

def validate(self, data, data_name, strict):
def validate(self, data: Any, data_name: Optional[str], strict: bool) -> ValidationResult:
path = DataPath()
try:
errors = self._validate(self._schema, data, path, strict)
except FatalValidationError as e:
errors = [e.error]
return ValidationResult(data_name, self.name, errors)

def _validate_item(self, validator, data, path, strict, key):
def _validate_item(self, validator: SchemaNode, data: Any, path: DataPath, strict: bool, key: Any) -> List[str]:
"""
Fetch item from data at the position key and validate with validator.

Expand All @@ -77,7 +90,7 @@ def _validate_item(self, validator, data, path, strict, key):

return self._validate(validator, data_item, path, strict)

def _validate(self, validator, data, path, strict):
def _validate(self, validator: SchemaNode, data: Any, path: DataPath, strict: bool) -> List[str]:
"""
Validate data with validator.
Special handling of non-primitive validators.
Expand Down Expand Up @@ -111,7 +124,7 @@ def _validate(self, validator, data, path, strict):

return errors

def _validate_static_map_list(self, validator, data, path, strict):
def _validate_static_map_list(self, validator: Container, data: Any, path: DataPath, strict: bool) -> List[str]:
if util.is_map(validator) and not util.is_map(data):
return ["%s : '%s' is not a map" % (path, data)]

Expand All @@ -131,7 +144,7 @@ def _validate_static_map_list(self, validator, data, path, strict):
errors += self._validate_item(sub_validator, data, path, strict, key)
return errors

def _validate_map_list(self, validator, data, path, strict):
def _validate_map_list(self, validator: Union[val.Map, val.List], data: Any, path: DataPath, strict: bool) -> List[str]:
errors = []

if not validator.validators:
Expand All @@ -151,14 +164,14 @@ def _validate_map_list(self, validator, data, path, strict):

return errors

def _validate_include(self, validator, data, path, strict):
def _validate_include(self, validator: val.Include, data: Any, path: DataPath, strict: bool) -> List[str]:
include_schema = self.includes.get(validator.include_name)
if not include_schema:
raise FatalValidationError("Include '%s' has not been defined." % validator.include_name)
strict = strict if validator.strict is None else validator.strict
return include_schema._validate(include_schema._schema, data, path, strict)

def _validate_any(self, validator, data, path, strict):
def _validate_any(self, validator: val.Any, data: Any, path: DataPath, strict: bool) -> List[str]:
if not validator.validators:
return []

Expand All @@ -177,8 +190,8 @@ def _validate_any(self, validator, data, path, strict):

return errors

def _validate_subset(self, validator, data, path, strict):
def _internal_validate(internal_data):
def _validate_subset(self, validator: val.Subset, data: Any, path: DataPath, strict: bool) -> List[str]:
def _internal_validate(internal_data: Any) -> List[str]:
sub_errors = []
for v in validator.validators:
err = self._validate(v, internal_data, path, strict)
Expand All @@ -203,7 +216,7 @@ def _internal_validate(internal_data):
errors += _internal_validate(data)
return errors

def _validate_primitive(self, validator, data, path):
def _validate_primitive(self, validator: val.Validator, data: Any, path: DataPath) -> List[str]:
errors = validator.validate(data)

for i, error in enumerate(errors):
Expand Down
13 changes: 8 additions & 5 deletions yamale/schema/validationresults.py
Original file line number Diff line number Diff line change
@@ -1,21 +1,24 @@
from typing import List, Optional


class Result(object):
def __init__(self, errors):
def __init__(self, errors: List[str]) -> None:
self.errors = errors

def __str__(self):
def __str__(self) -> str:
return "\n".join(self.errors)

def isValid(self):
def isValid(self) -> bool:
return len(self.errors) == 0


class ValidationResult(Result):
def __init__(self, data, schema, errors):
def __init__(self, data: Optional[str], schema: str, errors: List[str]) -> None:
super(ValidationResult, self).__init__(errors)
self.data = data
self.schema = schema

def __str__(self):
def __str__(self) -> str:
if self.isValid():
error_str = "'%s' is Valid" % self.data
else:
Expand Down
7 changes: 5 additions & 2 deletions yamale/syntax/parser.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
import ast
from typing import Dict, Optional, Type

from .. import validators as val

safe_globals = ("True", "False", "None")
safe_builtins = dict((f, __builtins__[f]) for f in safe_globals)


def _validate_expr(call_node, validators):
def _validate_expr(call_node: ast.Call, validators: Dict[str, Type[val.Validator]]) -> None:
# Validate that the expression uses a known, registered validator.
try:
func_name = call_node.func.id
Expand All @@ -28,7 +29,9 @@ def _validate_expr(call_node, validators):
raise SyntaxError("Argument values must either be constant literals, or else " "reference other validators.")


def parse(validator_string, validators=None):
def parse(
validator_string: str, validators: Optional[Dict[str, Type[val.Validator]]] = None
) -> val.Validator:
validators = validators or val.DefaultValidators
try:
tree = ast.parse(validator_string, mode="eval")
Expand Down
20 changes: 13 additions & 7 deletions yamale/util.py
Original file line number Diff line number Diff line change
@@ -1,37 +1,43 @@
from collections.abc import Mapping, Sequence
from typing import Any, Iterable, Iterator, KeysView, Optional, Set, Type, TypeVar, Union


def isstr(s):
T = TypeVar("T")


def isstr(s: Any) -> bool:
return isinstance(s, str)


def to_unicode(s):
def to_unicode(s: T) -> T:
return s


def is_list(obj):
def is_list(obj: Any) -> bool:
return isinstance(obj, Sequence) and not isstr(obj)


def is_map(obj):
def is_map(obj: Any) -> bool:
return isinstance(obj, Mapping)


def get_keys(obj):
def get_keys(obj: Any) -> Optional[Union[KeysView[Any], range]]:
if is_map(obj):
return obj.keys()
elif is_list(obj):
return range(len(obj))


def get_iter(iterable):
def get_iter(iterable: Union[Mapping[Any, Any], Iterable[Any]]) -> Iterable[tuple[Any, Any]]:
if isinstance(iterable, Mapping):
return iterable.items()
else:
return enumerate(iterable)


def get_subclasses(cls, _subclasses_yielded=None):
def get_subclasses(
cls: Type[Any], _subclasses_yielded: Optional[Set[Type[Any]]] = None
) -> Iterator[Type[Any]]:
"""
Generator recursively yielding all subclasses of the passed class (in
depth-first order).
Expand Down
Loading