Skip to content

Commit b742daa

Browse files
authored
feat: select tasks by marker metadata (#977)
1 parent 31f3269 commit b742daa

9 files changed

Lines changed: 438 additions & 40 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@ releases are available on [PyPI](https://pypi.org/project/pytask) and
77

88
## Unreleased
99

10+
- [#977](https://github.com/pytask-dev/pytask/pull/977) allows marker expressions
11+
passed to `-m` to select tasks by marker keyword arguments.
1012
- [#976](https://github.com/pytask-dev/pytask/pull/976) makes invalid marker and
1113
keyword expressions raise the standard `SyntaxError` instead of a custom parser
1214
exception.

‎docs/source/tutorials/selecting_tasks.md‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,26 @@ analysis.
3131
$ pytask -m "(data_management and not plots) or (analysis and plots)"
3232
```
3333

34+
Markers can also store metadata in keyword arguments. For example, the following marker
35+
records the data and model used by a prediction task.
36+
37+
```python
38+
@pytask.mark.prediction(data="customer_churn", model="random_forest")
39+
def task_create_predictions():
40+
pass
41+
```
42+
43+
Pass keyword arguments in a marker expression to select tasks by this metadata.
44+
45+
```console
46+
$ pytask -m 'prediction(data="customer_churn", model="random_forest")'
47+
```
48+
49+
A task matches if one marker with the requested name contains all keyword arguments in
50+
the expression. The marker can contain additional keyword arguments. Marker expressions
51+
support unescaped strings, positive and negative integers, `True`, `False`, and `None`.
52+
Keyword arguments are only supported in marker expressions passed to `-m`.
53+
3454
If you create your markers, use the
3555
[`pytask markers`](../reference_guides/commands.md#pytask-markers) command to register
3656
and document them.

‎src/_pytask/lock.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,9 @@ def _expression_filter(
8080
f"at column {e.offset}: {e.msg}"
8181
)
8282
raise ValueError(msg) from None
83+
if option == "-k" and compiled.has_keyword_arguments():
84+
msg = "Keyword expressions do not support call parameters."
85+
raise ValueError(msg)
8386

8487
return {
8588
task.signature for task in tasks if compiled.evaluate(matcher_from_task(task))

‎src/_pytask/mark/__init__.py‎

Lines changed: 24 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -151,11 +151,11 @@ def from_task(cls, task: PTask) -> KeywordMatcher:
151151

152152
return cls(mapped_names)
153153

154-
def __call__(self, subname: str) -> bool:
154+
def __call__(self, subname: str, /, **kwargs: str | int | bool | None) -> bool:
155155
subname = subname.lower()
156156
names = (name.lower() for name in self._names)
157157

158-
return any(subname in name for name in names)
158+
return not kwargs and any(subname in name for name in names)
159159

160160

161161
def select_by_keyword(session: Session, dag: DAG) -> set[str] | None:
@@ -171,6 +171,9 @@ def select_by_keyword(session: Session, dag: DAG) -> set[str] | None:
171171
f"Wrong expression passed to '-k': {e.text}: at column {e.offset}: {e.msg}"
172172
)
173173
raise ValueError(msg) from None
174+
if expression.has_keyword_arguments():
175+
msg = "Keyword expressions do not support call parameters."
176+
raise ValueError(msg)
174177

175178
remaining: set[str] = set()
176179
for task in session.tasks:
@@ -190,6 +193,9 @@ def select_by_after_keyword(session: Session, after: str) -> set[str]:
190193
f"at column {e.offset}: {e.msg}"
191194
)
192195
raise ValueError(msg) from None
196+
if expression.has_keyword_arguments():
197+
msg = "Keyword expressions do not support call parameters."
198+
raise ValueError(msg)
193199

194200
ancestors: set[str] = set()
195201
for task in session.tasks:
@@ -199,6 +205,9 @@ def select_by_after_keyword(session: Session, after: str) -> set[str]:
199205
return ancestors
200206

201207

208+
_NOT_SET = object()
209+
210+
202211
@dataclass(slots=True)
203212
class MarkMatcher:
204213
"""A matcher for markers which are present.
@@ -207,15 +216,22 @@ class MarkMatcher:
207216
208217
"""
209218

210-
own_mark_names: set[str]
219+
own_mark_name_mapping: dict[str, list[Mark]]
211220

212221
@classmethod
213222
def from_task(cls, task: PTask) -> MarkMatcher:
214-
mark_names = {mark.name for mark in task.markers}
215-
return cls(mark_names)
216-
217-
def __call__(self, name: str) -> bool:
218-
return name in self.own_mark_names
223+
mark_name_mapping: dict[str, list[Mark]] = {}
224+
for mark in task.markers:
225+
mark_name_mapping.setdefault(mark.name, []).append(mark)
226+
return cls(mark_name_mapping)
227+
228+
def __call__(self, name: str, /, **kwargs: str | int | bool | None) -> bool:
229+
for mark in self.own_mark_name_mapping.get(name, []):
230+
if all(
231+
mark.kwargs.get(key, _NOT_SET) == value for key, value in kwargs.items()
232+
):
233+
return True
234+
return False
219235

220236

221237
def select_by_mark(session: Session, dag: DAG) -> set[str] | None:

‎src/_pytask/mark/expression.py‎

Lines changed: 144 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -2,22 +2,22 @@
22
33
The grammar is:
44
5-
+------------+--------------------------------------------+
6-
| expression | expr? EOF |
7-
+------------+--------------------------------------------+
8-
| expr | and_expr ('or' and_expr)* |
9-
+------------+--------------------------------------------+
10-
| and_expr | not_expr ('and' not_expr)* |
11-
+------------+--------------------------------------------+
12-
| not_expr | ``'not' not_expr | '(' expr ')' | ident`` |
13-
+------------+--------------------------------------------+
14-
| ident | ``(\w|:|\+|-|\.|\[|\]|\\)+`` |
15-
+------------+--------------------------------------------+
5+
expression: expr? EOF
6+
expr: and_expr ('or' and_expr)*
7+
and_expr: not_expr ('and' not_expr)*
8+
not_expr: 'not' not_expr | '(' expr ')' | ident kwargs?
9+
10+
ident: (\w|:|\+|-|\.|\[|\]|\\|/)+
11+
kwargs: ('(' name '=' value (', ' name '=' value)* ')')
12+
name: a valid identifier that is not a reserved keyword
13+
value: unescaped string literal | (-)?[0-9]+ | 'False' | 'True' | 'None'
1614
1715
The semantics are:
1816
1917
- Empty expression evaluates to False.
20-
- ident evaluates to True of False according to a provided matcher function.
18+
- ident evaluates to True or False according to a provided matcher function.
19+
- ident with keyword arguments evaluates to True or False according to a provided
20+
matcher function.
2121
- or/and/not evaluate according to the usual boolean semantics.
2222
2323
This module is adapted from pytest's ``_pytest.mark.expression`` module:
@@ -29,20 +29,21 @@
2929

3030
import ast
3131
import enum
32+
import keyword
3233
import re
33-
from collections.abc import Callable
3434
from collections.abc import Iterator
3535
from collections.abc import Mapping
3636
from collections.abc import Sequence
3737
from dataclasses import dataclass
3838
from typing import TYPE_CHECKING
39+
from typing import Protocol
3940

4041
if TYPE_CHECKING:
4142
import types
4243
from typing import NoReturn
4344

4445

45-
__all__ = ["Expression"]
46+
__all__ = ["Expression", "ExpressionMatcher"]
4647

4748

4849
FILE_NAME = "<pytask match expression>"
@@ -56,6 +57,9 @@ class TokenType(enum.Enum):
5657
NOT = "not"
5758
IDENT = "identifier"
5859
EOF = "end of input"
60+
EQUAL = "="
61+
STRING = "string literal"
62+
COMMA = ","
5963

6064

6165
@dataclass(frozen=True, slots=True)
@@ -74,7 +78,7 @@ def __init__(self, input_: str) -> None:
7478
self.tokens = self.lex(input_)
7579
self.current = next(self.tokens)
7680

77-
def lex(self, input_: str) -> Iterator[Token]:
81+
def lex(self, input_: str) -> Iterator[Token]: # noqa: C901, PLR0912
7882
pos = 0
7983
while pos < len(input_):
8084
if input_[pos] in (" ", "\t"):
@@ -85,6 +89,29 @@ def lex(self, input_: str) -> Iterator[Token]:
8589
elif input_[pos] == ")":
8690
yield Token(TokenType.RPAREN, ")", pos)
8791
pos += 1
92+
elif input_[pos] == "=":
93+
yield Token(TokenType.EQUAL, "=", pos)
94+
pos += 1
95+
elif input_[pos] == ",":
96+
yield Token(TokenType.COMMA, ",", pos)
97+
pos += 1
98+
elif (quote_char := input_[pos]) in ("'", '"'):
99+
end_quote_pos = input_.find(quote_char, pos + 1)
100+
if end_quote_pos == -1:
101+
msg = f'closing quote "{quote_char}" is missing'
102+
raise SyntaxError(
103+
msg,
104+
(FILE_NAME, 1, pos + 1, input_),
105+
)
106+
value = input_[pos : end_quote_pos + 1]
107+
if (backslash_pos := value.find("\\")) != -1:
108+
msg = r'escaping with "\" not supported in marker expression'
109+
raise SyntaxError(
110+
msg,
111+
(FILE_NAME, 1, pos + backslash_pos + 1, input_),
112+
)
113+
yield Token(TokenType.STRING, value, pos)
114+
pos += len(value)
88115
else:
89116
match = re.match(r"(:?\w|:|\+|-|\.|\[|\]|/|\\)+", input_[pos:])
90117
if match:
@@ -168,19 +195,95 @@ def not_expr(s: Scanner) -> ast.expr:
168195
ident = s.accept(TokenType.IDENT)
169196
if ident:
170197
s.idents.add(ident.value)
171-
return ast.Name(IDENT_PREFIX + ident.value, ast.Load())
198+
name = ast.Name(IDENT_PREFIX + ident.value, ast.Load())
199+
if s.accept(TokenType.LPAREN):
200+
ret = ast.Call(func=name, args=[], keywords=all_kwargs(s))
201+
s.accept(TokenType.RPAREN, reject=True)
202+
else:
203+
ret = name
204+
return ret
172205
s.reject((TokenType.NOT, TokenType.LPAREN, TokenType.IDENT))
173206
return None # ty: ignore[invalid-return-type] # Unreachable: reject() raises
174207

175208

176-
class MatcherAdapter(Mapping[str, bool]):
209+
BUILTIN_MATCHERS = {"True": True, "False": False, "None": None}
210+
211+
212+
def single_kwarg(s: Scanner) -> ast.keyword:
213+
"""Parse one keyword argument."""
214+
keyword_name = s.accept(TokenType.IDENT, reject=True)
215+
assert keyword_name is not None
216+
if not keyword_name.value.isidentifier():
217+
msg = f"not a valid python identifier {keyword_name.value}"
218+
raise SyntaxError(
219+
msg,
220+
(FILE_NAME, 1, keyword_name.pos + 1, s.input),
221+
)
222+
if keyword.iskeyword(keyword_name.value):
223+
msg = f"unexpected reserved python keyword `{keyword_name.value}`"
224+
raise SyntaxError(
225+
msg,
226+
(FILE_NAME, 1, keyword_name.pos + 1, s.input),
227+
)
228+
s.accept(TokenType.EQUAL, reject=True)
229+
230+
if value_token := s.accept(TokenType.STRING):
231+
value: str | int | bool | None = value_token.value[1:-1]
232+
else:
233+
value_token = s.accept(TokenType.IDENT, reject=True)
234+
assert value_token is not None
235+
if (number := value_token.value).isdigit() or (
236+
number.startswith("-") and number[1:].isdigit()
237+
):
238+
value = int(number)
239+
elif value_token.value in BUILTIN_MATCHERS:
240+
value = BUILTIN_MATCHERS[value_token.value]
241+
else:
242+
msg = f'unexpected character/s "{value_token.value}"'
243+
raise SyntaxError(
244+
msg,
245+
(FILE_NAME, 1, value_token.pos + 1, s.input),
246+
)
247+
248+
return ast.keyword(keyword_name.value, ast.Constant(value))
249+
250+
251+
def all_kwargs(s: Scanner) -> list[ast.keyword]:
252+
"""Parse all keyword arguments."""
253+
kwargs = [single_kwarg(s)]
254+
while s.accept(TokenType.COMMA):
255+
kwargs.append(single_kwarg(s))
256+
return kwargs
257+
258+
259+
class ExpressionMatcher(Protocol):
260+
"""Match an identifier and optional keyword arguments."""
261+
262+
def __call__(self, name: str, /, **kwargs: str | int | bool | None) -> bool: ...
263+
264+
265+
@dataclass
266+
class MatcherNameAdapter:
267+
"""Adapt one matcher name to boolean and callable expression forms."""
268+
269+
matcher: ExpressionMatcher
270+
name: str
271+
272+
def __bool__(self) -> bool:
273+
return self.matcher(self.name)
274+
275+
def __call__(self, **kwargs: str | int | bool | None) -> bool:
276+
return self.matcher(self.name, **kwargs)
277+
278+
279+
class MatcherAdapter(Mapping[str, MatcherNameAdapter]):
177280
"""Adapts a matcher function to a locals mapping as required by eval()."""
178281

179-
def __init__(self, matcher: Callable[[str], bool]) -> None:
282+
def __init__(self, matcher: ExpressionMatcher) -> None:
180283
self.matcher = matcher
181284

182-
def __getitem__(self, key: str) -> bool:
183-
return self.matcher(key[len(IDENT_PREFIX) :])
285+
def __getitem__(self, key: str) -> MatcherNameAdapter:
286+
return MatcherNameAdapter(self.matcher, key[len(IDENT_PREFIX) :])
184287

185288
def __iter__(self) -> Iterator[str]: # pragma: no cover
186289
raise NotImplementedError
@@ -196,11 +299,17 @@ class Expression:
196299
197300
"""
198301

199-
__slots__ = ("_idents", "code")
302+
__slots__ = ("_has_keyword_arguments", "_idents", "code")
200303

201-
def __init__(self, code: types.CodeType, idents: frozenset[str]) -> None:
304+
def __init__(
305+
self,
306+
code: types.CodeType,
307+
idents: frozenset[str],
308+
has_keyword_arguments: bool,
309+
) -> None:
202310
self.code = code
203311
self._idents = idents
312+
self._has_keyword_arguments = has_keyword_arguments
204313

205314
@classmethod
206315
def compile_(cls, input_: str) -> Expression:
@@ -218,13 +327,20 @@ def compile_(cls, input_: str) -> Expression:
218327
filename="<pytask match expression>",
219328
mode="eval",
220329
)
221-
return cls(code, idents)
330+
has_keyword_arguments = any(
331+
isinstance(node, ast.Call) for node in ast.walk(astexpr)
332+
)
333+
return cls(code, idents, has_keyword_arguments)
222334

223335
def idents(self) -> frozenset[str]:
224336
"""Return all identifiers which appear in the expression."""
225337
return self._idents
226338

227-
def evaluate(self, matcher: Callable[[str], bool]) -> bool:
339+
def has_keyword_arguments(self) -> bool:
340+
"""Return whether the expression contains marker keyword arguments."""
341+
return self._has_keyword_arguments
342+
343+
def evaluate(self, matcher: ExpressionMatcher) -> bool:
228344
"""Evaluate the match expression.
229345
230346
Parameters
@@ -239,7 +355,8 @@ def evaluate(self, matcher: Callable[[str], bool]) -> bool:
239355
Whether the expression matches or not.
240356
241357
"""
242-
ret: bool = eval( # noqa: S307
243-
self.code, {"__builtins__": {}}, MatcherAdapter(matcher)
358+
return bool(
359+
eval( # noqa: S307
360+
self.code, {"__builtins__": {}}, MatcherAdapter(matcher)
361+
)
244362
)
245-
return ret

0 commit comments

Comments
 (0)