phaserEM 0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- phaser/__init__.py +0 -0
- phaser/__main__.py +5 -0
- phaser/engines/common/__init__.py +0 -0
- phaser/engines/common/noise_models.py +113 -0
- phaser/engines/common/output.py +189 -0
- phaser/engines/common/position_correction.py +62 -0
- phaser/engines/common/regularizers.py +403 -0
- phaser/engines/common/simulation.py +270 -0
- phaser/engines/conventional/__init__.py +0 -0
- phaser/engines/conventional/run.py +142 -0
- phaser/engines/conventional/solvers.py +476 -0
- phaser/engines/gradient/run.py +451 -0
- phaser/engines/gradient/solvers.py +139 -0
- phaser/execute.py +371 -0
- phaser/hooks/__init__.py +158 -0
- phaser/hooks/hook.py +159 -0
- phaser/hooks/io/empad.py +88 -0
- phaser/hooks/object.py +25 -0
- phaser/hooks/preprocessing.py +133 -0
- phaser/hooks/probe.py +24 -0
- phaser/hooks/regularization.py +97 -0
- phaser/hooks/scan.py +27 -0
- phaser/hooks/schedule.py +75 -0
- phaser/hooks/solver.py +169 -0
- phaser/io/__init__.py +0 -0
- phaser/io/empad.py +212 -0
- phaser/main.py +92 -0
- phaser/plan.py +184 -0
- phaser/py.typed +0 -0
- phaser/state.py +249 -0
- phaser/types.py +305 -0
- phaser/utils/__init__.py +0 -0
- phaser/utils/_cuda_kernels.py +213 -0
- phaser/utils/_jax_kernels.py +98 -0
- phaser/utils/analysis.py +263 -0
- phaser/utils/image.py +201 -0
- phaser/utils/io.py +402 -0
- phaser/utils/misc.py +295 -0
- phaser/utils/num.py +800 -0
- phaser/utils/object.py +578 -0
- phaser/utils/optics.py +377 -0
- phaser/utils/physics.py +88 -0
- phaser/utils/plotting.py +699 -0
- phaser/utils/scan.py +60 -0
- phaser/web/__init__.py +0 -0
- phaser/web/dist/03510a839ccb97b0da9f.module.wasm +0 -0
- phaser/web/dist/9573273f862f4f5d9644.module.wasm +0 -0
- phaser/web/dist/bundle-dashboard.js +712 -0
- phaser/web/dist/bundle-manager.js +210 -0
- phaser/web/dist/bundle-vendors-node_modules_wasm-array_wasm_array_js.js +106 -0
- phaser/web/dist/style.css +152 -0
- phaser/web/notebook.py +269 -0
- phaser/web/routes.py +211 -0
- phaser/web/server.py +540 -0
- phaser/web/slurm.py +180 -0
- phaser/web/templates/base.html +13 -0
- phaser/web/templates/dashboard.html +14 -0
- phaser/web/templates/manager.html +19 -0
- phaser/web/types.py +255 -0
- phaser/web/util.py +116 -0
- phaser/web/worker.py +200 -0
- phaserem-0.1.dist-info/METADATA +121 -0
- phaserem-0.1.dist-info/RECORD +67 -0
- phaserem-0.1.dist-info/WHEEL +5 -0
- phaserem-0.1.dist-info/entry_points.txt +2 -0
- phaserem-0.1.dist-info/licenses/LICENSE.txt +373 -0
- phaserem-0.1.dist-info/top_level.txt +1 -0
phaser/state.py
ADDED
|
@@ -0,0 +1,249 @@
|
|
|
1
|
+
import typing as t
|
|
2
|
+
|
|
3
|
+
import numpy
|
|
4
|
+
from numpy.typing import NDArray
|
|
5
|
+
from typing_extensions import Self
|
|
6
|
+
|
|
7
|
+
from phaser.utils.num import Sampling, to_numpy, get_array_module, Float
|
|
8
|
+
from phaser.utils.misc import jax_dataclass
|
|
9
|
+
from phaser.utils.object import ObjectSampling
|
|
10
|
+
|
|
11
|
+
if t.TYPE_CHECKING:
|
|
12
|
+
from phaser.utils.io import HdfLike
|
|
13
|
+
from phaser.utils.image import _BoundaryMode
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@jax_dataclass
|
|
17
|
+
class Patterns():
|
|
18
|
+
patterns: NDArray[numpy.floating]
|
|
19
|
+
"""Raw diffraction patterns, with 0-frequency sample in corner"""
|
|
20
|
+
pattern_mask: NDArray[numpy.floating]
|
|
21
|
+
"""Mask indicating which portions of the diffraction patterns contain data."""
|
|
22
|
+
|
|
23
|
+
def to_numpy(self) -> Self:
|
|
24
|
+
return self.__class__(
|
|
25
|
+
to_numpy(self.patterns), to_numpy(self.pattern_mask)
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@jax_dataclass
|
|
30
|
+
class IterState():
|
|
31
|
+
engine_num: int
|
|
32
|
+
"""Engine number. 1-indexed (0 means before any reconstruction)."""
|
|
33
|
+
engine_iter: int
|
|
34
|
+
"""Iteration number on this engine. 1-indexed (0 means before any iterations)."""
|
|
35
|
+
total_iter: int
|
|
36
|
+
"""Total iteration number. 1-indexed (0 means before any iterations)."""
|
|
37
|
+
|
|
38
|
+
def to_numpy(self) -> Self:
|
|
39
|
+
return self.__class__(
|
|
40
|
+
int(self.engine_num), int(self.engine_iter), int(self.total_iter)
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
def copy(self) -> Self:
|
|
44
|
+
import copy
|
|
45
|
+
return copy.deepcopy(self)
|
|
46
|
+
|
|
47
|
+
@staticmethod
|
|
48
|
+
def empty() -> 'IterState':
|
|
49
|
+
return IterState(0, 0, 0)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@jax_dataclass(static_fields=('sampling',))
|
|
53
|
+
class ProbeState():
|
|
54
|
+
sampling: Sampling
|
|
55
|
+
"""Probe coordinate system. See `Sampling` for more details."""
|
|
56
|
+
data: NDArray[numpy.complexfloating]
|
|
57
|
+
"""Probe wavefunction, in realspace. Shape (modes, y, x)"""
|
|
58
|
+
|
|
59
|
+
def resample(
|
|
60
|
+
self, new_samp: Sampling,
|
|
61
|
+
rotation: float = 0.0,
|
|
62
|
+
order: int = 1,
|
|
63
|
+
mode: '_BoundaryMode' = 'grid-constant',
|
|
64
|
+
) -> Self:
|
|
65
|
+
new_data = self.sampling.resample(
|
|
66
|
+
self.data, new_samp,
|
|
67
|
+
rotation=rotation,
|
|
68
|
+
order=order,
|
|
69
|
+
mode=mode,
|
|
70
|
+
)
|
|
71
|
+
return self.__class__(new_samp, new_data)
|
|
72
|
+
|
|
73
|
+
def to_xp(self, xp: t.Any) -> Self:
|
|
74
|
+
return self.__class__(
|
|
75
|
+
self.sampling, xp.array(self.data)
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
def to_numpy(self) -> Self:
|
|
79
|
+
return self.__class__(
|
|
80
|
+
self.sampling, to_numpy(self.data)
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
def copy(self) -> Self:
|
|
84
|
+
import copy
|
|
85
|
+
return copy.deepcopy(self)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
@jax_dataclass(static_fields=('sampling',))
|
|
89
|
+
class ObjectState():
|
|
90
|
+
sampling: ObjectSampling
|
|
91
|
+
"""Object coordinate system. See `ObjectSampling` for more details."""
|
|
92
|
+
data: NDArray[numpy.complexfloating]
|
|
93
|
+
"""Object wavefunction. Shape (z, y, x)"""
|
|
94
|
+
thicknesses: NDArray[numpy.floating]
|
|
95
|
+
"""
|
|
96
|
+
Slice thicknesses (in length units).
|
|
97
|
+
Length < 2 for single slice, equal to the number of slices otherwise.
|
|
98
|
+
"""
|
|
99
|
+
|
|
100
|
+
def to_xp(self, xp: t.Any) -> Self:
|
|
101
|
+
return self.__class__(
|
|
102
|
+
self.sampling, xp.array(self.data), xp.array(self.thicknesses)
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
def to_numpy(self) -> Self:
|
|
106
|
+
return self.__class__(
|
|
107
|
+
self.sampling, to_numpy(self.data), to_numpy(self.thicknesses)
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
def zs(self) -> NDArray[numpy.floating]:
|
|
111
|
+
xp = get_array_module(self.thicknesses)
|
|
112
|
+
if len(self.thicknesses) < 2:
|
|
113
|
+
return xp.array([0.], dtype=self.thicknesses.dtype)
|
|
114
|
+
return xp.cumsum(self.thicknesses) - self.thicknesses
|
|
115
|
+
|
|
116
|
+
def copy(self) -> Self:
|
|
117
|
+
import copy
|
|
118
|
+
return copy.deepcopy(self)
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
@jax_dataclass
|
|
122
|
+
class ProgressState:
|
|
123
|
+
iters: NDArray[numpy.integer]
|
|
124
|
+
"""Iterations error measurements were taken at."""
|
|
125
|
+
detector_errors: NDArray[numpy.floating]
|
|
126
|
+
"""Detector error measurements at those iterations"""
|
|
127
|
+
|
|
128
|
+
def to_numpy(self) -> Self:
|
|
129
|
+
return self.__class__(
|
|
130
|
+
to_numpy(self.iters), to_numpy(self.detector_errors)
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
def copy(self) -> Self:
|
|
134
|
+
import copy
|
|
135
|
+
return copy.deepcopy(self)
|
|
136
|
+
|
|
137
|
+
@staticmethod
|
|
138
|
+
def empty() -> 'ProgressState':
|
|
139
|
+
return ProgressState(
|
|
140
|
+
numpy.array([], dtype=numpy.uint64),
|
|
141
|
+
numpy.array([], dtype=numpy.float64),
|
|
142
|
+
)
|
|
143
|
+
|
|
144
|
+
# TODO: this is a hack to prevent JIT recompilation.
|
|
145
|
+
def __hash__(self) -> int:
|
|
146
|
+
return id(self)
|
|
147
|
+
|
|
148
|
+
def __eq__(self, other: t.Any) -> bool:
|
|
149
|
+
if type(self) is not type(other):
|
|
150
|
+
return False
|
|
151
|
+
xp = get_array_module(self.iters, other.iters)
|
|
152
|
+
return (
|
|
153
|
+
xp.array_equal(self.iters, other.iters) and
|
|
154
|
+
xp.array_equal(self.detector_errors, other.detector_errors)
|
|
155
|
+
)
|
|
156
|
+
|
|
157
|
+
@jax_dataclass(kw_only=True, static_fields=('progress',))
|
|
158
|
+
class ReconsState:
|
|
159
|
+
iter: IterState
|
|
160
|
+
wavelength: Float
|
|
161
|
+
|
|
162
|
+
probe: ProbeState
|
|
163
|
+
object: ObjectState
|
|
164
|
+
scan: NDArray[numpy.floating]
|
|
165
|
+
"""Scan coordinates (y, x), in length units. Shape (..., 2)"""
|
|
166
|
+
progress: ProgressState
|
|
167
|
+
|
|
168
|
+
def to_xp(self, xp: t.Any) -> Self:
|
|
169
|
+
return self.__class__(
|
|
170
|
+
iter=self.iter,
|
|
171
|
+
probe=self.probe.to_xp(xp),
|
|
172
|
+
object=self.object.to_xp(xp),
|
|
173
|
+
scan=xp.array(self.scan),
|
|
174
|
+
progress=self.progress,
|
|
175
|
+
wavelength=self.wavelength,
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
def to_numpy(self) -> Self:
|
|
179
|
+
return self.__class__(
|
|
180
|
+
iter=self.iter.to_numpy(),
|
|
181
|
+
probe=self.probe.to_numpy(),
|
|
182
|
+
object=self.object.to_numpy(),
|
|
183
|
+
scan=to_numpy(self.scan),
|
|
184
|
+
progress=self.progress.to_numpy(),
|
|
185
|
+
wavelength=float(self.wavelength),
|
|
186
|
+
)
|
|
187
|
+
|
|
188
|
+
def copy(self) -> Self:
|
|
189
|
+
import copy
|
|
190
|
+
return copy.deepcopy(self)
|
|
191
|
+
|
|
192
|
+
def write_hdf5(self, file: 'HdfLike'):
|
|
193
|
+
from phaser.utils.io import hdf5_write_state
|
|
194
|
+
hdf5_write_state(self, file)
|
|
195
|
+
|
|
196
|
+
@staticmethod
|
|
197
|
+
def read_hdf5(file: 'HdfLike') -> 'ReconsState':
|
|
198
|
+
from phaser.utils.io import hdf5_read_state
|
|
199
|
+
return hdf5_read_state(file).to_complete()
|
|
200
|
+
|
|
201
|
+
|
|
202
|
+
@jax_dataclass(kw_only=True, static_fields=('progress',))
|
|
203
|
+
class PartialReconsState:
|
|
204
|
+
iter: t.Optional[IterState] = None
|
|
205
|
+
wavelength: t.Optional[Float] = None
|
|
206
|
+
|
|
207
|
+
probe: t.Optional[ProbeState] = None
|
|
208
|
+
object: t.Optional[ObjectState] = None
|
|
209
|
+
scan: t.Optional[NDArray[numpy.floating]] = None
|
|
210
|
+
"""Scan coordinates (y, x), in length units. Shape (..., 2)"""
|
|
211
|
+
progress: t.Optional[ProgressState] = None
|
|
212
|
+
|
|
213
|
+
def to_numpy(self) -> Self:
|
|
214
|
+
return self.__class__(
|
|
215
|
+
iter=self.iter.to_numpy() if self.iter is not None else None,
|
|
216
|
+
probe=self.probe.to_numpy() if self.probe is not None else None,
|
|
217
|
+
object=self.object.to_numpy() if self.object is not None else None,
|
|
218
|
+
scan=to_numpy(self.scan) if self.scan is not None else None,
|
|
219
|
+
progress=self.progress.to_numpy() if self.progress is not None else None,
|
|
220
|
+
wavelength=float(self.wavelength) if self.wavelength is not None else None,
|
|
221
|
+
)
|
|
222
|
+
|
|
223
|
+
def to_complete(self) -> ReconsState:
|
|
224
|
+
missing = tuple(filter(lambda k: getattr(self, k) is None, ('probe', 'object', 'scan', 'wavelength')))
|
|
225
|
+
if len(missing):
|
|
226
|
+
raise ValueError(f"ReconsState missing {', '.join(map(repr, missing))}")
|
|
227
|
+
|
|
228
|
+
progress = self.progress if self.progress is not None else ProgressState.empty()
|
|
229
|
+
iter = self.iter if self.iter is not None else IterState.empty()
|
|
230
|
+
|
|
231
|
+
return ReconsState(
|
|
232
|
+
wavelength=t.cast(Float, self.wavelength),
|
|
233
|
+
probe=t.cast(ProbeState, self.probe),
|
|
234
|
+
object=t.cast(ObjectState, self.object),
|
|
235
|
+
scan=t.cast(NDArray[numpy.floating], self.scan),
|
|
236
|
+
progress=progress, iter=iter,
|
|
237
|
+
)
|
|
238
|
+
|
|
239
|
+
def write_hdf5(self, file: 'HdfLike'):
|
|
240
|
+
from phaser.utils.io import hdf5_write_state
|
|
241
|
+
hdf5_write_state(self, file)
|
|
242
|
+
|
|
243
|
+
@staticmethod
|
|
244
|
+
def read_hdf5(file: 'HdfLike') -> 'PartialReconsState':
|
|
245
|
+
from phaser.utils.io import hdf5_read_state
|
|
246
|
+
return hdf5_read_state(file)
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
StateObserver: t.TypeAlias = t.Callable[[t.Union[ReconsState, PartialReconsState]], t.Any]
|
phaser/types.py
ADDED
|
@@ -0,0 +1,305 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from functools import lru_cache
|
|
3
|
+
import typing as t
|
|
4
|
+
|
|
5
|
+
import numpy
|
|
6
|
+
import pane
|
|
7
|
+
from pane.converters import Converter, make_converter, ConverterHandlers, ErrorNode
|
|
8
|
+
from pane.annotations import ConvertAnnotation
|
|
9
|
+
from pane.errors import ParseInterrupt, WrongTypeError
|
|
10
|
+
from pane.util import pluralize, list_phrase
|
|
11
|
+
from typing_extensions import Self
|
|
12
|
+
|
|
13
|
+
if t.TYPE_CHECKING:
|
|
14
|
+
from phaser.hooks.schedule import FlagArgs, FlagLike, ScheduleLike
|
|
15
|
+
|
|
16
|
+
T = t.TypeVar('T')
|
|
17
|
+
|
|
18
|
+
@t.overload
|
|
19
|
+
def cast_length(val: t.Iterable[T], n: t.Literal[5]) -> t.Tuple[T, T, T, T, T]:
|
|
20
|
+
...
|
|
21
|
+
|
|
22
|
+
@t.overload
|
|
23
|
+
def cast_length(val: t.Iterable[T], n: t.Literal[4]) -> t.Tuple[T, T, T, T]:
|
|
24
|
+
...
|
|
25
|
+
|
|
26
|
+
@t.overload
|
|
27
|
+
def cast_length(val: t.Iterable[T], n: t.Literal[3]) -> t.Tuple[T, T, T]:
|
|
28
|
+
...
|
|
29
|
+
|
|
30
|
+
@t.overload
|
|
31
|
+
def cast_length(val: t.Iterable[T], n: t.Literal[2]) -> t.Tuple[T, T]:
|
|
32
|
+
...
|
|
33
|
+
|
|
34
|
+
@t.overload
|
|
35
|
+
def cast_length(val: t.Iterable[T], n: t.Literal[1]) -> t.Tuple[T]:
|
|
36
|
+
...
|
|
37
|
+
|
|
38
|
+
def cast_length(val: t.Iterable[T], n: int) -> t.Tuple[T, ...]:
|
|
39
|
+
return tuple(val)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class _EmptyDictAnnotation(ConvertAnnotation, Converter[t.Dict[t.NoReturn, t.NoReturn]]):
|
|
43
|
+
def _converter(self, inner_type: t.Any, *, handlers: ConverterHandlers):
|
|
44
|
+
return self
|
|
45
|
+
|
|
46
|
+
def __eq__(self, other):
|
|
47
|
+
return self.__class__ is other.__class__
|
|
48
|
+
|
|
49
|
+
def __hash__(self) -> int:
|
|
50
|
+
return hash(self.__class__.__name__)
|
|
51
|
+
|
|
52
|
+
def expected(self, plural: bool = False) -> str:
|
|
53
|
+
return pluralize("empty dict", plural, article='an')
|
|
54
|
+
|
|
55
|
+
def try_convert(self, val: t.Any) -> t.Dict[t.NoReturn, t.NoReturn]:
|
|
56
|
+
if isinstance(val, (dict, t.Mapping)) and len(val) == 0:
|
|
57
|
+
return {}
|
|
58
|
+
raise ParseInterrupt()
|
|
59
|
+
|
|
60
|
+
def collect_errors(self, val: t.Any) -> t.Optional[WrongTypeError]:
|
|
61
|
+
if isinstance(val, dict) and len(val) == 0:
|
|
62
|
+
return None
|
|
63
|
+
return WrongTypeError(self.expected(), val)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class _ReconsVarsAnnotation(ConvertAnnotation):
|
|
67
|
+
def _converter(self, inner_type: t.Any, *, handlers: ConverterHandlers):
|
|
68
|
+
return _ReconsVarsConverter(inner_type, handlers)
|
|
69
|
+
|
|
70
|
+
def __eq__(self, other):
|
|
71
|
+
return self.__class__ is other.__class__
|
|
72
|
+
|
|
73
|
+
def __hash__(self) -> int:
|
|
74
|
+
return hash(self.__class__.__name__)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
BackendName: t.TypeAlias = t.Literal['cuda', 'cupy', 'jax', 'cpu', 'numpy']
|
|
78
|
+
ReconsVar: t.TypeAlias = t.Literal['object', 'probe', 'positions']
|
|
79
|
+
|
|
80
|
+
ReconsVars: t.TypeAlias = t.Annotated[t.FrozenSet[ReconsVar], _ReconsVarsAnnotation()]
|
|
81
|
+
EmptyDict: t.TypeAlias = t.Annotated[t.Dict[t.NoReturn, t.NoReturn], _EmptyDictAnnotation()]
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
class Cancelled(BaseException):
|
|
85
|
+
...
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class Dataclass(pane.PaneBase, kw_only=True, allow_extra=False):
|
|
89
|
+
...
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
class SliceList(Dataclass):
|
|
93
|
+
thicknesses: t.List[float]
|
|
94
|
+
|
|
95
|
+
class SliceStep(Dataclass):
|
|
96
|
+
n: int
|
|
97
|
+
slice_thickness: float
|
|
98
|
+
|
|
99
|
+
@property
|
|
100
|
+
def zs(self) -> t.List[float]:
|
|
101
|
+
return [float(z) for z in numpy.arange(self.n) * self.slice_thickness]
|
|
102
|
+
|
|
103
|
+
@property
|
|
104
|
+
def thicknesses(self) -> t.List[float]:
|
|
105
|
+
return [self.slice_thickness] * self.n
|
|
106
|
+
|
|
107
|
+
class SliceTotal(Dataclass):
|
|
108
|
+
n: int
|
|
109
|
+
total_thickness: float
|
|
110
|
+
|
|
111
|
+
@property
|
|
112
|
+
def zs(self) -> t.List[float]:
|
|
113
|
+
return [float(z) for z in numpy.arange(self.n) * self.total_thickness / self.n]
|
|
114
|
+
|
|
115
|
+
@property
|
|
116
|
+
def thicknesses(self) -> t.List[float]:
|
|
117
|
+
return [self.total_thickness / self.n] * self.n
|
|
118
|
+
|
|
119
|
+
Slices: t.TypeAlias = t.Union[SliceList, SliceStep, SliceTotal]
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
class Flag(Dataclass):
|
|
123
|
+
after: int = 0
|
|
124
|
+
every: int = 1
|
|
125
|
+
before: t.Optional[int] = None
|
|
126
|
+
|
|
127
|
+
def any_true(self, niter: int) -> bool:
|
|
128
|
+
return (
|
|
129
|
+
self.after < niter
|
|
130
|
+
and (self.before is None or self.before > 0)
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
def __call__(self, args: 'FlagArgs') -> bool:
|
|
134
|
+
i = args['state'].iter.engine_iter
|
|
135
|
+
return (
|
|
136
|
+
(self.before is None or i < self.before)
|
|
137
|
+
and i > self.after
|
|
138
|
+
and (i - self.after) % self.every == 0
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
def resolve(self) -> Self:
|
|
142
|
+
return self
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
class _ConstFlag:
|
|
146
|
+
def __init__(self, val: bool):
|
|
147
|
+
self.val = val
|
|
148
|
+
|
|
149
|
+
def __call__(self, args: 'FlagArgs') -> bool:
|
|
150
|
+
return self.val
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
@lru_cache
|
|
154
|
+
def process_flag(flag: 'FlagLike') -> t.Callable[['FlagArgs'], bool]:
|
|
155
|
+
if isinstance(flag, bool):
|
|
156
|
+
return _ConstFlag(flag)
|
|
157
|
+
return flag
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
@lru_cache
|
|
161
|
+
def process_schedule(schedule: 'ScheduleLike') -> t.Callable[['FlagArgs'], float]:
|
|
162
|
+
if isinstance(schedule, (int, float)):
|
|
163
|
+
return lambda _: schedule
|
|
164
|
+
return schedule
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def flag_any_true(flag: t.Callable[['FlagArgs'], bool], niter: int) -> bool:
|
|
168
|
+
if isinstance(flag, Flag):
|
|
169
|
+
return flag.any_true(niter)
|
|
170
|
+
elif isinstance(flag, _ConstFlag):
|
|
171
|
+
return flag.val
|
|
172
|
+
# assume flag will return true
|
|
173
|
+
return True
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
@dataclass(init=False, frozen=True)
|
|
177
|
+
class IsVersion(ConvertAnnotation):
|
|
178
|
+
min: t.Optional[t.Tuple[int, ...]] = None
|
|
179
|
+
max: t.Optional[t.Tuple[int, ...]] = None
|
|
180
|
+
|
|
181
|
+
def __init__(self, *,
|
|
182
|
+
min: t.Union[str, t.Tuple[int, ...], None] = None,
|
|
183
|
+
max: t.Union[str, t.Tuple[int, ...], None] = None,
|
|
184
|
+
exactly: t.Union[str, t.Tuple[int, ...], None] = None):
|
|
185
|
+
if exactly is not None:
|
|
186
|
+
if min is not None or max is not None:
|
|
187
|
+
raise TypeError("'exactly' cannot be specified with 'min' or 'max'")
|
|
188
|
+
min = max = _VersionConverter.parse_version(exactly)
|
|
189
|
+
if min is not None:
|
|
190
|
+
min = _VersionConverter.parse_version(min)
|
|
191
|
+
if max is not None:
|
|
192
|
+
max = _VersionConverter.parse_version(max)
|
|
193
|
+
object.__setattr__(self, 'min', min)
|
|
194
|
+
object.__setattr__(self, 'max', max)
|
|
195
|
+
|
|
196
|
+
def _converter(self, inner_type: t.Any, *,
|
|
197
|
+
handlers: ConverterHandlers):
|
|
198
|
+
return _VersionConverter(inner_type, handlers, min=self.min, max=self.max)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
Version: t.TypeAlias = t.Annotated[str, IsVersion()]
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
class _VersionConverter(Converter[t.Any]):
|
|
205
|
+
def __init__(self, inner_type: t.Any, handlers: ConverterHandlers,
|
|
206
|
+
min: t.Optional[t.Tuple[int, ...]] = None,
|
|
207
|
+
max: t.Optional[t.Tuple[int, ...]] = None):
|
|
208
|
+
self.inner = make_converter(inner_type, handlers)
|
|
209
|
+
self.min = min
|
|
210
|
+
self.max = max
|
|
211
|
+
|
|
212
|
+
def expected(self, plural: bool = False) -> str:
|
|
213
|
+
if self.min is not None and self.max is not None:
|
|
214
|
+
min_s = '.'.join(map(str, self.min))
|
|
215
|
+
max_s = '.'.join(map(str, self.max))
|
|
216
|
+
if self.min == self.max:
|
|
217
|
+
return f"{pluralize('version', plural)} {min_s}"
|
|
218
|
+
return f"{pluralize('version', plural)} between {min_s} and {max_s}"
|
|
219
|
+
elif self.min is not None:
|
|
220
|
+
min_s = '.'.join(map(str, self.min))
|
|
221
|
+
return f"{pluralize('version', plural)} at least {min_s}"
|
|
222
|
+
elif self.max is not None:
|
|
223
|
+
max_s = '.'.join(map(str, self.max))
|
|
224
|
+
return f"{pluralize('version', plural)} at most {max_s}"
|
|
225
|
+
else:
|
|
226
|
+
return f"version {pluralize('string', plural)}"
|
|
227
|
+
|
|
228
|
+
def into_data(self, val: t.Any) -> str:
|
|
229
|
+
return str(val)
|
|
230
|
+
|
|
231
|
+
@staticmethod
|
|
232
|
+
def parse_version(val: t.Union[str, t.Tuple[int, ...]]) -> t.Tuple[int, ...]:
|
|
233
|
+
def to_int(seg: t.Union[int, str]) -> int:
|
|
234
|
+
if isinstance(seg, int):
|
|
235
|
+
return seg
|
|
236
|
+
seg = seg.strip()
|
|
237
|
+
if not seg.isdigit():
|
|
238
|
+
raise ValueError()
|
|
239
|
+
return int(seg)
|
|
240
|
+
|
|
241
|
+
return tuple(map(to_int, val.split('.') if isinstance(val, str) else val))
|
|
242
|
+
|
|
243
|
+
def check_version(self, val: t.Tuple[int, ...]):
|
|
244
|
+
if self.min == self.max:
|
|
245
|
+
if self.min is not None and val != self.min:
|
|
246
|
+
raise ValueError(f"Version {'.'.join(map(str, val))} is not supported version {'.'.join(map(str, self.min))}")
|
|
247
|
+
elif self.min is not None and val < self.min:
|
|
248
|
+
raise ValueError(f"Version {'.'.join(map(str, val))} less than minimum supported version {'.'.join(map(str, self.min))}")
|
|
249
|
+
elif self.max is not None and val > self.max:
|
|
250
|
+
raise ValueError(f"Version {'.'.join(map(str, val))} greater than maximum supported version {'.'.join(map(str, self.max))}")
|
|
251
|
+
|
|
252
|
+
def try_convert(self, val: t.Any) -> str:
|
|
253
|
+
s = self.inner.try_convert(val)
|
|
254
|
+
try:
|
|
255
|
+
version = self.parse_version(val)
|
|
256
|
+
self.check_version(version)
|
|
257
|
+
except ValueError:
|
|
258
|
+
raise ParseInterrupt()
|
|
259
|
+
return s
|
|
260
|
+
|
|
261
|
+
def collect_errors(self, val: t.Any) -> t.Optional[ErrorNode]:
|
|
262
|
+
try:
|
|
263
|
+
s = self.inner.try_convert(val)
|
|
264
|
+
except ParseInterrupt:
|
|
265
|
+
return WrongTypeError(self.expected(), val)
|
|
266
|
+
try:
|
|
267
|
+
version = self.parse_version(s)
|
|
268
|
+
except ValueError:
|
|
269
|
+
return WrongTypeError(self.expected(), val, info='Invalid version string')
|
|
270
|
+
try:
|
|
271
|
+
self.check_version(version)
|
|
272
|
+
except ValueError as e:
|
|
273
|
+
return WrongTypeError(self.expected(), val, info=e.args[0])
|
|
274
|
+
return None
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
class _ReconsVarsConverter(Converter[t.FrozenSet[ReconsVar]]):
|
|
278
|
+
def __init__(self, ty: type, handlers: ConverterHandlers):
|
|
279
|
+
self.inner = make_converter(ty, handlers)
|
|
280
|
+
|
|
281
|
+
def expected(self, plural: bool = False) -> str:
|
|
282
|
+
known_params = t.get_args(ReconsVar)
|
|
283
|
+
return f"{pluralize('set', plural, article='a')} of comma-separated " \
|
|
284
|
+
f"variables ({list_phrase(tuple(map(repr, known_params)))})"
|
|
285
|
+
|
|
286
|
+
def into_data(self, val: t.Any) -> str:
|
|
287
|
+
return ", ".join(self.inner.into_data(val)) # type: ignore
|
|
288
|
+
|
|
289
|
+
def try_convert(self, val: t.Any) -> t.FrozenSet[ReconsVar]:
|
|
290
|
+
if isinstance(val, str):
|
|
291
|
+
val = tuple(v.strip() for v in val.split(","))
|
|
292
|
+
|
|
293
|
+
return self.inner.try_convert(val)
|
|
294
|
+
|
|
295
|
+
def collect_errors(self, val: t.Any) -> t.Optional[ErrorNode]:
|
|
296
|
+
if isinstance(val, str):
|
|
297
|
+
val = tuple(v.strip() for v in val.split(","))
|
|
298
|
+
|
|
299
|
+
return self.inner.collect_errors(val)
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
__all__ = [
|
|
303
|
+
'BackendName', 'Dataclass', 'Slices', 'Flag',
|
|
304
|
+
'process_flag', 'flag_any_true',
|
|
305
|
+
]
|
phaser/utils/__init__.py
ADDED
|
File without changes
|