diff --git a/.style.yapf b/.style.yapf new file mode 100644 index 0000000..afdf3ea --- /dev/null +++ b/.style.yapf @@ -0,0 +1,17 @@ +[style] +based_on_style = pep8 + +align_closing_bracket_with_visual_indent = false +allow_split_before_default_or_named_assigns = false +allow_split_before_dict_value = false +coalesce_brackets = true +column_limit = 120 +continuation_align_style = fixed +dedent_closing_brackets = true +indent_dictionary_value = true +split_before_arithmetic_operator = true +split_before_dot = true +split_before_expression_after_opening_paren = true +split_before_first_argument = true +split_complex_comprehension = true +use_tabs = true diff --git a/.vscode/settings.json b/.vscode/settings.json index 721a059..290efb9 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -10,5 +10,8 @@ ], "python.testing.pytestEnabled": false, "python.testing.nosetestsEnabled": false, - "python.testing.unittestEnabled": true + "python.testing.unittestEnabled": true, + "python.testing.autoTestDiscoverOnSaveEnabled": true, + "python.formatting.provider": "yapf", + "python.formatting.yapfArgs": [] } diff --git a/imperial/cache.py b/imperial/cache.py deleted file mode 100644 index ef7cb2f..0000000 --- a/imperial/cache.py +++ /dev/null @@ -1,23 +0,0 @@ -from typing import Any - -class Cache: - value: Any - is_valid: bool = False - - def __init__(self): - self._invalidations = [] - - def invalidate(self): - self.is_valid = False - - def add_invalidation(self, invalidation): - if hasattr(invalidation, "_cache"): - invalidation = invalidation._cache - if isinstance(invalidation, Cache): - self._invalidations.append(invalidation) - else: - raise ValueError(invalidation) - - def cache(self, value): - self.value = value - self.is_valid = True diff --git a/imperial/core/base.py b/imperial/core/base.py index 0abf4f5..9e58838 100644 --- a/imperial/core/base.py +++ b/imperial/core/base.py @@ -1,47 +1,51 @@ +from __future__ import annotations + from copy import deepcopy -from typing import Callable, ClassVar, Dict, List, Optional, overload, Sequence, Union -from collections import OrderedDict +from typing import Any, Callable, cast, ClassVar, Dict, Iterator, List, Optional, overload, Sequence, Set, Union +from collections import defaultdict, OrderedDict +from ..magic import SpecialRef, NAME, BASIC +from ..linkmap import Linkable, LinkNode, StringLinkNode, LinkMap from ..exceptions import ImperialKeyError -PythonValue = Union[int, str, list] +PythonValue = Union[int, str, list, float, bool, bytes, tuple] EitherValue = Union[PythonValue, "ImperialType"] -class KeyMap(OrderedDict): + +class KeyMap(OrderedDict[str, Linkable]): """ Storage system for keys and values in a struct. Provides access to key data as well as methods to query the information. """ - def contains(self, name: str, memo: Dict[str, bool]) -> bool: - if name in memo: - return memo[name] + _owner: ImperialType + _reference_staging: Dict[str, Set[ImperialType]] + + def __init__(self, *args, owner: ImperialType): + super().__init__(*args) + self._owner = owner + self._reference_staging = defaultdict(set) + + def __setitem__(self, name: str, value: Linkable): + if name in self._reference_staging: + value.add_links(self._reference_staging[name]) + del self._reference_staging[name] + super().__setitem__(name, value) + + def __deepcopy__(self, memo: Dict[int, Any]): + ret = type(self)(owner=None) + memo[id(self)] = ret + for key, value in self.items(): + OrderedDict.__setitem__(ret, key, deepcopy(value, memo)) - ret = memo[key] = self.contains_quick(name) return ret - def contains_quick(self, name: str) -> bool: - if self.is_special_ref(name): - return self.special_ref_exists_quick(name) - return super().__contains__(name) - - def is_special_ref(self, ref: str) -> bool: - """ - Special refs are references to things that are either - not normally able to be referenced or locations that - are relative to this location, such as siblings or parents. - """ - return name.startswith("!") or "." in name - - def special_ref_exists(self, ref: str, memo: Dict[str, bool]) -> bool: - pass - - def special_ref_exists_quick(self, ref: str) -> bool: - pass + def is_ready(self, key: str) -> bool: + return super().__contains__(key) class Meta(type): def __new__(cls, name, bases, dct): - ret = type.__new__(cls, name, bases, dct) + ret = cast("ImperialType", type.__new__(cls, name, bases, dct)) for name, prop in dct.items(): if callable(prop) and getattr(prop, "_do_propagate", False): if name in ret.propagated_methods: @@ -51,52 +55,87 @@ def __new__(cls, name, bases, dct): class ImperialType(metaclass=Meta): - nocopy: ClassVar[List[str]] = ["clones", "parent", "container", "donor"] + # Whether or not get_basic can function for this class under some condition + # Override has_special_ref if there are any conditions in order to specify them + has_basic: ClassVar[bool] = False + + nocopy: ClassVar[List[str]] = [ + "propagated_methods", "clones", "_this", "parent", "benefactor", "container", "donor", "manager", "linkmap", + "caches" + ] propagated_methods: ClassVar[Dict[str, Callable]] = {} - name: Optional[str] - parent: Optional["ImperialType"] - context: Optional["ImperialType"] - container: Optional["ImperialType"] + name: LinkNode + link_prefix: str + + _this: Optional[ImperialType] + parent: Optional[ImperialType] + benefactor: Optional[ImperialType] + container: Optional[Any] + manager: Optional[ImperialType] keys: KeyMap - children: Dict[str, "ImperialType"] + children: Dict[str, ImperialType] + + donor: Optional[ImperialType] + clones: List[ImperialType] - donor: "ImperialType" - clones: List["ImperialType"] + linkmap: LinkMap + caches: Dict[str, LinkNode] - frozen: bool + # Pulled from basic node + add_link: Callable[[Any], None] + add_links: Callable[[Sequence], None] + invalidate: Callable[[Optional[Set[int]]], None] - def __init__(self, + def __init__( + self, data=None, *, name: Optional[str] = None, source=None, # TODO: type - children: Sequence["ImperialType"] = (), - hidden: bool = False, - parent: Optional["ImperialType"] = None, - context: Optional["ImperialType"] = None, - container: Optional["ImperialType"] = None, - donor: Optional["ImperialType"] = None + children: Sequence[ImperialType] = (), + this: Optional[ImperialType] = None, + parent: Optional[ImperialType] = None, + benefactor: Optional[ImperialType] = None, + container: Optional[Any] = None, + donor: Optional[ImperialType] = None, + manager: Optional[ImperialType] = None, ): """ + this: What @this should point to; None means self. parent: What @parent should point to. - context: Context that manages this struct - container: What this should inherit keys from first. + benefactor: What this should inherit keys from first. + container: Parent in a literal sense. ImperialType or Key. donor: What this was cloned from. + manager: The struct which controls this one's locator keys, if any. """ - self.name = name + self._this = this self.parent = parent - self.context = context + self.benefactor = benefactor self.container = container + self.manager = manager - self.keys = KeyMap() + self.keys = KeyMap(owner=self) self.children = OrderedDict() self.donor = donor self.clones = [] - self.frozen = False + n = name or str(id(self)) + lp = self.link_prefix = n if parent is None else parent.link_prefix + "{%s}" % (n, ) + lm = self.linkmap = LinkMap() if parent is None else self.root.linkmap + self.caches = {} + + lm[f"{lp}/name"] = self.name = StringLinkNode(name, rigid=True) + + if self.has_basic: + b = lm[f"{lp}/basic"] = self.caches["basic"] = LinkNode(refresh=self.refresh_basic) + self.add_link = b.add_link + self.add_links = b.add_links + self.invalidate = b.invalidate + + self.post_init() if data is not None: self.set(data) @@ -106,6 +145,12 @@ def __init__(self, self.add_children(children) + def post_init(self): + """ + Override this to hook into __init__ after setup before data assignment. + """ + pass + def __call__(self, data=None, **kwargs): """ Copy this struct with overridden settings. @@ -119,13 +164,18 @@ def __call__(self, data=None, **kwargs): if data is not None: new.set(data) - for kw in ( - "name", "source", "children", "hidden", - "parent", "container", "donor" - ): + if "this" in kwargs: + new._this = kwargs["this"] + + for kw in ("source", "children", "parent", "benefactor", "container", "donor", "manager"): if kw in kwargs: setattr(new, kw, kwargs[kw]) + if "name" in kwargs: + n = kwargs["name"] or str(id(self)) + new.link_prefix = n if kwargs["parent"] is None else kwargs["parent"].link_prefix + "{%s}" % (n, ) + new.name = StringLinkNode(new.link_prefix + "/name", kwargs["name"]) + return new def clone(self): @@ -142,16 +192,32 @@ def clone(self): return new def __deepcopy__(self, memo): + # Skip calling __init__ ret = object.__new__(self.__class__) memo[id(self)] = ret for attr, value in self.__dict__.items(): if attr in self.nocopy: setattr(ret, attr, value) else: - setattr(ret, attr, deepcopy(value, memo)) + v = deepcopy(value, memo) + if isinstance(v, KeyMap): + v._owner = self + setattr(ret, attr, v) return ret + @property + def root(self) -> ImperialType: + if self.parent is None: + return self + return self.parent.root + + @property + def this(self) -> ImperialType: + if self._this is None: + return self + return self._this + def get(self, names: Union[None, str, Sequence[str]] = None) -> PythonValue: """ Get the python value of a key. @@ -217,7 +283,7 @@ def set(self, *args): if len(names) == 1: names = names[0] elif names: - self.resolve(names[:-1]).set_key(names[-1], value) + self.resolve(names[:-1]).set_by_key(names[-1], value) return if names: self.set_by_key(names, value) @@ -230,7 +296,7 @@ def set(self, *args): else: self.set_basic(value) - def resolve(self, names: Union[None, str, Sequence[str]] = None) -> "ImperialType": + def resolve(self, names: Union[None, str, Sequence[str]] = None) -> ImperialType: """ Get the ImperialType value of a key. @@ -266,6 +332,11 @@ def get_by_key(self, name: str) -> PythonValue: return self.resolve_by_key(name).get() def get_basic(self) -> PythonValue: + if "basic" in self.caches: + return self.caches["basic"].value + raise NotImplementedError(f"no basic for {self.__class__.__name__}") + + def refresh_basic(self) -> PythonValue: """ Override this to implement retrieving a basic value for this struct. """ @@ -277,7 +348,7 @@ def set_by_key(self, name: str, value: EitherValue): the value of a single key. Typically will not need to be overridden. """ - self.keys[name] = ImperialType.normalize(value) + self.keys[name] = self.imperialize(value) def set_basic(self, value: EitherValue): """ @@ -289,7 +360,7 @@ def set_all(self, values: Dict[str, EitherValue]): for key, value in values.items(): self.set(key, value) - def resolve_by_key(self, name: str) -> "ImperialType": + def resolve_by_key(self, name: str) -> ImperialType: """ Override this to change the behavior of retrieving the ImperialType value of a single key. @@ -297,13 +368,20 @@ def resolve_by_key(self, name: str) -> "ImperialType": """ return self.key(name).data.resolve_basic() - def resolve_basic(self) -> "ImperialType": + def resolve_basic(self) -> ImperialType: """ Only Reference really needs to override this. Typically will not need to be overridden. """ return self + def add_links_to_keys(self, keys: Sequence[str], *, invalidates: Linkable): + for key in keys: + if self.keys.is_ready(key): + self.keys[key].add_link(invalidates) + else: + self.keys._reference_staging[key].add(invalidates) + def key(self, names: Union[str, Sequence[str]]): """ Get a key's instance from its name. @@ -321,17 +399,29 @@ def key(self, names: Union[str, Sequence[str]]): def has_keys(self, keys: Sequence[str]) -> bool: for key in keys: - if key not in self.keys: + if isinstance(key, SpecialRef): + if not self.has_special_ref(key): + return False + elif key not in self.keys: return False return True - def add_child(self, child: "ImperialType"): + def has_special_ref(self, ref: SpecialRef) -> bool: + if ref is NAME: + if self.name: + return True + elif ref is BASIC: + return self.has_basic + # If it's not known by us it's not here + return False + + def add_child(self, child: ImperialType): """ Override this in order to support having substructs. """ raise NotImplementedError(f"{self.__class__.__name__} must implement add_child") - def add_children(self, children: Sequence["ImperialType"]): + def add_children(self, children: Sequence[ImperialType]): """ Add multiple substructs at one time. """ @@ -342,15 +432,16 @@ def set_source(self, source): raise NotImplementedError("TODO: set_source") @classmethod - def normalize(cls, value) -> "ImperialType": + def imperialize(cls, value) -> ImperialType: if isinstance(value, ImperialType): return value elif isinstance(value, int): + from .number import Number return Number(value) elif isinstance(value, str): t = "string" elif isinstance(value, bytes): - return Bin(value) + t = "bin" elif isinstance(value, (list, tuple)): t = "list" elif isinstance(value, dict): @@ -359,6 +450,24 @@ def normalize(cls, value) -> "ImperialType": raise ValueError(value) raise NotImplementedError(t) + def parents(self) -> Iterator[ImperialType]: + parent = self.parent + while parent is not None: + yield parent + parent = parent.parent + + def benefactors(self) -> Iterator[ImperialType]: + benefactor = self.benefactor + while benefactor is not None: + yield benefactor + benefactor = benefactor.container + + def containers(self) -> Iterator[Any]: + container = self.container + while container is not None: + yield container + container = container.container + def __getattr__(self, name: str): if name in self.propagated_methods: self.propagated_methods[name] diff --git a/imperial/core/dynamic.py b/imperial/core/dynamic.py index b673e8f..8c2e94f 100644 --- a/imperial/core/dynamic.py +++ b/imperial/core/dynamic.py @@ -1,10 +1,12 @@ -from typing import Callable, ClassVar, Dict, Iterator, List, Optional, Sequence, Type -from collections import defaultdict, OrderedDict +from typing import Callable, ClassVar, Dict, List, Optional, Type +from collections import defaultdict from .key import Key -from .base import ImperialType, KeyMap, EitherValue +from .base import ImperialType, KeyMap, EitherValue, PythonValue from ..util import DotMap -from ..exceptions import ImperialKeyError +from ..exceptions import ImperialKeyError, ImperialLibraryError, ImperialTypeError + +Converter = Callable[[ImperialType], ImperialType] class DynamicKeyMap(KeyMap): @@ -15,22 +17,11 @@ class DynamicKeyMap(KeyMap): Keys may be inherited, have default values, or come from the result of calculations composed of other, defined keys. """ - _owner: Optional["Dynamic"] - - def __init__(self, *args, owner: Optional["Dynamic"] = None): - super().__init__(*args) - self._owner = owner - def __getitem__(self, name: str) -> Key: if name in self: return super().__getitem__(name) - inherited = self.find_inherited(name) - if inherited is not None: - ref = Reference(to=inherited) - ret = self[name] = self._owner._make_key(name, ref) - else: - ret = self[name] = self._owner._make_key(name) + ret = self[name] = self._owner._make_key(name) return ret def __setitem__(self, key: str, value: Key): @@ -47,54 +38,8 @@ def __contains__(self, name: str) -> bool: # TODO: works if super().__contains__(name): return True - return self.find_inherited(name) is not None + return self._owner.find_inherited(name) is not None - def contains(self, name: str, memo: Dict[str, bool]) -> bool: - """ - If you requested self[name], would it return a value? - """ - # TODO: cache result - if name in memo: - return memo[name] - - memo[name] = None - - if self.is_special_ref(name): - ret = memo[name] = self.special_ref_exists(name) - return ret - - if super().__contains__(name) or self.find_inherited(name) is not None: - return True - - kt = self._owner.key_type(name) - rems: List[set] = [] - for refs in kt._calc_links: - r = {ref for ref in refs if not self.contains_quick(ref)} - if r: - rems.append(r) - else: - return True - - rems.sort(key=len) - - for refs in rems: - if all(self.contains(name, memo) for ref in refs): - return True - return False - - def find_inherited(self, name: str) -> Key: - # TODO: do we want to consider calculated values of parents? - aliases = self._owner.localize_key(name) - for parent in self._owner.containers(): - for n in aliases: - if isinstance(parent, Dynamic): - n = parent.key_name_from_localization(n) - if n in parent.keys: - key = parent.keys[n] - if not key.hidden and not key.defaulted: - return key - return None - class Dynamic(ImperialType): """ @@ -103,7 +48,7 @@ class Dynamic(ImperialType): definition. Typically, the data is retrieved from an external source, but it may also be algorithmic, for example. - In a dynamic struct, keys are implicitely typed, may have + In a dynamic struct, keys are implicitly typed, may have defaults, can inherit their values from parent structs, and can have their values calculated from other, defined keys. @@ -119,17 +64,14 @@ class Dynamic(ImperialType): locators: ClassVar[Optional[Dict[str, Key]]] = None _overrides: ClassVar[Optional[Dict[Type[ImperialType], Dict[str, Key]]]] = None + _converters_to: ClassVar[Dict[Type, Converter]] + _converters_from: ClassVar[Dict[Type, Converter]] + keys: DynamicKeyMap - def __init__(self, data=None, *, children=(), **kwargs): - super().__init__(**kwargs) + def post_init(self): self._register() self.keys = DynamicKeyMap(owner=self) - - if data is not None: - self.set(data) - - self.add_children(children) @classmethod def register(cls, key: Type[Key]) -> Type[Key]: @@ -146,7 +88,7 @@ def register(cls, key: Type[Key]) -> Type[Key]: cls._keys[key.keyname] = key setattr(cls, key.__name__, key) return key - + @classmethod def register_locator(cls, key: Type[Key]) -> Type[Key]: """ @@ -178,75 +120,126 @@ def registrar(key: Type[Key]) -> Type[Key]: cls._overrides = defaultdict(dict) cls._overrides[context][key.keyname] = key return key + return registrar - + + @classmethod + def register_converter( + cls, fun: Optional[Converter] = None, *, source: Optional[Type] = None, target: Optional[Type] = None + ): + """ + Register a converter from a source to this or from this to a target. + The conversion function must take in an instance of the source type + and return a corresponding version of the target type. + + This can be used for simple conversions like string to number or + it can be used for more complex conversions like BMP to PNG. + + Currently only meant for reversible conversions. + TODO: Support recoverable, lossy, irreversible + TODO: Coercion vs conversion? + """ + if source is not None and target is not None: + raise ImperialLibraryError("cannot register a converter with both a source and target") + elif source is None and target is None: + raise ImperialLibraryError("registering a converter must specify either a source or target") + + def handler(fun: Converter): + if target is not None: + cls._converters_to[target] = fun + elif source is not None: + cls._converters_from[source] = fun + + if fun is not None: + handler(fun) + return + + return handler + def key_type(self, name: str) -> Type[Key]: """ Get a key's class from its name. """ - if self.context is not None: - ctx = type(self.context) + if self.manager is not None: + ctx = type(self.manager) if ctx in self._overrides: overrides = self._overrides[ctx] if name in overrides: return overrides[name] - elif name in self.context.locators: - return self.context.locators[name] + elif name in self.manager.locators: + return self.manager.locators[name] if name in self._keys: return self._keys[name] - raise ImperialKeyError(name) + raise ImperialKeyError(f"{name} of {self}") def _make_key(self, name: str, data=None) -> Key: return self.key_type(name)(data, name=name, container=self) - + def localize_key(self, name: str) -> List[str]: """ Get all localizations of a key name. """ # TODO: this return [name] - + def key_name_from_localization(self, name: str) -> str: """ Retrieve the internal name of a key from a localized name. """ # TODO: this return name - - def run_calculations(self, name: str): + + def check_constraints(self, name: Optional[str] = None): """ Run all registered calculations and assert that they have the same result. If the key was unset, set it. Raises ImperialSanityError if they do not. """ - if name in self.keys: - self.keys[name].run_calculations(self) - else: - key = self.keys[name] = self._make_key(name) - key.run_calculations(self) + if name is None: + for name in self.keys: + self.check_constraints(name) + return + + self.keys[name].check_constraints() @classmethod - def normalize(cls, value) -> ImperialType: + def normalize(cls, value: EitherValue) -> PythonValue: + """ + Unify multiple possible basic values into a single form + of basic value or a dict of keys. + """ raise NotImplementedError(f"{cls.__name__} must implement normalize") def set_by_key(self, name: str, value: EitherValue): - # TODO: assert when frozen - key = self._make_key(name) - key.data = key.type(value) - self.keys[name] = key - - def containers(self) -> Iterator[ImperialType]: - container = self.container - while container is not None: - yield container - container = container.container - - def parents(self) -> Iterator[ImperialType]: - parent = self.parent - while parent is not None: - yield parent - parent = parent.parent - - def add_link(self, key: str, *, invalidates=None): - if invalidates is not None: - pass # TODO + self.keys[name].set(value) + + def convert_to(self, type: Type[ImperialType]) -> ImperialType: + """ + Convert this struct into another struct type. + Override this in order to do more generalized conversions. + """ + if type in self._converters_to: + return self._converters_to[type](self) + raise ImperialTypeError(f"no conversion from {self.__class__.__name__} to {type.__name__} known") + + def convert_from(self, data: ImperialType) -> ImperialType: + """ + Convert another struct into this struct type. + Override this in order to do more generalized conversions. + """ + type_data = type(data) + if type_data in self._converters_from: + return self._converters_from[type_data](data) + return data.convert_to(type(self)) + + def find_inherited(self, name: str) -> Key: + aliases = self.localize_key(name) + for benefactor in self.benefactors(): + for n in aliases: + if isinstance(benefactor, Dynamic): + n = benefactor.key_name_from_localization(n) + if n in benefactor.keys: + key = benefactor.keys[n] + if not key.hidden and not key.defaulted: + return key + return None diff --git a/imperial/core/key.py b/imperial/core/key.py index f34f1a9..519b5a2 100644 --- a/imperial/core/key.py +++ b/imperial/core/key.py @@ -1,9 +1,15 @@ -from typing import Any, Callable, ClassVar, List, Optional, overload, Set, Tuple +from typing import Any, Callable, ClassVar, List, Optional, overload, Sequence, Set, Tuple, Type +from operator import attrgetter +from functools import reduce from .base import ImperialType, EitherValue -from ..magic import add_help, make_refs_only_resolver +from ..magic import add_help, make_container_resolver, ReferenceHandler +from ..linkmap import LinkNode from ..exceptions import ImperialKeyError, ImperialSanityError +NO_DEFAULT = object() + + class KeyMeta(type): def __new__(cls, name, bases, dct): ret = type.__new__(cls, name, bases, dct) @@ -17,10 +23,12 @@ def __new__(cls, name, bases, dct): elif getattr(value, "_is_calculation", False): if ret._calculations: ret._calculations.append(value) - ret._calc_links |= value._refs + ret._calc_links.add(value._refs) else: ret._calculations = [value] ret._calc_links = value._refs.copy() + ret._estimations = tuple(ret._estimations) + ret._calculations = tuple(ret._calculations) return ret @@ -29,93 +37,127 @@ class Key(metaclass=KeyMeta): A key for a dynamic struct. """ # Must define keyname + type: ClassVar[Type[ImperialType]] keyname: ClassVar[Optional[str]] - default: ClassVar[Optional[Any]] = None + default: ClassVar[Any] = NO_DEFAULT _estimations: ClassVar[List[Callable]] = [] _calculations: ClassVar[List[Callable]] = [] - _calc_links: Set[str] = set() + _calc_links: ClassVar[Set[str]] = set() - _data: ImperialType = None + _data: LinkNode defaulted: bool = False - def __init__( - self, - data=None, - *, - name: str = "", - container: Optional[ImperialType] = None - ): + name: str + container: Optional[ImperialType] + + # Pulled from node + add_link: Callable[[Any], None] + add_links: Callable[[Sequence], None] + invalidate: Callable[[Optional[Set[int]]], None] + + def __init__(self, data=None, *, name: str = "", container: Optional[ImperialType] = None): self.name = name or self.keyname self.container = container + self._data = LinkNode(refresh=self._refresh_basic, rigid=data is not None) + self.add_link = self._data.add_link + self.add_links = self._data.add_links + self.invalidate = self._data.invalidate if data is not None: self.set(data) @property def data(self): - if self._data is None: - if self._calculations: - self.run_calculations(self.container) - if self._data is None: - if self.default is None: - raise ImperialKeyError(self.name) - self._data = self.type(self.default) - self.defaulted = True - return self._data - + return self._data.value + @data.setter def data(self, value: ImperialType): # TODO: type checking, superset casting? - self._data = value + self._data.value = value(container=self) - def set(self, value: EitherValue): - if self.frozen >= 1: - self.data.freeze() - self.data.set(self.normalize(value)) + @data.deleter + def data(self): + self._data.invalidate() + + def _refresh_basic(self): + inherited = self.container.find_inherited(self.name) + if inherited is not None: + self._data.set_links_out({inherited}, self.container.linkmap.parents()) + # TODO: conversions, or is this in the DynamicKeyMap? + return self.imperialize(inherited()) + else: + is_valid, default = self.get_default() + if not is_valid: + raise ImperialKeyError(self.name) + self._data.value = default + self.defaulted = True + default.add_links_to_keys(self._calc_links, invalidates=default) + return default + + def get_default(self) -> Tuple[bool, ImperialType]: + """ + Returns the validity of the value and the default value. + Only override this in order to produce complex default values. + For calculated values you can define @calculate methods. + For static values, assign the value to the ClassVar `default` + and an ImperialType to the ClassVar `type`. + """ + it = (calc(self) for calc in self._calculations if self.container.has_keys(calc._refs)) + try: + res = next(it) + except StopIteration: + pass else: - self.data = self.type(self.normalize(value)) - self.defaulted = False - + first_value = self.imperialize(res) + # TODO: should this set this, parent, etc? + base = self.type(first_value, benefactor=None, container=self) + if any(base != x for x in it): + raise ImperialSanityError() + return True, base + + if self.default is NO_DEFAULT: + return False, None + return True, self.type(self.default) + + def set(self, value: EitherValue): + self.data = self.imperialize(value) + self.defaulted = False + def resolve(self) -> ImperialType: return self.data.resolve() @classmethod - def normalize(cls, value): - # Pass-through to type - return value - - def run_calculations(self, source): - iterator = ( - calc(self, source) - for calc in self._calculations - if source.has_keys(calc._refs) - ) - - if self._data is not None: - if any(self._data != x for x in iterator): - raise ImperialSanityError() - else: - try: - res = next(iterator) - except StopIteration: - return - - first_value = self.normalize(res) - base = self.type(first_value) - if any(base != x for x in iterator): - raise ImperialSanityError() - self.data = base + def imperialize(cls, value: EitherValue) -> ImperialType: + """ + This is run by set in order to convert the values to a standard + ImperialType via the `type` ClassVar. + Override this entirely if there are any translations to be done + between what is set to a key vs the struct type it actually holds. + For instance, if this should return a tuple-type list, but can + accept one of the members when defined and default the other(s). + """ + if isinstance(value, cls.type): + return value + return cls.type(value) + + def check_constraints(self): + it = (calc(self) for calc in self._calculations if self.container.has_keys(calc._refs)) + if not reduce(lambda x, y: x == y, it): + raise ImperialSanityError() + @overload def calculate(*args: str) -> Callable[[Callable], Callable[[Optional[ImperialType]], Any]]: ... + @overload def calculate(fun: Callable) -> Callable[[Optional[ImperialType]], Any]: ... -def calculate(*args): + +def calculate(*args, estimation=False): """ Decorator for defining a method of calculating a Key from other keys. The arguments defined in the method @@ -138,19 +180,40 @@ def from_whatever(basic, *, k1, k2): refs: Tuple[str] = () def wrapper(fun: Callable) -> Callable[[Optional[ImperialType]], Any]: - resolver = make_refs_only_resolver(fun, refs) - def handler(self, source: Optional[ImperialType] = None) -> Any: - if source is None: - source = self.container - assert source is not None - return resolver(source).run() - handler._is_calculation = True - resolver.add_to(handler) + resolver = ReferenceHandler.from_method_using_args(fun, positional=refs) + handler = make_container_resolver(resolver) + + if estimation: + handler._is_estimation = True + else: + handler._is_calculation = True + add_help(handler, fun) + resolver.add_to(handler) return handler - + if len(args) == 1 and callable(args[0]): return wrapper(args[0]) else: refs = args return wrapper + + +@overload +def estimate(*args: str) -> Callable[[Callable], Callable[[Optional[ImperialType]], Any]]: + ... + + +@overload +def estimate(fun: Callable) -> Callable[[Optional[ImperialType]], Any]: + ... + + +def estimate(*args): + """ + Works like calculate, but does not assert anything about + the sanity of the data. They're simply used if there's no + other recourse for data inferrence and it's necessary to + do so. + """ + return calculate(*args, estimation=True) diff --git a/imperial/core/number.py b/imperial/core/number.py new file mode 100644 index 0000000..8e704d2 --- /dev/null +++ b/imperial/core/number.py @@ -0,0 +1,108 @@ +from .key import Key, estimate +from .base import PythonValue +from .value import Value, number +from .serializable import Serializable, serialize, unserialize +from ..util import BytesBuffer +from ..exceptions import ImperialSanityError, ImperialSerializationError, ImperialTypeError + + +class BaseNumber(Value): + """ + For types to inhereit when they represent abstract numbers. + That is, when they're not packable. + However, note that Number does subclass this. + """ + @classmethod + def _register(cls): + super()._register() + + # Intuitions + @cls.register + class Min(Key): + """ + The actual value of this {type} cannot be less than min. + This is an inclusive lower bound. + """ + + type = BaseNumber + keyname = "min" + + @estimate + def strictly_equal(self, data: BaseNumber) -> int: + return data.number() + + @cls.register + class Max(Key): + """ + The actual value of this {type} cannot be greater than max. + This is an inclusive upper bound. + """ + + type = BaseNumber + keyname = "max" + + @estimate + def strictly_equal(self, data: BaseNumber) -> int: + return data.number() + + @classmethod + def normalize(cls, value: PythonValue): + if isinstance(value, int): + return value + raise ImperialTypeError(value, expects=int) + + @number + def number(self, data: PythonValue) -> int: + return data + + +class Number(BaseNumber, Serializable): + """ + The fundamental numerical atom. + """ + @serialize + def serialize(self, blob: BytesBuffer, *, data, endian, sign): + # Whole bytes only + nbits = blob.len_bits() + if nbits % 8 != 0: + raise ImperialSerializationError("{type} cannot serialize partial bytes") + + value = data.number() + nbytes = len(blob) + signed = sign.get() is Number.Sign.SIGNED + + # Only signed values can be negative + if not signed and value < 0: + raise ImperialSanityError("unsigned numbers cannot be negative") + + try: + b = value.to_bytes(nbytes, endian.string(), signed=signed) + except OverflowError: + raise ImperialSanityError( + "{type} {extra.value} is too big for {extra.bytes}", + extra={ + "value": value, + "bytes": self.string("size"), + }, + ) from None + + blob.write(b) + + @unserialize + def unserialize(self, blob: BytesBuffer, *, endian, sign): + signed = sign.get() is Number.Sign.SIGNED + value = int.from_bytes(blob.read(), endian.string(), signed=signed) + return value + + # @stringify + # def stringify(self, data): + # # TODO: form + # return str(data.number()) + + # @parse + # def parse(self, string: str): + # # TODO: form + # try: + # return int(string) + # except ValueError: + # raise ImperialParsingError.GenericNotValid(string) from None diff --git a/imperial/core/packable.py b/imperial/core/packable.py new file mode 100644 index 0000000..ade5d66 --- /dev/null +++ b/imperial/core/packable.py @@ -0,0 +1,33 @@ +from collections import defaultdict + +from .base import Meta +from .dynamic import Dynamic +from ..magic import SpecialRef, PACKED + + +class PackableMeta(Meta): + def __new__(cls, name, bases, dct): + ret = type.__new__(cls, name, bases, dct) + tmp = defaultdict(lambda: [None, None]) + for value in dct.values(): + if callable(value): + mode: str + direction: int + refs: set = getattr(value, "_refs", set()) + mode, direction = getattr(value, "_pack", ("", -10)) + if mode and refs: + tmp[mode][direction] = refs + + # TODO: frozen dict + ret._pack_links = dict(tmp) + + return ret + + +class Packable(Dynamic, metaclass=PackableMeta): + _pack_links = {} + + def has_special_ref(self, ref: SpecialRef) -> bool: + if ref is PACKED: + return True + return super().has_special_ref(ref) diff --git a/imperial/core/serializable.py b/imperial/core/serializable.py new file mode 100644 index 0000000..e110d14 --- /dev/null +++ b/imperial/core/serializable.py @@ -0,0 +1,110 @@ +from typing import Any, Iterator, Optional, Set, Tuple, Union + +from .packable import Packable +from ..util import BytesBuffer +from ..magic import ReferenceHandler, CachingReferenceHandler +from ..linkmap import BigBlobLinkNode + + +def serialize(fun): + resolver = CachingReferenceHandler.from_method_using_kwargs(fun, "packed") + + def handler(self: "Serializable", *args: BytesBuffer) -> Optional[bytes]: + """ + Serialize this struct into a bytes sequence. + """ + if len(args) == 1: + resolver(self, args[0]) + return + elif not args: + blob = BytesBuffer(bits=self.number(("size", "bits"))) + resolver(self, blob) + blob.seek(0) + return blob.readall() + raise TypeError(f"serialize() takes from 0 to 1 positional arguments but {len(args)} were given") + + resolver.add_to(handler) + handler._pack = ("Serializable", 0) + return handler + + +def unserialize(fun): + resolver = ReferenceHandler.from_method_using_kwargs(fun) + + def handler(self: "Serializable", blob: Union[bytes, BytesBuffer] = b"", until: Set[str] = {""}): + """ + Unserialize this struct from a bytes sequence. + """ + if not blob: + return + + if isinstance(blob, bytes): + blob = BytesBuffer(blob) + value = resolver(self, blob) + self.set(value) + return self + + resolver.add_to(handler) + handler._pack = ("Serializable", 1) + return handler + + +def unserialize_yield(fun): + last_blob: BytesBuffer + last_position: int + last_generator: Iterator[Tuple[str, Any]] + + resolver = ReferenceHandler.from_method_using_kwargs(fun) + + def handler(self: "Serializable", blob: Union[bytes, BytesBuffer] = b"", until: Set[str] = {""}): + """ + Unserialize this struct from a bytes sequence. + Stop unserializing when all keys in "until" are satisfied. + By default, pulls all keys it can. + """ + nonlocal last_blob, last_position, last_generator + if not blob: + try: + blob = last_blob + except NameError: + return + blob.seek(last_position) + elif isinstance(blob, bytes): + last_blob = blob = BytesBuffer(blob) + last_generator = resolver(self, blob) + + # Clear what's already been defined + until = {key for key in until if key not in self.keys} + + if until: + for key, value in last_generator: + if key: + self.set(key, value) + else: + self.set(value) + until.discard(key) + + if not until: + break + + last_position = blob.tell() + return self + + resolver.add_to(handler) + handler._pack = ("Serializable", 1) + return handler + + +class Serializable(Packable): + def post_init(self): + super().post_init() + cp = self.caches["packed"] = BigBlobLinkNode(refresh=self.serialize) + self.linkmap[self.link_prefix + "/packed"] = cp + + @serialize + def serialize(self, blob: BytesBuffer): + raise NotImplementedError(f"{self.__class__.__name__} must implement serialize") + + @unserialize + def unserialize(self, blob: BytesBuffer): + raise NotImplementedError(f"{self.__class__.__name__} must implement unserialize") diff --git a/imperial/core/value.py b/imperial/core/value.py new file mode 100644 index 0000000..94f1b5e --- /dev/null +++ b/imperial/core/value.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +from typing import Callable, cast, List, Optional, Sequence, Tuple, Union + +from .key import Key +from .base import propagate, EitherValue, ImperialType, PythonValue +from .dynamic import Dynamic +from ..magic import add_help +from ..exceptions import ImperialKeyError, ImperialValueError + +OptionalStrSeq = Union[None, str, Sequence[str]] + + +def get_data(self: Value) -> Optional[ImperialType]: + data = self.key("data").data + if data is self: + return None + if isinstance(data, type(self)): + return data + return self.convert_from(data) + + +def number(fun: Callable[[Value, PythonValue], int]) -> Callable[[Value, OptionalStrSeq], int]: + def handler(self: Value, names: OptionalStrSeq = None) -> int: + if names is not None: + return self.resolve(names).number() + + data = get_data(self) + if data is not None: + return data.number() + return fun(self, self.caches["basic"].value) + + add_help(handler, fun) + return handler + + +def string(fun: Callable[[Value, PythonValue], str]) -> Callable[[Value, OptionalStrSeq], str]: + def handler(self: Value, names: OptionalStrSeq = None) -> str: + if names is not None: + return self.resolve(names).string() + + data = get_data(self) + if data is not None: + return data.string() + return fun(self, self.caches["basic"].value) + + add_help(handler, fun) + return handler + + +def list(fun: Callable[[Value, PythonValue], List]) -> Callable[[Value, OptionalStrSeq], List]: + def handler(self: Value, names: OptionalStrSeq = None) -> List: + if names is not None: + return self.resolve(names).list() + + data = get_data(self) + if data is not None: + return data.list() + return fun(self, self.caches["basic"].value) + + add_help(handler, fun) + return handler + + +class Value(Dynamic): + has_basic = True + + @propagate + @number + def number(self, data: PythonValue) -> int: + """ + Return the python int form of this value. + """ + raise ImperialValueError("not a number") + + @propagate + @string + def string(self, data: PythonValue) -> str: + """ + Return the python str form of this value. + """ + raise ImperialValueError("not a string") + + @propagate + @list + def list(self, data: PythonValue) -> List: + """ + Return the python list form of this value. + """ + raise ImperialValueError("not a list") + + # TODO: value proxy support + @classmethod + def _register(cls): + @cls.register + class DataKey(Key): + keyname = "data" + + _py_data: Optional[PythonValue] = None + + def set(self, value: EitherValue): + if isinstance(value, ImperialType): + self._py_data = None + self.data = value + # TODO: reference group + self.container.add_link(value.caches["basic"]) + value.add_link(self.container.caches["basic"]) + else: + # Container will normalize this when it retrieves it + self._py_data = value + cb = self.container.caches["basic"] + vb = cast(ImperialType, self.data).caches["basic"] + cb.remove_link(vb) + vb.remove_link(cb) + del self.data + self.defaulted = False + + def get_default(self) -> Tuple[bool, ImperialType]: + return True, self.container + + def refresh_basic(self) -> PythonValue: + data = self.key("data") + if data._py_data is None: + return self.get_by_proxy(data.data) + return self.get_primitive(data._py_data) + + def get_by_proxy(self, data: ImperialType) -> PythonValue: + if isinstance(data, type(self)): + if data is self: + raise ImperialKeyError("data") + return data.get_basic() + return self.convert_from(data).get() + + def get_primitive(self, data: PythonValue) -> PythonValue: + return data + + def set_basic(self, value: EitherValue): + if isinstance(value, ImperialType): + self.set_by_key("data", value) + else: + self.caches["basic"].value = self.normalize(value) + + __int__ = number + __str__ = string diff --git a/imperial/exceptions.py b/imperial/exceptions.py index 443cbc5..639cd59 100644 --- a/imperial/exceptions.py +++ b/imperial/exceptions.py @@ -1,14 +1,21 @@ # TODO: be able to interpret container, key name, line/col, etc as appropriate +# some sort of nice exceptions system + class ImperialError(Exception): + """ + Base exception for errors from the Imperial system. + """ pass + class ImperialSanityError(ImperialError): """ Raised when there are conflicts in the description. """ pass + class ImperialLibraryError(ImperialError): """ Raised by an error caused by a problem in a @@ -16,8 +23,34 @@ class ImperialLibraryError(ImperialError): """ pass + class ImperialKeyError(ImperialError): """ Raised when a non-existent key was requested. """ pass + + +class ImperialTypeError(ImperialError): + """ + Like a TypeError but caused by the Imperial system. + """ + def __init__(self, value, expects): + self.value = value + self.expects = expects + super().__init__(value, expects) + + +class ImperialValueError(ImperialError): + """ + Like a ValueError but caused by the Imperial system. + """ + pass + + +class ImperialSerializationError(ImperialError): + """ + Raised when un/serialization is made impossible by the + current configuration of the struct. + """ + pass diff --git a/imperial/linkmap.py b/imperial/linkmap.py new file mode 100644 index 0000000..20c161c --- /dev/null +++ b/imperial/linkmap.py @@ -0,0 +1,220 @@ +from __future__ import annotations + +from typing import Any, Callable, Dict, Optional, Protocol, Set, Sequence +from weakref import ref, WeakValueDictionary, WeakSet +from collections import defaultdict + + +class Linkable(Protocol): + # TODO: origin/origins should be Linkables, pylance can't handle it tho + def add_link(self, origin): + ... + + def add_links(self, origins: Sequence): + ... + + def invalidate(self, memo: Optional[Set[int]] = None): + ... + + +class LinkMap(WeakValueDictionary[str, "BaseLinkNode"]): + """ + Access the overall link tree and wait for nodes to be created. + Keys should be of the form: + BaseStruct{SubStruct}.keyname/basic + RHS of the slash can be: + * name - the name of the struct as specified after the type + * basic - the basic value of the struct/key + * packed - the packed form of the struct (serialiazed, stringified, etc) + """ + parent: Optional[LinkMap] + + def __init__(self, *args, parent: Optional[LinkMap] = None, **kwargs): + super().__init__(*args, **kwargs) + self.parent = parent + self.staged = defaultdict(WeakSet) + + def __setitem__(self, name: str, node: BaseLinkNode): + if name in self: + del self[name] + + # TODO: check if name exists in parent, and steal nodes with self in map hierarchy + + super().__setitem__(name, node) + + # Link anything waiting to be linked to this + if name in self.staged: + node.references_in.update(self.staged[name]) + del self.staged[name] + # TODO: delete out of parents too + + def __delitem__(self, name: str): + # Remove any node that starts with name and throw links in staged + # TODO: more efficient? + name, *ex = name.rsplit("/", 1) + deletables = [] + for key, node in self.items(): + if key.startswith(name): + # TODO: also add to parents? + self.staged[key].update(node.references_in) + deletables.append(key) + + for key in deletables: + super().__delitem__(key) + + def add_reference(self, origin: BaseLinkNode, target: str): + if target in self.nodes: + self[target].add_reference(origin) + else: + self.staged[target].add(origin) + if self.parent is not None: + self.parent.add_reference(origin, target) + + def add_references(self, origin: BaseLinkNode, targets: Sequence[str]): + for name in targets: + if name in self.nodes: + self[name].add_reference(origin) + else: + self.staged[name].add(origin) + if self.parent is not None: + self.parent.add_reference(origin, name) + + def remove_references(self, origin: BaseLinkNode, targets: Sequence[str]): + for name in targets: + if name in self.nodes: + self[name].remove_reference(origin) + if name in self.staged: + self.staged[name].remove(origin) + if self.parent is not None: + self.parent.remove_references(origin, targets) + + def parents(self): + parent = self.parent + while parent is not None: + yield parent + parent = parent.parent + + +class BaseLinkNode: + value: Any = None + valid: bool = False + rigid: bool + + def __init__(self, *, rigid: bool = False): + self.rigid = rigid + + self.links_in: WeakSet[BaseLinkNode] = WeakSet() + self.references_in: WeakSet[BaseLinkNode] = WeakSet() + self.references_out: Set[str] = set() + + def add_link(self, origin: Linkable): + self.links_in.add(origin) + + def add_links(self, origins: Sequence[Linkable]): + self.links_in.update(origins) + + def add_reference(self, origin): + self.references_in.add(origin) + + def remove_link(self, origin): + self.links_in.remove(origin) + + def remove_reference(self, origin): + self.references_in.remove(origin) + + def set_links_out(self, names: Sequence[str], maps: Sequence[LinkMap]): + snames = set(names) + added = snames - self.references_out + removed = self.references_out - snames + + if removed: + # Clear them out of the targets + for lmap in maps: + lmap.remove_references(self, removed) + + if added: + for lmap in maps: + lmap.add_references(self, added) + + self.references_out = snames + + def invalidate(self, memo: Optional[Set[int]] = None): + if not self.rigid: + if memo is None: + memo = set() + + if id(self) not in memo: + memo.add(id(self)) + self.valid = False + + for x in self.links_in: + x.invalidate(memo) + for x in self.references_in: + x.invalidate(memo) + + +class LinkNode(BaseLinkNode): + def __init__(self, *value, refresh, **kwargs): + self.refresh = refresh + + if value: + self._value = value[0] + self.valid = True + else: + self._value = None + self.valid = False + + super().__init__(**kwargs) + + @property + def value(self): + if not self.valid: + self._value = self.refresh() + self.valid = True + return self._value + + @value.setter + def value(self, value): + if self._value is not value: + self._value = value + self.valid = True + + memo = {id(self)} + for x in self.links_in: + x.invalidate(memo) + for x in self.references_in: + x.invalidate(memo) + + +class StringLinkNode(BaseLinkNode): + def __init__(self, value: str, **kwargs): + super().__init__(**kwargs) + self._value = value + + @property + def value(self) -> str: + return self._value + + @property + def valid(self) -> bool: + return True + + +class BigBlobLinkNode(LinkNode): + def __init__(self, *value, refresh, **kwargs): + super().__init__(*map(ref, value), refresh=refresh, **kwargs) + + @property + def value(self): + if not self.valid: + self._value = ref(self.refresh()) + self.valid = True + return self._value() + + @value.setter + def value(self, value): + if self._value is not value: + self._value = ref(value) + self.valid = True + for x in self.connections: + x.invalidate() diff --git a/imperial/magic.py b/imperial/magic.py index dd60313..c6a3fc1 100644 --- a/imperial/magic.py +++ b/imperial/magic.py @@ -1,87 +1,121 @@ +from __future__ import annotations + import inspect -from typing import Callable, Dict, Set, Tuple, Union +from typing import Any, Callable, cast, Dict, Optional, Protocol, Set, Tuple, Union, TYPE_CHECKING +from operator import attrgetter + +if TYPE_CHECKING: + from .core.base import ImperialType + + +class SpecialRef: + def __init__(self, name: str, getter: Optional[Callable] = None): + self.name = name + self.getter = getter + + def __repr__(self): + return f"SpecialRef('{self.name}')" + + +NAME = SpecialRef("name", attrgetter("name")) +BASIC = SpecialRef("basic", lambda x: x.caches["basic"]) +PACKED = SpecialRef("packed", lambda x: x.caches["packed"]) + +REF = Union[str, SpecialRef] +InsConv = Optional[Callable[[Any], "ImperialType"]] -from .cache import Cache -from .core.base import ImperialType -from .exceptions import ImperialLibraryError def add_help(to: Union[Callable, type], source: Union[Callable, type]): to.__name__ = source.__name__ to.__doc__ = source.__doc__ -def make_refs_resolver(fun: Callable) -> "ReferenceHandler": - """ - Transform all keyword-only arguments into requests for - keys by name. Sets up cache links, too. - """ - keys = {key: key for key in inspect.getfullargspec(fun).kwonlyargs} - return ReferenceHandler(fun, keys) - -def make_refs_only_resolver(fun: Callable, positional: Tuple[str] = ()) -> "ReferenceHandler": - """ - Transform all arguments into requests for - keys by name. Sets up cache links, too. - """ - spec = inspect.getfullargspec(fun) - nonself_args = spec.args[1:] - keys = dict(zip(positional, nonself_args)) - for key in (spec.kwonlyargs if keys else nonself_args): - keys[key] = key - return ReferenceHandler(fun, keys) -class ReferenceHandler: - _has_run = False - _other: ImperialType +class HasContainer(Protocol): + container: ImperialType + + +def make_container_resolver(rh: ReferenceHandler) -> Callable: + def handler(thing: HasContainer, *args, **kwargs): + return rh(thing.container, *args, **kwargs) + + return handler + +class ReferenceHandler: _fun: Callable - _keys: Dict[str, str] - _cache: Cache + _keys: Dict[REF, str] + _refs: Set[REF] + _to_instance: InsConv - def __init__(self, fun: Callable, keys: Dict[str, str]): + def __init__(self, fun: Callable, keys: Dict[REF, str], *, to_instance: InsConv = None): self._fun = fun self._keys = keys - self._cache = Cache() - - def __call__(self, other: ImperialType): - if self._has_run: - if other is not self._other: - raise ImperialLibraryError( - f"{self.__class__.__name__} must only be used for a single struct") - else: - for key in self._keys.keys(): - other.add_link(key, invalidates=self._cache) - self._has_run = True - self._other = other - return self - - def run(self, *args, **kwargs): - if self._cache.is_valid: - return self._cache.value + self._refs = set(keys.keys()) + self._to_instance = to_instance + + @classmethod + def from_method_using_args(cls, fun: Callable, *args, positional: Tuple[REF] = (), **kwargs) -> ReferenceHandler: + """ + Transform all arguments into requests for + keys by name. Sets up cache links, too. + """ + spec = inspect.getfullargspec(fun) + nonself_args = spec.args[1:] + keys = dict(zip(positional, nonself_args)) + for key in (spec.kwonlyargs if keys else nonself_args): + keys[key] = key + return cls(fun, keys, *args, **kwargs) + + @classmethod + def from_method_using_kwargs(cls, fun: Callable, *args, **kwargs) -> ReferenceHandler: + """ + Transform all keyword-only arguments into requests for + keys by name. Sets up cache links, too. + """ + keys = {key: key for key in inspect.getfullargspec(fun).kwonlyargs} + return cls(fun, keys, *args, **kwargs) + + def __call__(self, instance: ImperialType, *args, **kwargs): keyargs = {} - other = self._other for origin, kwarg in self._keys.items(): - if origin.startswith("!"): - if origin == "!children": - if other.has_children(): - keyargs[kwarg] = other.children() - else: - return None - elif origin == "!basic": - keyargs[kwarg] = other.resolve_basic() - else: - raise ImperialLibraryError(f"Unknown internal reference {key}") - elif origin in other.keys: - # TODO: use link map to make this safer? - keyargs[kwarg] = other.keys[origin].resolve() + if isinstance(origin, SpecialRef): + # Referencing something of the struct that's not a key + keyargs[kwarg] = origin.getter(instance).value + elif origin in instance.keys: + keyargs[kwarg] = instance.keys[origin].resolve() else: return None - ret = self._fun(self._other, *args, **kwargs, **keyargs) - self._cache.cache(ret) - return ret - - def keys(self) -> Set[str]: - return set(self._keys.keys()) + + return self._fun(instance, *args, **kwargs, **keyargs) def add_to(self, handler): - handler._refs = self.keys() - handler._cache = self._cache + handler._refs = self._refs + + +class CachingReferenceHandler(ReferenceHandler): + _cache_name: str + + def __init__(self, fun: Callable, keys: Dict[REF, str], cache_name: str, **kwargs): + super().__init__(fun, keys, **kwargs) + self._cache_name = cache_name + + def __call__(self, instance: ImperialType, *args, **kwargs): + node = instance.caches[self._cache_name] + if node.valid: + return node.value + + keyargs = {} + for origin, kwarg in self._keys.items(): + if isinstance(origin, SpecialRef): + # Referencing something of the struct that's not a key + to_link = origin.getter(instance) + keyargs[kwarg] = to_link.value + to_link.add_link(node) + elif origin in instance.keys: + to_link = keyargs[kwarg] = instance.keys[origin].resolve() + to_link.add_link(node) + else: + return None + + node.value = ret = self._fun(instance, *args, **kwargs, **keyargs) + return ret diff --git a/imperial/util.py b/imperial/util.py index 702be35..537aff1 100644 --- a/imperial/util.py +++ b/imperial/util.py @@ -1,5 +1,230 @@ -from typing import Any +from io import BufferedIOBase, RawIOBase, SEEK_SET, SEEK_CUR, SEEK_END +from typing import Any, Callable, Literal, Optional, Union + +SeekWhence = Literal[SEEK_SET, SEEK_CUR, SEEK_END] + class DotMap(dict): def __getattr__(self, name: str) -> Any: return self[name] + + +class RawBytesIO(RawIOBase): + """ + A BytesIO-like object based on RawIOBase. + """ + blob: bytearray + _cursor: int + _length: int + closed: bool + + def __init__(self, blob: bytes): + self.blob = bytearray(blob) + self._length = len(blob) + self._cursor = 0 + + def _raise_if_closed(self): + if self.closed: + raise ValueError("I/O operation on closed file.") + + def readable(self) -> bool: + return True + + def writable(self) -> bool: + return True + + def seekable(self) -> bool: + return True + + def isatty(self) -> bool: + return False + + def read(self, size: int = -1) -> bytes: + self._raise_if_closed() + if size == -1: + return self.readall() + elif size < 0: + raise ValueError("size") + start = self._cursor + end = self._cursor = start + size + return bytes(self.blob[start:end]) + + def readall(self) -> bytes: + self._raise_if_closed() + start = self._cursor + self._cursor = self._length + return bytes(self.blob[start:]) + + def readinto(self, b) -> int: + self._raise_if_closed() + start = self._cursor + to_write = min(len(b), self._length - start) + b[:] = self.blob[start:start + to_write] + self._cursor += to_write + return to_write + + def write(self, b): + self._raise_if_closed() + start = self._cursor + to_write = min(len(b), self._length - start) + self.blob[start:start + to_write] = b[:to_write] + self._cursor += to_write + return to_write + + def truncate(self, size: Optional[int] = None) -> int: + if size is None: + size = self._cursor + elif size < 0: + raise ValueError("size") + + if size > self._length: + self.blob.extend(b'\0' * (size - self._length)) + elif size < self._length: + self.blob[size:] = b'' + return size + + def seek(self, offset: int, whence: SeekWhence = SEEK_SET) -> int: + if whence == SEEK_SET: + self._cursor = offset + elif whence == SEEK_CUR: + self._cursor += offset + elif whence == SEEK_END: + self._cursor = self._length + offset + else: + raise ValueError("whence") + if self._cursor > self._length: + self._cursor = self._length + return self._cursor + + def tell(self) -> int: + return self._cursor + + def close(self): + super().close() + try: + del self.blob + except AttributeError: + pass + + +class BytesBuffer(BufferedIOBase): + """ + Access bytes from some location in a safe and sane manner. + """ + raw: RawIOBase + _base: int + _end: int + _cursor: int + _unbounded: bool + + # Proxied methods + flush: Callable[[], None] + readable: Callable[[], bool] + writable: Callable[[], bool] + seekable: Callable[[], bool] + isatty: Callable[[], bool] + truncate: Callable[[Optional[int]], int] + + def __init__(self, blob: Union[bytes, RawIOBase] = b'', *, base: int = 0, size: int = -1, bits: int = -1): + if base < 0: + raise ValueError("base") + + self._base = base + self._cursor = 0 + + if size >= 0: + if bits >= 0: + bits += size * 8 + else: + bits = size * 8 + + if bits < 0: + self._unbounded = True + if isinstance(blob, bytes): + self.raw = RawBytesIO(blob) + self._end = len(blob) + else: + self.raw = blob + cur = blob.tell() + blob.seek(0, SEEK_END) + self._end = blob.tell() + blob.seek(cur) + else: + # TODO: support any granularity + if bits % 8: + raise ValueError(f"{self.__class__.__name__} currently only supports byte-bounded streams") + + size = bits // 8 + + if isinstance(blob, bytes): + self.raw = RawBytesIO(blob + b"\0" * (size - (len(blob) - base))) + self._end = base + size + else: + self.raw = blob + self._end = base + size + + self.flush = self.raw.flush + self.readable = self.raw.readable + self.writable = self.raw.writable + self.seekable = self.raw.seekable + self.isatty = self.raw.isatty + self.truncate = self.raw.truncate + + def read(self, size=-1) -> bytes: + if size == -1: + return self.readall() + elif size < 0: + raise ValueError("size") + start = self._base + self._cursor + self.raw.seek(start) + if start + size > self._end: + size = self._end - start + if size < 0: + return b"" + ret = self.raw.read(size) + self._cursor = self.raw.tell() - self._base + return ret + + def readall(self) -> bytes: + start = self._base + self._cursor + self.raw.seek(start) + ret = self.raw.read(self._end - start) + self._cursor = self._end + return ret + + def readinto(self, b) -> int: + start = self._base + self._cursor + self.raw.seek(start) + if self._end - start > len(b): + ret = self._end - start + if ret: + b[:ret] = self.raw.read(ret) + else: + ret = self.raw.readinto(b) + self._cursor += ret + return ret + + def write(self, b: bytes) -> int: + start = self._base + self._cursor + self.raw.seek(start) + if start + len(b) > self._end: + b = b[:self._end - start] + ret: int = self.raw.write(b) + self._cursor += ret + return ret + + def seek(self, offset: int, whence: SeekWhence = SEEK_SET): + if whence == SEEK_SET: + self._cursor = offset + elif whence == SEEK_CUR: + self._cursor += offset + elif whence == SEEK_END: + self._cursor = self._end - self._base + offset + else: + raise ValueError("whence") + if self._base + self._cursor > self._end: + self._cursor = self._end - self._base + return self._cursor + + def tell(self) -> int: + return self._cursor diff --git a/test/core/test_dynamic.py b/test/core/test_dynamic.py index 9bc84c4..c36996e 100644 --- a/test/core/test_dynamic.py +++ b/test/core/test_dynamic.py @@ -3,17 +3,21 @@ from imperial import exceptions from imperial.core import base, dynamic, key + class Int(base.ImperialType): @classmethod def normalize(cls, value): + if isinstance(value, Int): + return value._data return int(value) def get_basic(self): return self._data - + def set_basic(self, value): self._data = self.normalize(value) + class Adder(dynamic.Dynamic): @classmethod def _register(cls): @@ -43,7 +47,7 @@ class Data(key.Key): @key.calculate def from_b(self, a, b): return a.get() + b.get() - + def get_basic(self): return self.get("data") diff --git a/test/core/test_number.py b/test/core/test_number.py new file mode 100644 index 0000000..4892067 --- /dev/null +++ b/test/core/test_number.py @@ -0,0 +1,20 @@ +import unittest + +from imperial.core import number +from imperial.exceptions import ImperialValueError + +class TestNumber(unittest.TestCase): + def test_create_primitive(self): + n = number.Number(100) + self.assertIsNotNone(n) + + def test_get_basic(self): + n = number.Number(100) + self.assertEqual(n.get(), 100) + self.assertEqual(n.number(), 100) + + with self.assertRaises(ImperialValueError): + n.string() + + with self.assertRaises(ImperialValueError): + n.list() diff --git a/test/core/test_serializable.py b/test/core/test_serializable.py new file mode 100644 index 0000000..cde2513 --- /dev/null +++ b/test/core/test_serializable.py @@ -0,0 +1,127 @@ +import unittest + +from imperial import exceptions +from imperial.core import key, serializable +from imperial.util import BytesBuffer, RawBytesIO +from imperial.magic import BASIC + + +class PosInt(serializable.Serializable): + has_basic = True + + @classmethod + def _register(cls): + @cls.register + class SizeKey(key.Key): + type = Size + default = 4 + keyname = "size" + + @classmethod + def normalize(cls, value): + if isinstance(value, PosInt): + return value._data + + i = int(value) + if i < 0: + raise exceptions.ImperialTypeError(value, expects=cls) + return i + + number = serializable.Serializable.get + + def get_basic(self): + return self._data + + def refresh_basic(self): + return self._data + + def set_basic(self, value): + self._data = self.normalize(value) + + @serializable.serialize + def serialize(self, blob: BytesBuffer): + blob.write(self._data.to_bytes(self.get("size"), "little")) + + @serializable.unserialize + def unserialize(self, blob: BytesBuffer) -> int: + return int.from_bytes(blob.read(self.get("size")), "little") + + +class Size(PosInt): + has_basic = True + + @classmethod + def _register(cls): + @cls.register + class Bits(key.Key): + type = PosInt + keyname = "bits" + + @key.calculate(BASIC) + def from_basic(self, basic: int): + return basic * 8 + + +class Pair(serializable.Serializable): + @classmethod + def _register(cls): + @cls.register + class Left(key.Key): + type = PosInt + keyname = "left" + + @cls.register + class Right(key.Key): + type = PosInt + keyname = "right" + + @serializable.serialize + def serialize(self, blob: BytesBuffer): + blob.write(self.number("left").to_bytes(2, "little")) + blob.write(self.number("right").to_bytes(2, "little")) + + @serializable.unserialize_yield + def unserialize(self, blob: BytesBuffer) -> int: + left = int.from_bytes(blob.read(2), "little") + yield "left", {"": left, "size": 2} + + right = int.from_bytes(blob.read(2), "little") + yield "right", {"": right, "size": 2} + + +class TestSerializable(unittest.TestCase): + def test_serialize_bytes(self): + one = PosInt(1) + self.assertEqual(one.serialize(), b"\x01\x00\x00\x00") + + def test_serialize_stream(self): + b = RawBytesIO(b"\x01\x02\x03\x04\x05\x06\x07\x08") + bb = BytesBuffer(b, base=1, size=4) + ten = PosInt(10) + self.assertIsNone(ten.serialize(bb)) + b.seek(0) + self.assertEqual(b.readall(), b"\x01\x0a\x00\x00\x00\x06\x07\x08") + + def test_unserialize_bytes(self): + one = PosInt() + one.unserialize(b"\x01\x00\x00\x00") + self.assertEqual(one.get(), 1) + + def test_unserialize_stream(self): + bb = BytesBuffer(b"\x01\x02\x03\x04\x05\x06\x07\x08", base=1, size=4) + num = PosInt() + num.unserialize(bb) + self.assertEqual(num.get(), 0x05040302) + + def test_unserialize_yield_bytes(self): + pair = Pair() + pair.unserialize(b"\x01\x00\x02\x00", {"left"}) + self.assertEqual(pair.get("left"), 1) + self.assertFalse("right" in dict(pair.keys)) + + pair.unserialize() + self.assertEqual(pair.get("right"), 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/core/test_value.py b/test/core/test_value.py new file mode 100644 index 0000000..72c8ad0 --- /dev/null +++ b/test/core/test_value.py @@ -0,0 +1,77 @@ +import unittest +from typing import cast + +from imperial.core import value +from imperial.core.base import ImperialType, PythonValue +from imperial.exceptions import ImperialKeyError, ImperialTypeError + + +class Int(value.Value): + def normalize(self, value: PythonValue) -> PythonValue: + if isinstance(value, Int): + return value._data + if isinstance(value, int): + return value + raise ImperialTypeError("value must be an int") + + @value.number + def number(self, data: PythonValue) -> int: + return cast(int, data) + + +class Str(value.Value): + @classmethod + def _register(cls): + @cls.register_converter(target=Int) + def convert_to_int(data: Str) -> Int: + return Int(int(data.string())) + + @cls.register_converter(source=Int) + def convert_from_int(data: Int) -> Str: + return Str(str(data.number())) + + def normalize(self, value: PythonValue) -> PythonValue: + if isinstance(value, Str): + return value._data + if isinstance(value, str): + return value + raise ImperialTypeError("value must be a str") + + @value.string + def string(self, data: PythonValue) -> str: + return cast(str, data) + + +class TestValuePrimitives(unittest.TestCase): + def test_create_with_primitive(self): + i = Int(1) + self.assertIs(i, i.resolve("data")) + self.assertEqual(1, i.get()) + self.assertEqual(1, i.get("data")) + + def test_set_value_with_primitive(self): + i = Int() + i.set(1) + self.assertIs(i, i.resolve("data")) + self.assertEqual(1, i.get()) + self.assertEqual(1, i.get("data")) + + def test_get_primitive_as_number(self): + i = Int(1) + self.assertEqual(i.number(), 1) + self.assertEqual(i.number("data"), 1) + + +class TestValueProxies(unittest.TestCase): + def test_create_with_redundant_proxy(self): + i = Int(Int(1)) + self.assertEqual(i.number(), 1) + self.assertEqual(i.number("data"), 1) + + # TODO: need reference groups to fix this + def test_set_with_redundant_proxy(self): + i = Int(Int()) + i.set(1) + self.assertEqual(i.number(), 1) + self.assertIsNot(i, i.resolve("data")) + self.assertEqual(i.number("data"), 1) diff --git a/test/test_util.py b/test/test_util.py new file mode 100644 index 0000000..fcef383 --- /dev/null +++ b/test/test_util.py @@ -0,0 +1,101 @@ +import unittest + +from imperial import util + +class TestRawBytesIO(unittest.TestCase): + def test_create_bytes(self): + util.RawBytesIO(b"123abc") + + def test_read_bytes(self): + b = util.RawBytesIO(b"123abc") + self.assertEqual(b.read(3), b"123") + self.assertEqual(b.read(3), b"abc") + + def test_readall_bytes(self): + b = util.RawBytesIO(b"123abc") + self.assertEqual(b.readall(), b"123abc") + + def test_seek_tell(self): + b = util.RawBytesIO(b"123abc") + b.seek(3) + self.assertEqual(b.read(3), b"abc") + + b.seek(1) + b.seek(2, util.SEEK_CUR) + self.assertEqual(b.tell(), 3) + + b.seek(-1, util.SEEK_CUR) + self.assertEqual(b.read(1), b"3") + + b.seek(-1, util.SEEK_END) + self.assertEqual(b.read(1), b"c") + + def test_write_bytes(self): + b = util.RawBytesIO(b"123abc") + b.write(b"pp") + b.seek(0) + self.assertEqual(b.read(), b"pp3abc") + + def test_close(self): + b = util.RawBytesIO(b"123abc") + b.close() + self.assertTrue(b.closed) + + with self.assertRaises(ValueError): + b.read(1) + + +class TestBytesBuffer(unittest.TestCase): + def test_create_bytes(self): + util.BytesBuffer(b"abc123") + + def test_create_window(self): + b = util.RawBytesIO(b"123abc456") + return util.BytesBuffer(b, base=3, size=3) + + def test_read_bytes(self): + bb = util.BytesBuffer(b"abc123") + self.assertEqual(bb.read(3), b"abc") + + def test_read_from_window(self): + b = util.RawBytesIO(b"123abc456") + w = util.BytesBuffer(b, base=3, size=3) + self.assertEqual(w.read(3), b"abc") + + def test_read_to_end_window(self): + b = util.RawBytesIO(b"123abc456") + w = util.BytesBuffer(b, base=3, size=3) + self.assertEqual(w.read(), b"abc") + self.assertEqual(w.read(1), b"") + + def test_readall_bytes(self): + bb = util.BytesBuffer(b"abc123") + self.assertEqual(bb.readall(), b"abc123") + + def test_readall_window(self): + b = util.RawBytesIO(b"123abc456") + w = util.BytesBuffer(b, base=3, size=3) + self.assertEqual(w.readall(), b"abc") + + def test_seek_tell_window(self): + b = util.RawBytesIO(b"123abc456") + w = util.BytesBuffer(b, base=3, size=3) + self.assertEqual(w.tell(), 0) + w.seek(0) + self.assertEqual(w.tell(), 0) + self.assertEqual(w.read(1), b"a") + self.assertEqual(w.tell(), 1) + + self.assertEqual(w.seek(-1, util.SEEK_CUR), 0) + self.assertEqual(w.tell(), 0) + + self.assertEqual(w.seek(0, util.SEEK_END), 3) + self.assertEqual(w.tell(), 3) + + def test_multiple_windows(self): + b = util.RawBytesIO(b"123abc456") + w1 = util.BytesBuffer(b, base=0, size=3) + w2 = util.BytesBuffer(b, base=3, size=3) + + self.assertEqual(w1.read(1), b"1") + self.assertEqual(w2.read(1), b"a")