From 5ff97844643b59d313d56bd613408c10ecc29c09 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Tue, 15 Feb 2022 23:05:54 +0100 Subject: [PATCH 01/12] Add metaresolver --- src/class_resolver/metaresolver.py | 44 +++++++++++ tests/test_metaresolver.py | 119 +++++++++++++++++++++++++++++ 2 files changed, 163 insertions(+) create mode 100644 src/class_resolver/metaresolver.py create mode 100644 tests/test_metaresolver.py diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py new file mode 100644 index 0000000..8e017c2 --- /dev/null +++ b/src/class_resolver/metaresolver.py @@ -0,0 +1,44 @@ +from typing import Any, Iterable, Mapping, Type, TypeVar + +from .api import ClassResolver +from .utils import Hint, OptionalKwargs + +X = TypeVar("X") + +__all__ = [ + "Metaresolver", +] + + +def is_hint(hint: Any, cls: Type[X]) -> bool: + return hint == Hint[cls] + + +class Metaresolver: + """A resolver of resolvers.""" + + def __init__(self, resolvers: Iterable[ClassResolver]): + self.resolvers: Mapping[Type, ClassResolver] = { + resolver.base: resolver + for resolver in resolvers + } + self.names = { + resolver.normalize_cls(cls): resolver + for cls, resolver in self.resolvers.items() + } + + def check_kwargs(self, base: Type[X], query: Hint[X], kwargs: OptionalKwargs) -> bool: + main = self.resolvers[base] + signature = main.signature(query) + parameters = signature.parameters + + for name, parameter in parameters.items(): + annotation = parameter.annotation + + next_resolver = self.names.get(name) + if next_resolver is None: + raise NotImplementedError + elif not is_hint(annotation, next_resolver.base): + raise TypeError + else: + raise NotImplementedError diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py new file mode 100644 index 0000000..c1fd8e4 --- /dev/null +++ b/tests/test_metaresolver.py @@ -0,0 +1,119 @@ +import unittest +from typing import Optional, Type + +from class_resolver import ClassResolver, Hint, OptionalKwargs +from class_resolver.metaresolver import Metaresolver, is_hint + + +class Baz: + def __init__(self, value: bool = False): + self.value = value + + +class XBaz(Baz): + pass + + +class YBaz(Baz): + pass + + +baz_resolver = ClassResolver.from_subclasses(Baz) + + +class Bar: + def __init__( + self, + baz: Hint[Baz] = None, + baz_kwargs: OptionalKwargs = None, + ): + self.baz = baz_resolver.make(baz, baz_kwargs) + + +class AlphaBar(Bar): + pass + + +class BetaBar(Bar): + pass + + +bar_resolver = ClassResolver.from_subclasses(Bar) + + +class Foo: + def __init__( + self, + *, + bar: Hint[Bar] = None, + bar_kwargs: OptionalKwargs = None, + param_1: float, + param_2: Optional[int] = None, + ): + self.bar = bar_resolver.make(bar, bar_kwargs) + self.param_1 = param_1 + self.param_2 = param_2 or 5 + + +class AFoo(Foo): + pass + + +class BFoo(Foo): + pass + + +foo_resolver = ClassResolver.from_subclasses(Foo) + + +class TestMetaResolver(unittest.TestCase): + """""" + + def setUp(self) -> None: + self.meta_resolver = Metaresolver([baz_resolver, bar_resolver, foo_resolver]) + + def test_is_hint(self): + """Test hint predicate.""" + self.assertTrue(is_hint(Hint[Foo], Foo)) + self.assertFalse(is_hint(Hint[Foo], Bar)) + self.assertFalse(is_hint(Type[Bar], Bar)) + self.assertFalse(is_hint(str, Bar)) + self.assertFalse(is_hint(None, Bar)) + + def test_check(self): + self.assertTrue(self.meta_resolver.check_kwargs( + Foo, AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + "baz_kwargs": { + "value": True, + } + }, + "param_1": 3.0, + # Param 2 is optional, so not necessary to give + }, + )) + self.assertTrue(self.meta_resolver.check_kwargs( + Foo, AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + # baz_kwargs value not necessary since has default + }, + "param_1": 3.0, + # Param 2 is optional, so not necessary to give + }, + )) + self.assertFalse(self.meta_resolver.check_kwargs( + Foo, AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + }, + # Missing param_1 !! + }, + )) From 3dab3fc5464ba999a82d6287bf1ccdacbda574e0 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Tue, 15 Feb 2022 23:45:26 +0100 Subject: [PATCH 02/12] Cleanup --- src/class_resolver/metaresolver.py | 51 ++++++++++++++------ tests/test_metaresolver.py | 75 +++++++++++++++++------------- 2 files changed, 79 insertions(+), 47 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index 8e017c2..d3cc837 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -1,4 +1,5 @@ -from typing import Any, Iterable, Mapping, Type, TypeVar +import inspect +from typing import Any, Iterable, Mapping, Optional, Tuple, Type, TypeVar from .api import ClassResolver from .utils import Hint, OptionalKwargs @@ -19,26 +20,48 @@ class Metaresolver: def __init__(self, resolvers: Iterable[ClassResolver]): self.resolvers: Mapping[Type, ClassResolver] = { - resolver.base: resolver - for resolver in resolvers - } - self.names = { - resolver.normalize_cls(cls): resolver - for cls, resolver in self.resolvers.items() + resolver.base: resolver for resolver in resolvers } + self.names = {resolver.suffix: resolver for cls, resolver in self.resolvers.items()} def check_kwargs(self, base: Type[X], query: Hint[X], kwargs: OptionalKwargs) -> bool: main = self.resolvers[base] signature = main.signature(query) parameters = signature.parameters - - for name, parameter in parameters.items(): + for key, parameter, related_key, related_parameter in _iter_params(parameters): annotation = parameter.annotation - - next_resolver = self.names.get(name) + next_resolver = self.names.get(key) if next_resolver is None: - raise NotImplementedError + if key not in kwargs: + if parameter.default is parameter.empty: + raise ValueError(f"{key} without default not given") + else: + if not isinstance(kwargs[key], parameter.annotation): + raise TypeError elif not is_hint(annotation, next_resolver.base): - raise TypeError + raise TypeError( + f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" + ) + elif not self.check_kwargs( + next_resolver.base, + next_resolver.lookup(kwargs[key]), + kwargs[related_key], + ): + return False else: - raise NotImplementedError + continue + + +def _iter_params( + parameters, +) -> Iterable[Tuple[str, inspect.Parameter, str, Optional[inspect.Parameter]]]: + kwarg_map = {} + for key in parameters.items(): + related_key = f"{key}_kwargs" + if related_key in parameters: + kwarg_map[key] = related_key + for key in parameters: + related_key = f"{key}_kwargs" + if key in kwarg_map: + continue + yield key, parameters[key], related_key, parameters.get(related_key) diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index c1fd8e4..79191a9 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -81,39 +81,48 @@ def test_is_hint(self): self.assertFalse(is_hint(None, Bar)) def test_check(self): - self.assertTrue(self.meta_resolver.check_kwargs( - Foo, AFoo, - { - "bar": "alpha", - "bar_kwargs": { - "baz": "x", - "baz_kwargs": { - "value": True, - } + self.assertTrue( + self.meta_resolver.check_kwargs( + Foo, + AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + "baz_kwargs": { + "value": True, + }, + }, + "param_1": 3.0, + # Param 2 is optional, so not necessary to give }, - "param_1": 3.0, - # Param 2 is optional, so not necessary to give - }, - )) - self.assertTrue(self.meta_resolver.check_kwargs( - Foo, AFoo, - { - "bar": "alpha", - "bar_kwargs": { - "baz": "x", - # baz_kwargs value not necessary since has default + ) + ) + self.assertTrue( + self.meta_resolver.check_kwargs( + Foo, + AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + # baz_kwargs value not necessary since has default + }, + "param_1": 3.0, + # Param 2 is optional, so not necessary to give }, - "param_1": 3.0, - # Param 2 is optional, so not necessary to give - }, - )) - self.assertFalse(self.meta_resolver.check_kwargs( - Foo, AFoo, - { - "bar": "alpha", - "bar_kwargs": { - "baz": "x", + ) + ) + self.assertFalse( + self.meta_resolver.check_kwargs( + Foo, + AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + }, + # Missing param_1 !! }, - # Missing param_1 !! - }, - )) + ) + ) From 4b4a76ae3ea1afb61f6deddea7d2cfd82b26d4bf Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 00:22:25 +0100 Subject: [PATCH 03/12] Fix tests --- src/class_resolver/metaresolver.py | 50 ++++++++++++++++-------------- tests/test_metaresolver.py | 38 +++++++++-------------- 2 files changed, 40 insertions(+), 48 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index d3cc837..bbde1f2 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -1,5 +1,5 @@ import inspect -from typing import Any, Iterable, Mapping, Optional, Tuple, Type, TypeVar +from typing import Any, Callable, Iterable, Mapping, Optional, Tuple, Type, TypeVar from .api import ClassResolver from .utils import Hint, OptionalKwargs @@ -24,44 +24,46 @@ def __init__(self, resolvers: Iterable[ClassResolver]): } self.names = {resolver.suffix: resolver for cls, resolver in self.resolvers.items()} - def check_kwargs(self, base: Type[X], query: Hint[X], kwargs: OptionalKwargs) -> bool: - main = self.resolvers[base] - signature = main.signature(query) - parameters = signature.parameters + def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: + signature = inspect.signature(func) + parameters = dict(signature.parameters) + if kwargs is None: + kwargs = {} for key, parameter, related_key, related_parameter in _iter_params(parameters): annotation = parameter.annotation next_resolver = self.names.get(key) - if next_resolver is None: + if next_resolver is not None: + if not is_hint(annotation, next_resolver.base): + raise TypeError( + f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" + ) + self.check_kwargs( + next_resolver.lookup(kwargs[key]), + kwargs.get(related_key, {}), + ) + else: if key not in kwargs: if parameter.default is parameter.empty: raise ValueError(f"{key} without default not given") else: - if not isinstance(kwargs[key], parameter.annotation): - raise TypeError - elif not is_hint(annotation, next_resolver.base): - raise TypeError( - f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" - ) - elif not self.check_kwargs( - next_resolver.base, - next_resolver.lookup(kwargs[key]), - kwargs[related_key], - ): - return False - else: - continue + try: + instance_flag = isinstance(kwargs[key], parameter.annotation) + except TypeError: + raise TypeError(f"{key} {kwargs[key]} {parameter.annotation}") from None + if not instance_flag: + raise ValueError + return True def _iter_params( parameters, ) -> Iterable[Tuple[str, inspect.Parameter, str, Optional[inspect.Parameter]]]: kwarg_map = {} - for key in parameters.items(): + for key in parameters: related_key = f"{key}_kwargs" if related_key in parameters: - kwarg_map[key] = related_key + kwarg_map[related_key] = key for key in parameters: - related_key = f"{key}_kwargs" if key in kwarg_map: continue - yield key, parameters[key], related_key, parameters.get(related_key) + yield key, parameters[key], f"{key}_kwargs", parameters.get(f"{key}_kwargs") diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index 79191a9..95831f4 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -81,9 +81,8 @@ def test_is_hint(self): self.assertFalse(is_hint(None, Bar)) def test_check(self): - self.assertTrue( - self.meta_resolver.check_kwargs( - Foo, + true_kwargs = [ + ( AFoo, { "bar": "alpha", @@ -96,26 +95,14 @@ def test_check(self): "param_1": 3.0, # Param 2 is optional, so not necessary to give }, - ) - ) - self.assertTrue( - self.meta_resolver.check_kwargs( - Foo, - AFoo, - { - "bar": "alpha", - "bar_kwargs": { - "baz": "x", - # baz_kwargs value not necessary since has default - }, - "param_1": 3.0, - # Param 2 is optional, so not necessary to give - }, - ) - ) - self.assertFalse( - self.meta_resolver.check_kwargs( - Foo, + ), + ] + for func, kwargs in true_kwargs: + with self.subTest(): + self.assertTrue(self.meta_resolver.check_kwargs(func, kwargs)) + + false_kwargs = [ + ( AFoo, { "bar": "alpha", @@ -125,4 +112,7 @@ def test_check(self): # Missing param_1 !! }, ) - ) + ] + for func, kwargs in false_kwargs: + with self.subTest(), self.assertRaises(ValueError): + self.meta_resolver.check_kwargs(func, kwargs) From 85ec0de29bd29fc139a86e131d39582ed84e08ec Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 00:34:31 +0100 Subject: [PATCH 04/12] Add additional tests to round out coverage --- src/class_resolver/metaresolver.py | 7 +++++-- tests/test_metaresolver.py | 16 ++++++++++++++-- 2 files changed, 19 insertions(+), 4 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index bbde1f2..cce3bc5 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -37,9 +37,12 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: raise TypeError( f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" ) + query = kwargs.get(key) + if query is None: + raise KeyError self.check_kwargs( - next_resolver.lookup(kwargs[key]), - kwargs.get(related_key, {}), + next_resolver.lookup(query), + kwargs.get(related_key), ) else: if key not in kwargs: diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index 95831f4..66a01be 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -111,8 +111,20 @@ def test_check(self): }, # Missing param_1 !! }, - ) + ), + ( + AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + }, + "param_1": "3.0", # wrong type, should be float + }, + ), + (AFoo, None), + (AFoo, {}), ] for func, kwargs in false_kwargs: - with self.subTest(), self.assertRaises(ValueError): + with self.subTest(), self.assertRaises((ValueError, KeyError)): self.meta_resolver.check_kwargs(func, kwargs) From 3f0d0aa0deadb4bd145950ca9fef4191d0bb00c3 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 00:49:41 +0100 Subject: [PATCH 05/12] Additional cleanup --- src/class_resolver/metaresolver.py | 38 +++++++++++++++++++++++------- tests/test_metaresolver.py | 36 +++++++++++++++++++--------- 2 files changed, 55 insertions(+), 19 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index cce3bc5..2a99189 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -1,35 +1,56 @@ +# -*- coding: utf-8 -*- + +"""An argument checker.""" + import inspect from typing import Any, Callable, Iterable, Mapping, Optional, Tuple, Type, TypeVar from .api import ClassResolver from .utils import Hint, OptionalKwargs -X = TypeVar("X") - __all__ = [ + "is_hint", "Metaresolver", ] +X = TypeVar("X") + def is_hint(hint: Any, cls: Type[X]) -> bool: - return hint == Hint[cls] + """Check if the hint is applicable to the given class. + + :param hint: The hint type + :param cls: The class to check + :returns: If the hint is appropriate for the class + """ + if not isinstance(cls, type): + raise TypeError + return hint == Hint[cls] # type: ignore class Metaresolver: """A resolver of resolvers.""" def __init__(self, resolvers: Iterable[ClassResolver]): + """Instantiate a meta-resolver. + + :param resolvers: A set of resolvers to index for checking kwargs + """ self.resolvers: Mapping[Type, ClassResolver] = { resolver.base: resolver for resolver in resolvers } self.names = {resolver.suffix: resolver for cls, resolver in self.resolvers.items()} def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: - signature = inspect.signature(func) - parameters = dict(signature.parameters) + """Check the appropriate of the kwargs with a given function. + + :param func: A function or class to check + :param kwargs: The keyword arguments to pass to the function + :returns: True if there are no issues, raises if there are. + """ if kwargs is None: kwargs = {} - for key, parameter, related_key, related_parameter in _iter_params(parameters): + for key, parameter, related_key in _iter_params(func): annotation = parameter.annotation next_resolver = self.names.get(key) if next_resolver is not None: @@ -59,8 +80,9 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: def _iter_params( - parameters, + func, ) -> Iterable[Tuple[str, inspect.Parameter, str, Optional[inspect.Parameter]]]: + parameters = inspect.signature(func).parameters kwarg_map = {} for key in parameters: related_key = f"{key}_kwargs" @@ -69,4 +91,4 @@ def _iter_params( for key in parameters: if key in kwarg_map: continue - yield key, parameters[key], f"{key}_kwargs", parameters.get(f"{key}_kwargs") + yield key, parameters[key], f"{key}_kwargs" diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index 66a01be..b94df57 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -1,3 +1,7 @@ +# -*- coding: utf-8 -*- + +"""Tests for the argument checker.""" + import unittest from typing import Optional, Type @@ -6,23 +10,27 @@ class Baz: - def __init__(self, value: bool = False): + """A dummy class.""" + + def __init__(self, value: bool = False): # noqa:D107 self.value = value class XBaz(Baz): - pass + """A dummy child of the Baz class.""" class YBaz(Baz): - pass + """A dummy child of the Baz class.""" baz_resolver = ClassResolver.from_subclasses(Baz) class Bar: - def __init__( + """A dummy class.""" + + def __init__( # noqa:D107 self, baz: Hint[Baz] = None, baz_kwargs: OptionalKwargs = None, @@ -31,18 +39,20 @@ def __init__( class AlphaBar(Bar): - pass + """A dummy child of the Bar class.""" class BetaBar(Bar): - pass + """A dummy child of the Bar class.""" bar_resolver = ClassResolver.from_subclasses(Bar) class Foo: - def __init__( + """A dummy class.""" + + def __init__( # noqa:D107 self, *, bar: Hint[Bar] = None, @@ -56,20 +66,21 @@ def __init__( class AFoo(Foo): - pass + """A dummy child of the Foo class.""" class BFoo(Foo): - pass + """A dummy child of the Foo class.""" foo_resolver = ClassResolver.from_subclasses(Foo) class TestMetaResolver(unittest.TestCase): - """""" + """A test case for the argument checker.""" def setUp(self) -> None: + """Set up the test case.""" self.meta_resolver = Metaresolver([baz_resolver, bar_resolver, foo_resolver]) def test_is_hint(self): @@ -79,8 +90,11 @@ def test_is_hint(self): self.assertFalse(is_hint(Type[Bar], Bar)) self.assertFalse(is_hint(str, Bar)) self.assertFalse(is_hint(None, Bar)) + with self.assertRaises(TypeError): + is_hint(..., 5) - def test_check(self): + def test_argument_checker(self): + """Test the argument checker.""" true_kwargs = [ ( AFoo, From cd5d0890cce9ba95bb26561b04f724edafdedb55 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 01:01:27 +0100 Subject: [PATCH 06/12] More cleanup --- src/class_resolver/metaresolver.py | 35 +++++++++++++++++++++++++----- tests/test_metaresolver.py | 14 +++++++----- 2 files changed, 38 insertions(+), 11 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index 2a99189..62b5b1d 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -9,6 +9,7 @@ from .utils import Hint, OptionalKwargs __all__ = [ + "check_kwargs", "is_hint", "Metaresolver", ] @@ -22,12 +23,33 @@ def is_hint(hint: Any, cls: Type[X]) -> bool: :param hint: The hint type :param cls: The class to check :returns: If the hint is appropriate for the class + :raises TypeError: If the ``cls`` is not a type """ if not isinstance(cls, type): raise TypeError return hint == Hint[cls] # type: ignore +def check_kwargs( + func: Callable, + kwargs: OptionalKwargs = None, + *, + resolvers: Iterable[ClassResolver], +) -> bool: + """Check the appropriate of the kwargs with a given function. + + :param func: A function or class to check + :param kwargs: The keyword arguments to pass to the function + :param resolvers: A set of resolvers to index for checking kwargs + :returns: True if there are no issues, raises if there are. + """ + return Metaresolver(resolvers).check_kwargs(func, kwargs) + + +class ArgumentError(TypeError): + """A custom argument error.""" + + class Metaresolver: """A resolver of resolvers.""" @@ -47,6 +69,7 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: :param func: A function or class to check :param kwargs: The keyword arguments to pass to the function :returns: True if there are no issues, raises if there are. + :raises ArgumentError: If there is an error in the kwargs """ if kwargs is None: kwargs = {} @@ -55,12 +78,12 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: next_resolver = self.names.get(key) if next_resolver is not None: if not is_hint(annotation, next_resolver.base): - raise TypeError( + raise ArgumentError( f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" ) query = kwargs.get(key) if query is None: - raise KeyError + raise ArgumentError(f"{key} is missing from the arguments") self.check_kwargs( next_resolver.lookup(query), kwargs.get(related_key), @@ -68,20 +91,20 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: else: if key not in kwargs: if parameter.default is parameter.empty: - raise ValueError(f"{key} without default not given") + raise ArgumentError(f"{key} without default not given") else: try: instance_flag = isinstance(kwargs[key], parameter.annotation) except TypeError: - raise TypeError(f"{key} {kwargs[key]} {parameter.annotation}") from None + raise ArgumentError(f"{key} {kwargs[key]} {parameter.annotation}") from None if not instance_flag: - raise ValueError + raise ArgumentError return True def _iter_params( func, -) -> Iterable[Tuple[str, inspect.Parameter, str, Optional[inspect.Parameter]]]: +) -> Iterable[Tuple[str, inspect.Parameter, str]]: parameters = inspect.signature(func).parameters kwarg_map = {} for key in parameters: diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index b94df57..9f45655 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -6,7 +6,7 @@ from typing import Optional, Type from class_resolver import ClassResolver, Hint, OptionalKwargs -from class_resolver.metaresolver import Metaresolver, is_hint +from class_resolver.metaresolver import ArgumentError, check_kwargs, is_hint class Baz: @@ -81,7 +81,11 @@ class TestMetaResolver(unittest.TestCase): def setUp(self) -> None: """Set up the test case.""" - self.meta_resolver = Metaresolver([baz_resolver, bar_resolver, foo_resolver]) + self.resolvers = [baz_resolver, bar_resolver, foo_resolver] + + def check_kwargs(self, func, kwargs) -> bool: + """Check the kwargs.""" + return check_kwargs(func, kwargs, resolvers=self.resolvers) def test_is_hint(self): """Test hint predicate.""" @@ -113,7 +117,7 @@ def test_argument_checker(self): ] for func, kwargs in true_kwargs: with self.subTest(): - self.assertTrue(self.meta_resolver.check_kwargs(func, kwargs)) + self.assertTrue(self.check_kwargs(func, kwargs)) false_kwargs = [ ( @@ -140,5 +144,5 @@ def test_argument_checker(self): (AFoo, {}), ] for func, kwargs in false_kwargs: - with self.subTest(), self.assertRaises((ValueError, KeyError)): - self.meta_resolver.check_kwargs(func, kwargs) + with self.subTest(), self.assertRaises(ArgumentError): + self.check_kwargs(func, kwargs) From a00b3f665845de0b303098a29ba1d52bd34af94a Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 01:04:08 +0100 Subject: [PATCH 07/12] Finish docs --- src/class_resolver/metaresolver.py | 2 +- tests/test_metaresolver.py | 9 ++++++--- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index 62b5b1d..eca9bd8 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -3,7 +3,7 @@ """An argument checker.""" import inspect -from typing import Any, Callable, Iterable, Mapping, Optional, Tuple, Type, TypeVar +from typing import Any, Callable, Iterable, Mapping, Tuple, Type, TypeVar from .api import ClassResolver from .utils import Hint, OptionalKwargs diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index 9f45655..dc94162 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -12,7 +12,8 @@ class Baz: """A dummy class.""" - def __init__(self, value: bool = False): # noqa:D107 + def __init__(self, value: bool = False): + """Instantiate the dummy class.""" self.value = value @@ -30,11 +31,12 @@ class YBaz(Baz): class Bar: """A dummy class.""" - def __init__( # noqa:D107 + def __init__( self, baz: Hint[Baz] = None, baz_kwargs: OptionalKwargs = None, ): + """Instantiate the dummy class.""" self.baz = baz_resolver.make(baz, baz_kwargs) @@ -52,7 +54,7 @@ class BetaBar(Bar): class Foo: """A dummy class.""" - def __init__( # noqa:D107 + def __init__( self, *, bar: Hint[Bar] = None, @@ -60,6 +62,7 @@ def __init__( # noqa:D107 param_1: float, param_2: Optional[int] = None, ): + """Instantiate the dummy class.""" self.bar = bar_resolver.make(bar, bar_kwargs) self.param_1 = param_1 self.param_2 = param_2 or 5 From 320c788629f4901002de65e1dfa8ee08b8585ae3 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 01:52:32 +0100 Subject: [PATCH 08/12] Add additional logic --- src/class_resolver/metaresolver.py | 38 +++++++++++++++++++++--------- tests/test_metaresolver.py | 38 ++++++++++++++++++++++++++++++ 2 files changed, 65 insertions(+), 11 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index eca9bd8..b832e60 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -3,7 +3,9 @@ """An argument checker.""" import inspect -from typing import Any, Callable, Iterable, Mapping, Tuple, Type, TypeVar +from typing import Any, Callable, Iterable, Mapping, Tuple, Type, TypeVar, Union + +from typing_extensions import get_args, get_origin from .api import ClassResolver from .utils import Hint, OptionalKwargs @@ -30,6 +32,9 @@ def is_hint(hint: Any, cls: Type[X]) -> bool: return hint == Hint[cls] # type: ignore +SIMPLE_TYPES = {float, int, bool, str, type(None)} + + def check_kwargs( func: Callable, kwargs: OptionalKwargs = None, @@ -76,29 +81,40 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: for key, parameter, related_key in _iter_params(func): annotation = parameter.annotation next_resolver = self.names.get(key) + value = kwargs.get(key) if next_resolver is not None: if not is_hint(annotation, next_resolver.base): raise ArgumentError( f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" ) - query = kwargs.get(key) - if query is None: + if value is None: raise ArgumentError(f"{key} is missing from the arguments") self.check_kwargs( - next_resolver.lookup(query), + next_resolver.lookup(value), kwargs.get(related_key), ) else: - if key not in kwargs: + if value is None: if parameter.default is parameter.empty: raise ArgumentError(f"{key} without default not given") else: - try: - instance_flag = isinstance(kwargs[key], parameter.annotation) - except TypeError: - raise ArgumentError(f"{key} {kwargs[key]} {parameter.annotation}") from None - if not instance_flag: - raise ArgumentError + origin = get_origin(annotation) + if origin is Union: + args = get_args(annotation) + if all(arg in SIMPLE_TYPES for arg in args): + if isinstance(value, args): + pass + else: + raise ArgumentError + else: + raise NotImplementedError(f"{args}") + elif origin is None: + if isinstance(value, annotation): + pass + else: + raise ArgumentError + else: + raise NotImplementedError(f"origin: {origin}") return True diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index dc94162..bdb0e02 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -100,6 +100,30 @@ def test_is_hint(self): with self.assertRaises(TypeError): is_hint(..., 5) + def test_annotation_resolver_mismatch(self): + """Test when the key for a resolver is annotated incorrectly.""" + + class Bax: + """A dummy class.""" + + def __init__(self, baz: int): + """Instantiate the dummy class.""" + + with self.assertRaises(ArgumentError): + self.check_kwargs(Bax, {"baz": "x"}) + + def test_annotation_mismatch(self): + """Test when a random value is annotated incorrectly.""" + + class Bax: + """A dummy class.""" + + def __init__(self, value: Optional[int]): + """Instantiate the dummy class.""" + self.value = value + + self.assertTrue(self.check_kwargs(Bax, {"value": 5})) + def test_argument_checker(self): """Test the argument checker.""" true_kwargs = [ @@ -117,6 +141,20 @@ def test_argument_checker(self): # Param 2 is optional, so not necessary to give }, ), + ( + AFoo, + { + "bar": "alpha", + "bar_kwargs": { + "baz": "x", + "baz_kwargs": { + "value": True, + }, + }, + "param_1": 3.0, + "param_2": 2, + }, + ), ] for func, kwargs in true_kwargs: with self.subTest(): From 187686e98b1d4c38c735af378f8552f2d4d40608 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 01:54:46 +0100 Subject: [PATCH 09/12] Update setup.cfg --- setup.cfg | 3 +++ 1 file changed, 3 insertions(+) diff --git a/setup.cfg b/setup.cfg index 383d892..7c4a46c 100644 --- a/setup.cfg +++ b/setup.cfg @@ -46,6 +46,9 @@ keywords = configurability [options] +install_requires = + typing_extensions + # Random options zip_safe = false include_package_data = True From 38b791f5fc4cfb6f0b333bb0e3c75921e02932bb Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 16 Feb 2022 02:08:17 +0100 Subject: [PATCH 10/12] Add remaining tests --- src/class_resolver/metaresolver.py | 4 ++-- tests/test_metaresolver.py | 37 ++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 2 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index b832e60..714aeba 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -107,14 +107,14 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: else: raise ArgumentError else: - raise NotImplementedError(f"{args}") + raise ArgumentError elif origin is None: if isinstance(value, annotation): pass else: raise ArgumentError else: - raise NotImplementedError(f"origin: {origin}") + raise ArgumentError(f"origin: {origin}") return True diff --git a/tests/test_metaresolver.py b/tests/test_metaresolver.py index bdb0e02..9712a08 100644 --- a/tests/test_metaresolver.py +++ b/tests/test_metaresolver.py @@ -2,6 +2,7 @@ """Tests for the argument checker.""" +import typing as t import unittest from typing import Optional, Type @@ -124,6 +125,42 @@ def __init__(self, value: Optional[int]): self.assertTrue(self.check_kwargs(Bax, {"value": 5})) + def test_annotation_extended_origin_mismatch(self): + """Test when a random value is annotated incorrectly.""" + + class Bax: + """A dummy class.""" + + def __init__(self, value: t.Union[int, float, None]): + """Instantiate the dummy class.""" + + with self.assertRaises(ArgumentError): + self.check_kwargs(Bax, {"value": "nope"}) + + def test_annotation_extended_origin_mismatch_2(self): + """Test when a random value is annotated incorrectly.""" + + class Bax: + """A dummy class.""" + + def __init__(self, value: t.Union[int, float, None, t.Iterable[int]]): + """Instantiate the dummy class.""" + + with self.assertRaises(ArgumentError): + self.check_kwargs(Bax, {"value": 1}) + + def test_annotation_invalid_origin(self): + """Test when a random value is annotated incorrectly.""" + + class Bax: + """A dummy class.""" + + def __init__(self, value: t.Iterable[int]): + """Instantiate the dummy class.""" + + with self.assertRaises(ArgumentError): + self.check_kwargs(Bax, {"value": 5}) + def test_argument_checker(self): """Test the argument checker.""" true_kwargs = [ From 4ba06bbf5878164e2d7a3de915798515a4cd2830 Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Thu, 17 Feb 2022 11:34:49 +0100 Subject: [PATCH 11/12] Update metaresolver.py --- src/class_resolver/metaresolver.py | 141 +++++++++++++++++++++++++---- 1 file changed, 121 insertions(+), 20 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index 714aeba..d870a6a 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -3,12 +3,24 @@ """An argument checker.""" import inspect -from typing import Any, Callable, Iterable, Mapping, Tuple, Type, TypeVar, Union +from typing import ( + Any, + Callable, + Iterable, + Mapping, + Optional, + Tuple, + Type, + TypeVar, + Union, +) from typing_extensions import get_args, get_origin -from .api import ClassResolver -from .utils import Hint, OptionalKwargs +from class_resolver.api import ClassResolver +from class_resolver.base import BaseResolver +from class_resolver.func import FunctionResolver +from class_resolver.utils import Hint, HintOrType, HintType, OptionalKwargs __all__ = [ "check_kwargs", @@ -29,12 +41,22 @@ def is_hint(hint: Any, cls: Type[X]) -> bool: """ if not isinstance(cls, type): raise TypeError - return hint == Hint[cls] # type: ignore + if hint == Hint[cls]: # type: ignore + return True + if hint == HintType[cls]: # type: ignore + return True + if hint == HintOrType[cls]: # type: ignore + return True + return False SIMPLE_TYPES = {float, int, bool, str, type(None)} +def _is_simple_union(u) -> bool: + return all(arg in SIMPLE_TYPES for arg in get_args(u)) + + def check_kwargs( func: Callable, kwargs: OptionalKwargs = None, @@ -54,19 +76,30 @@ def check_kwargs( class ArgumentError(TypeError): """A custom argument error.""" + def __init__(self, func, key, text): + self.func = func + self.key = key + self.text = text + + def __str__(self): + return f"{self.func} {self.key}: {self.text}" + class Metaresolver: """A resolver of resolvers.""" - def __init__(self, resolvers: Iterable[ClassResolver]): + def __init__( + self, + resolvers: Iterable[ClassResolver], + extras: Optional[Mapping[str, BaseResolver]] = None, + ): """Instantiate a meta-resolver. :param resolvers: A set of resolvers to index for checking kwargs """ - self.resolvers: Mapping[Type, ClassResolver] = { - resolver.base: resolver for resolver in resolvers - } - self.names = {resolver.suffix: resolver for cls, resolver in self.resolvers.items()} + self.names = {resolver.suffix: resolver for resolver in resolvers} + if extras: + self.names.update(extras) def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: """Check the appropriate of the kwargs with a given function. @@ -79,16 +112,22 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: if kwargs is None: kwargs = {} for key, parameter, related_key in _iter_params(func): + if key == "kwargs": + continue annotation = parameter.annotation next_resolver = self.names.get(key) value = kwargs.get(key) if next_resolver is not None: + if isinstance(next_resolver, FunctionResolver): + continue if not is_hint(annotation, next_resolver.base): raise ArgumentError( - f"{key} has bad annotation {annotation} wrt resolver {next_resolver}" + func, + key, + f"has bad annotation {annotation} wrt resolver {next_resolver}", ) if value is None: - raise ArgumentError(f"{key} is missing from the arguments") + raise ArgumentError(func, key, f"is missing from the arguments") self.check_kwargs( next_resolver.lookup(value), kwargs.get(related_key), @@ -96,25 +135,45 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: else: if value is None: if parameter.default is parameter.empty: - raise ArgumentError(f"{key} without default not given") + raise ArgumentError( + func, + key, + f"no default, value not given. Signature: {inspect.signature(func)}", + ) else: origin = get_origin(annotation) if origin is Union: - args = get_args(annotation) - if all(arg in SIMPLE_TYPES for arg in args): + if not _is_simple_union(annotation): + raise ArgumentError( + func, key, f"has inappropriate annotation: {annotation}" + ) + else: + args = get_args(annotation) if isinstance(value, args): pass else: - raise ArgumentError - else: - raise ArgumentError + raise ArgumentError( + func, + key, + f"value {value} does not match annotation {annotation}", + ) elif origin is None: - if isinstance(value, annotation): + try: + instance_check = isinstance(value, annotation) + except TypeError: + raise ArgumentError( + func, key, f"invalid annotation {annotation} ({type(annotation)})" + ) from None + if instance_check: pass + elif annotation == float and isinstance(value, int): + pass # log a warning? else: - raise ArgumentError + raise ArgumentError( + func, key, f"{value} mismatched annotation {annotation}" + ) else: - raise ArgumentError(f"origin: {origin}") + raise ArgumentError(func, key, f"unhandled origin {origin}") return True @@ -131,3 +190,45 @@ def _iter_params( if key in kwarg_map: continue yield key, parameters[key], f"{key}_kwargs" + + +def _main(): + import json + + from pykeen.datasets import dataset_resolver + from pykeen.experiments.cli import HERE + from pykeen.losses import loss_resolver + from pykeen.models import model_resolver + from pykeen.nn.emb import constrainer_resolver, normalizer_resolver + from pykeen.nn.init import initializer_resolver + from pykeen.pipeline import pipeline + from pykeen.regularizers import regularizer_resolver + + print(HERE) + r = Metaresolver( + [ + model_resolver, + regularizer_resolver, + loss_resolver, + dataset_resolver, + constrainer_resolver, + normalizer_resolver, + initializer_resolver, + ], + extras={ + "entity_initializer": initializer_resolver, + "entity_normalizer": normalizer_resolver, + "entity_constrainer": constrainer_resolver, + "relation_initializer": initializer_resolver, + "relation_normalizer": normalizer_resolver, + "relation_constrainer": constrainer_resolver, + }, + ) + for path in HERE.glob("*/*.json"): + data = json.loads(path.read_text()) + kwargs = data["pipeline"] + r.check_kwargs(pipeline, kwargs) + + +if __name__ == "__main__": + _main() From 282ca37298d8e4cd58333cf0c63cdb8f3a50d86b Mon Sep 17 00:00:00 2001 From: Charles Tapley Hoyt Date: Wed, 22 Feb 2023 00:39:22 +0100 Subject: [PATCH 12/12] Update metaresolver.py --- src/class_resolver/metaresolver.py | 216 ++++++++++++++++++----------- 1 file changed, 134 insertions(+), 82 deletions(-) diff --git a/src/class_resolver/metaresolver.py b/src/class_resolver/metaresolver.py index d870a6a..b90f7c8 100644 --- a/src/class_resolver/metaresolver.py +++ b/src/class_resolver/metaresolver.py @@ -6,6 +6,7 @@ from typing import ( Any, Callable, + Dict, Iterable, Mapping, Optional, @@ -54,6 +55,7 @@ def is_hint(hint: Any, cls: Type[X]) -> bool: def _is_simple_union(u) -> bool: + """Check if the arguments to type U are all simple.""" return all(arg in SIMPLE_TYPES for arg in get_args(u)) @@ -76,18 +78,20 @@ def check_kwargs( class ArgumentError(TypeError): """A custom argument error.""" - def __init__(self, func, key, text): + def __init__(self, func, parameter, text): self.func = func - self.key = key + self.parameter = parameter self.text = text def __str__(self): - return f"{self.func} {self.key}: {self.text}" + return f"{self.func} {self.parameter.name}: {self.text}" class Metaresolver: """A resolver of resolvers.""" + parameter_name_to_resolver: Dict[str, BaseResolver] + def __init__( self, resolvers: Iterable[ClassResolver], @@ -96,10 +100,107 @@ def __init__( """Instantiate a meta-resolver. :param resolvers: A set of resolvers to index for checking kwargs + :param extras: A dictionary of additional resolvers, e.g. class resolvers + that don't have suffixes or function resolvers given explicitly with + parameter names """ - self.names = {resolver.suffix: resolver for resolver in resolvers} + self.parameter_name_to_resolver = {resolver.suffix: resolver for resolver in resolvers} if extras: - self.names.update(extras) + self.parameter_name_to_resolver.update(extras) + + def check_kwarg(self, func: Callable, parameter: inspect.Parameter, related_key: str, kwargs): + annotation = parameter.annotation + value_not_given = parameter.name not in kwargs + value = kwargs.get(parameter.name) + + resolver = self.parameter_name_to_resolver.get(parameter.name) + if resolver is not None: + if isinstance(resolver, FunctionResolver): + # TODO more careful checks of functions' arguments + return True + + # If the annotation for this parameter does not match the base + # class type inside the resolver, raise an exception + if not is_hint(annotation, resolver.base): + raise ArgumentError( + func, + parameter, + f"has bad annotation {annotation} wrt resolver {resolver}", + ) + + # If there's no value given and the resolver does not have a default, + # raise an exception + if not value and not resolver.default: + raise ArgumentError(func, parameter, "is missing from the arguments") + + # Look up the class. If value is None, we already checked the resolver + # should have a default and this will be okay + cls = resolver.lookup(value) + + # Recur on the __init__() function of the class looked up by the resolver, + # optionally using the related keyword arguments, if available. + return self.check_kwargs(cls, kwargs.get(related_key)) + + # If there's no value given and there's no default, raise an exception + if value_not_given: + if parameter.default is not parameter.empty: + # We're going to go ahead and assume that the default value + # matches the type annotation and not do any further checking + return True + raise ArgumentError( + func, + parameter, + f"no default, value not given. Signature: {inspect.signature(func)}", + ) + + # Now comes the nitty-gritty part of checking - if + + # get_origin() checks if the annotation is inside a Union, List, etc. + origin = get_origin(annotation) + + if origin is Union: + # If the things inside the union are not all simple types, + # i.e., float, int, bool, str, or None, then raise an error. + # This is because anything other than these types can't live + # inside JSON. + if not _is_simple_union(annotation): + raise ArgumentError(func, parameter, f"has inappropriate annotation: {annotation}") + + # Get the list of arguments inside the Union + args = get_args(annotation) + + # if the value is one of the args, then we're done. + if isinstance(value, args): + return True + # otherwise, raise a type error + raise ArgumentError( + func, + parameter, + f"value {value} does not match annotation {annotation}", + ) + + # If the origin is None, this means that it's just a single type + if origin is None: + try: + instance_check = isinstance(value, annotation) + except TypeError: + # This type error gets thrown if there's some reason you can't + # type check on the given annotation (i.e., it's not a subclass of `type`) + raise ArgumentError( + func, parameter, f"invalid annotation {annotation} ({type(annotation)})" + ) from None + + if instance_check: + return True + elif annotation == float and isinstance(value, int): + # you can coerce an int into a float, so just say this is fine + return True + else: + raise ArgumentError(func, parameter, f"{value} mismatched annotation {annotation}") + + # Unknown origin type + # TODO might need to extend to handle list/dict + raise ArgumentError(func, parameter, f"unhandled origin {origin}") def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: """Check the appropriate of the kwargs with a given function. @@ -111,85 +212,29 @@ def check_kwargs(self, func: Callable, kwargs: OptionalKwargs = None) -> bool: """ if kwargs is None: kwargs = {} - for key, parameter, related_key in _iter_params(func): - if key == "kwargs": - continue - annotation = parameter.annotation - next_resolver = self.names.get(key) - value = kwargs.get(key) - if next_resolver is not None: - if isinstance(next_resolver, FunctionResolver): - continue - if not is_hint(annotation, next_resolver.base): - raise ArgumentError( - func, - key, - f"has bad annotation {annotation} wrt resolver {next_resolver}", - ) - if value is None: - raise ArgumentError(func, key, f"is missing from the arguments") - self.check_kwargs( - next_resolver.lookup(value), - kwargs.get(related_key), - ) - else: - if value is None: - if parameter.default is parameter.empty: - raise ArgumentError( - func, - key, - f"no default, value not given. Signature: {inspect.signature(func)}", - ) - else: - origin = get_origin(annotation) - if origin is Union: - if not _is_simple_union(annotation): - raise ArgumentError( - func, key, f"has inappropriate annotation: {annotation}" - ) - else: - args = get_args(annotation) - if isinstance(value, args): - pass - else: - raise ArgumentError( - func, - key, - f"value {value} does not match annotation {annotation}", - ) - elif origin is None: - try: - instance_check = isinstance(value, annotation) - except TypeError: - raise ArgumentError( - func, key, f"invalid annotation {annotation} ({type(annotation)})" - ) from None - if instance_check: - pass - elif annotation == float and isinstance(value, int): - pass # log a warning? - else: - raise ArgumentError( - func, key, f"{value} mismatched annotation {annotation}" - ) - else: - raise ArgumentError(func, key, f"unhandled origin {origin}") + for parameter_name, parameter, related_key in _iter_params(func): + self.check_kwarg( + func=func, + parameter=parameter, + related_key=related_key, + kwargs=kwargs, + ) return True -def _iter_params( - func, -) -> Iterable[Tuple[str, inspect.Parameter, str]]: - parameters = inspect.signature(func).parameters - kwarg_map = {} - for key in parameters: - related_key = f"{key}_kwargs" - if related_key in parameters: - kwarg_map[related_key] = key - for key in parameters: - if key in kwarg_map: +def _iter_params(func) -> Iterable[Tuple[str, inspect.Parameter, str]]: + parameter_names: Mapping[str, inspect.Parameter] = inspect.signature(func).parameters + parameter_name_to_kwargs = {} + for parameter_name in parameter_names: + kwargs_key = f"{parameter_name}_kwargs" + if kwargs_key in parameter_names: + parameter_name_to_kwargs[kwargs_key] = parameter_name + for parameter_name, parameter in parameter_names.items(): + if parameter_name in parameter_name_to_kwargs: + continue + if parameter_name == "kwargs": continue - yield key, parameters[key], f"{key}_kwargs" + yield parameter_name, parameter, f"{parameter_name}_kwargs" def _main(): @@ -199,8 +244,9 @@ def _main(): from pykeen.experiments.cli import HERE from pykeen.losses import loss_resolver from pykeen.models import model_resolver - from pykeen.nn.emb import constrainer_resolver, normalizer_resolver from pykeen.nn.init import initializer_resolver + from pykeen.nn.representation import constrainer_resolver, normalizer_resolver + from pykeen.optimizers import optimizer_resolver from pykeen.pipeline import pipeline from pykeen.regularizers import regularizer_resolver @@ -222,12 +268,18 @@ def _main(): "relation_initializer": initializer_resolver, "relation_normalizer": normalizer_resolver, "relation_constrainer": constrainer_resolver, + # skip optimizer since its instantiation is dynamic + # "optimizer": optimizer_resolver, }, ) for path in HERE.glob("*/*.json"): data = json.loads(path.read_text()) kwargs = data["pipeline"] - r.check_kwargs(pipeline, kwargs) + try: + r.check_kwargs(pipeline, kwargs) + except Exception as e: + print(path) + raise e if __name__ == "__main__":