Coverage for src/chebpy/decorators.py: 100%
78 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:43 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:43 +0000
1"""Decorator functions for the ChebPy package.
3This module provides various decorators used throughout the ChebPy package to
4implement common functionality such as caching, handling empty objects,
5pre- and post-processing of function inputs/outputs, and type conversion.
6These decorators help reduce code duplication and ensure consistent behavior
7across the package.
8"""
10from collections.abc import Callable
11from functools import wraps
12from typing import Any
14import numpy as np
17def cache(f: Callable[..., Any]) -> Callable[..., Any]:
18 """Object method output caching mechanism.
20 This decorator caches the output of zero-argument methods to speed up repeated
21 execution of relatively expensive operations such as .roots(). Cached computations
22 are stored in a dictionary called _cache which is bound to self using keys
23 corresponding to the method name.
25 Args:
26 f (callable): The method to be cached. Must be a zero-argument method.
28 Returns:
29 callable: A wrapped version of the method that implements caching.
31 Note:
32 Can be used in principle on arbitrary objects.
33 """
35 # TODO: look into replacing this with one of the functools cache decorators
36 @wraps(f)
37 def wrapper(self: Any) -> Any:
38 """Return the cached method result, computing and storing it on first call."""
39 try:
40 # f has been executed previously
41 out = self._cache[f.__name__] # ty: ignore[unresolved-attribute]
42 except AttributeError:
43 # f has not been executed previously and self._cache does not exist
44 self._cache = {}
45 out = self._cache[f.__name__] = f(self) # ty: ignore[unresolved-attribute]
46 except KeyError:
47 # f has not been executed previously, but self._cache exists
48 out = self._cache[f.__name__] = f(self) # ty: ignore[unresolved-attribute]
49 return out
51 return wrapper
54def self_empty(resultif: Any = None) -> Callable[..., Any]:
55 """Factory method to produce a decorator for handling empty objects.
57 This factory creates a decorator that checks whether the object whose method
58 is being wrapped is empty. If the object is empty, it returns either the supplied
59 resultif value or a copy of the object. Otherwise, it executes the wrapped method.
61 Args:
62 resultif: Value to return when the object is empty. If None, returns a copy
63 of the object instead.
65 Returns:
66 callable: A decorator function that implements the empty-checking logic.
68 Note:
69 This decorator is primarily used in chebtech.py.
70 """
72 def decorator(f: Callable[..., Any]) -> Callable[..., Any]:
73 """Wrap *f* with the empty-object short-circuit logic."""
75 @wraps(f)
76 def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any:
77 """Return the empty-case result if *self* is empty, else call *f*."""
78 if self.isempty:
79 if resultif is not None:
80 return resultif
81 else:
82 return self.copy()
83 else:
84 return f(self, *args, **kwargs)
86 return wrapper
88 return decorator
91def preandpostprocess(f: Callable[..., Any]) -> Callable[..., Any]:
92 """Decorator for pre- and post-processing tasks common to bary and clenshaw.
94 This decorator handles several edge cases for functions like bary and clenshaw:
95 - Empty arrays in input arguments
96 - Constant functions
97 - NaN values in coefficients
98 - Scalar vs. array inputs
100 Args:
101 f (callable): The function to be wrapped.
103 Returns:
104 callable: A wrapped version of the function with pre- and post-processing.
105 """
107 @wraps(f)
108 def thewrapper(*args: Any, **kwargs: Any) -> Any:
109 """Handle empty/constant/NaN/scalar edge cases around *f*."""
110 xx, akfk = args[:2]
111 # are any of the first two arguments empty arrays?
112 if (np.asarray(xx).size == 0) | (np.asarray(akfk).size == 0):
113 return np.array([])
114 # is the function constant?
115 elif akfk.size == 1:
116 if np.isscalar(xx):
117 return akfk[0]
118 else:
119 return akfk * np.ones(xx.size)
120 # are there any NaNs in the second argument?
121 elif np.any(np.isnan(akfk)):
122 return np.nan * np.ones(xx.size)
123 # convert first argument to an array if it is a scalar and then
124 # return the first (and only) element of the result if so
125 else:
126 args_list = list(args)
127 args_list[0] = np.array([xx]) if np.isscalar(xx) else args_list[0]
128 out = f(*args_list, **kwargs)
129 return out[0] if np.isscalar(xx) else out
131 return thewrapper
134def float_argument(f: Callable[..., Any]) -> Callable[..., Any]:
135 """Decorator to ensure consistent input/output types for Chebfun __call__ method.
137 This decorator ensures that when a Chebfun object is called with a float input,
138 it returns a float output, and when called with an array input, it returns an
139 array output. It handles various input formats including scalars, numpy arrays,
140 and nested arrays.
142 Args:
143 f (callable): The __call__ method to be wrapped.
145 Returns:
146 callable: A wrapped version of the method that ensures type consistency.
147 """
149 @wraps(f)
150 def thewrapper(self: Any, *args: Any, **kwargs: Any) -> Any:
151 """Coerce the first argument to an array and match scalar/array output to it."""
152 x = args[0]
153 xx = np.array([x]) if np.isscalar(x) else np.array(x)
154 # discern between the array(0.1) and array([0.1]) cases
155 if xx.size == 1 and xx.ndim == 0:
156 xx = np.array([xx])
157 args_list = list(args)
158 args_list[0] = xx
159 out = f(self, *args_list, **kwargs)
160 return out[0] if np.isscalar(x) else out
162 return thewrapper
165def cast_arg_to_chebfun(f: Callable[..., Any]) -> Callable[..., Any]:
166 """Decorator to cast the first argument to a chebfun object if needed.
168 This decorator attempts to convert the first argument to a chebfun object
169 if it is not already one. Currently, only numeric types can be cast to chebfun.
171 Args:
172 f (callable): The method to be wrapped.
174 Returns:
175 callable: A wrapped version of the method that ensures the first argument
176 is a chebfun object.
177 """
179 @wraps(f)
180 def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any:
181 """Cast the first argument to a chebfun (if needed) before calling *f*."""
182 other = args[0]
183 if not isinstance(other, self.__class__):
184 fun = self.initconst(args[0], self.support)
185 args_list = list(args)
186 args_list[0] = fun
187 return f(self, *args_list, **kwargs)
188 return f(self, *args, **kwargs)
190 return wrapper
193def cast_other(f: Callable[..., Any]) -> Callable[..., Any]:
194 """Decorator to cast the first argument to the same type as self.
196 This generic decorator is applied to binary operator methods to ensure that
197 the first positional argument (typically 'other') is cast to the same type
198 as the object on which the method is called.
200 Args:
201 f (callable): The binary operator method to be wrapped.
203 Returns:
204 callable: A wrapped version of the method that ensures type consistency
205 between self and the first argument.
206 """
208 @wraps(f)
209 def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any:
210 """Cast the first argument to ``type(self)`` (if needed) before calling *f*."""
211 cls = self.__class__
212 other = args[0]
213 if not isinstance(other, cls):
214 args_list = list(args)
215 args_list[0] = cls(other)
216 return f(self, *args_list, **kwargs)
217 return f(self, *args, **kwargs)
219 return wrapper