Skip to content

Commit ec30207

Browse files
authored
Support tuple destructuring in implied for's. (#220)
* Add docs and coverage for tuple-target comprehension lowering # Conflicts: # README.md # docs/source/generic/query_structure.md # func_adl/ast/syntatic_sugar.py * Fix up
1 parent 83390ac commit ec30207

4 files changed

Lines changed: 87 additions & 10 deletions

File tree

‎README.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,7 @@ There are several python expressions and idioms that are translated behind your
5858
|List Comprehension | `[j.pt() for j in jets]` | `jets.Select(lambda j: j.pt())` |
5959
|List Comprehension | `[j.pt() for j in jets if abs(j.eta()) < 2.4]` | `jets.Where(lambda j: abs(j.eta()) < 2.4).Select(lambda j: j.pt())` |
6060
|Multi-generator comprehension|`[j.pt() + e.pt() for j in jets for e in electrons]`|`jets.SelectMany(lambda j: electrons.Select(lambda e: j.pt() + e.pt()))`|
61+
|Destructuring List Comprehension|`[a+b for a,b in pairs]`|`pairs.Select(lambda tmp: tmp[0] + tmp[1])`|
6162
|Literal List Comprehension|`[i for i in [1, 2, 3]]`|`[1, 2, 3]`|
6263
| Data Classes<br>(typed) | `@dataclass`<br>`class my_data:`<br>`x: ObjectStream[Jets]`<br><br>`Select(lambda e: my_data(x=e.Jets()).x)` | `Select(lambda e: {'x': e.Jets()}.x)` |
6364
| Named Tuple<br>(typed) | `class my_data(NamedTuple):`<br>`x: ObjectStream[Jets]`<br><br>`Select(lambda e: my_data(x=e.Jets()).x)` | `Select(lambda e: {'x': e.Jets()}.x)` |

‎docs/source/generic/query_structure.md‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,9 @@ expressions:
4242
- List/generator comprehensions over streams are lowered to query operators.
4343
Multi-generator forms (`for ... for ...`) flatten outer levels via `.SelectMany(...)`
4444
so the stream shape matches Python iteration semantics.
45+
- Comprehensions with tuple/list destructuring targets are supported. For non-literal
46+
iterables, destructured names are rewritten as indexed access on an internal temporary
47+
value (for example, `[a+b for a,b in pairs]` behaves like `pairs.Select(lambda tmp: tmp[0] + tmp[1])`).
4548
- List comprehensions over literal iterables are expanded directly. For example,
4649
`[i for i in [1, 2, 3]]` becomes `[1, 2, 3]`.
4750
- `any`/`all` over literal lists/tuples are reduced to boolean `or`/`and` expressions.

‎func_adl/ast/syntatic_sugar.py‎

Lines changed: 53 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -70,6 +70,37 @@ def _target_bindings(
7070
f" - {ast.unparse(node)}"
7171
)
7272

73+
def _target_bindings_from_value(
74+
self, target: ast.AST, value: ast.expr, node: ast.AST
75+
) -> Dict[str, ast.expr]:
76+
"""Build bindings for comprehension targets, using indexing for non-literals."""
77+
78+
# Reuse literal destructuring where possible to preserve exact tuple/list nodes.
79+
literal_bindings = self._target_bindings(target, value, node)
80+
if literal_bindings is not None:
81+
return literal_bindings
82+
83+
if isinstance(target, ast.Name):
84+
return {target.id: copy.deepcopy(value)}
85+
86+
if isinstance(target, (ast.Tuple, ast.List)):
87+
bindings: Dict[str, ast.expr] = {}
88+
for index, target_elt in enumerate(target.elts):
89+
element_value = ast.Subscript(
90+
value=copy.deepcopy(value),
91+
slice=ast.Constant(value=index),
92+
ctx=ast.Load(),
93+
)
94+
bindings.update(
95+
self._target_bindings_from_value(target_elt, element_value, node)
96+
)
97+
return bindings
98+
99+
raise ValueError(
100+
f"Comprehension variable must be a name or tuple/list, but found {target}"
101+
f" - {ast.unparse(node)}"
102+
)
103+
73104
def _substitute_names(self, expr: ast.expr, bindings: Dict[str, ast.expr]) -> ast.expr:
74105
class _name_replacer(ast.NodeTransformer):
75106
def __init__(self, loop_bindings: Dict[str, ast.expr]):
@@ -371,25 +402,42 @@ def resolve_generator(
371402
"""
372403
a = node
373404
generator_count: int = len(generators)
405+
temp_counter = 0
374406
for index, c in enumerate(reversed(generators)):
375407
target = c.target
376-
if not isinstance(target, ast.Name):
377-
# Keep original comprehension for unsupported lowering cases.
378-
return node
379408
if c.is_async:
380409
raise ValueError(f"Comprehension can't be async - {ast.unparse(node)}.")
381410
source_collection = c.iter
382411

412+
if isinstance(target, ast.Name):
413+
lambda_arg_name = target.id
414+
elif isinstance(target, (ast.Tuple, ast.List)):
415+
lambda_arg_name = f"__fa_tmp_{temp_counter}"
416+
temp_counter += 1
417+
else:
418+
raise ValueError(
419+
"Comprehension variable must be a name or tuple/list, "
420+
f"but found {target} - {ast.unparse(node)}"
421+
)
422+
423+
target_bindings = self._target_bindings_from_value(
424+
target,
425+
ast.Name(id=lambda_arg_name, ctx=ast.Load()),
426+
node,
427+
)
428+
383429
# Turn the if clauses into Where statements
384430
for a_if in c.ifs:
385-
where_function = lambda_build(target.id, a_if)
431+
where_body = self._substitute_names(a_if, target_bindings)
432+
where_function = lambda_build(lambda_arg_name, where_body)
386433
source_collection = ast.Call(
387434
func=ast.Attribute(attr="Where", value=source_collection, ctx=ast.Load()),
388435
args=[where_function],
389436
keywords=[],
390437
)
391438

392-
lambda_function = lambda_build(target.id, lambda_body)
439+
rewritten_lambda_body = self._substitute_names(lambda_body, target_bindings)
440+
lambda_function = lambda_build(lambda_arg_name, rewritten_lambda_body)
393441
use_select_many = generator_count > 1 and index > 0
394442
a = ast.Call(
395443
func=ast.Attribute(

‎tests/ast/test_syntatic_sugar.py‎

Lines changed: 30 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -90,13 +90,38 @@ def test_resolve_3generator_list_comp_flattened_shape():
9090
) == ast.dump(a_new)
9191

9292

93-
def test_resolve_bad_iterator():
94-
a = ast.parse("[j.pt() for idx,j in enumerate(jets)]")
93+
def test_resolve_tuple_target():
94+
a = ast.parse("[a+b for a,b in pairs]")
9595
a_new = resolve_syntatic_sugar(a)
9696

97-
# Unsupported lowering (tuple target with non-literal source) should be
98-
# preserved for downstream processing.
99-
assert ast.unparse(a_new) == ast.unparse(a)
97+
assert ast.dump(
98+
ast.parse("pairs.Select(lambda __fa_tmp_0: __fa_tmp_0[0] + __fa_tmp_0[1])")
99+
) == ast.dump(a_new)
100+
101+
102+
def test_resolve_tuple_target_nested_with_if():
103+
a = ast.parse("[a + c for (a, (b, c)) in triples if b > 0 if c < 10]")
104+
a_new = resolve_syntatic_sugar(a)
105+
106+
assert ast.dump(
107+
ast.parse(
108+
"triples.Where(lambda __fa_tmp_0: __fa_tmp_0[1][0] > 0)"
109+
".Where(lambda __fa_tmp_0: __fa_tmp_0[1][1] < 10)"
110+
".Select(lambda __fa_tmp_0: __fa_tmp_0[0] + __fa_tmp_0[1][1])"
111+
)
112+
) == ast.dump(a_new)
113+
114+
115+
def test_resolve_tuple_target_from_enumerate_with_if():
116+
a = ast.parse("[idx + j.pt() for idx, j in enumerate(jets) if idx > 0]")
117+
a_new = resolve_syntatic_sugar(a)
118+
119+
assert ast.dump(
120+
ast.parse(
121+
"enumerate(jets).Where(lambda __fa_tmp_0: __fa_tmp_0[0] > 0)"
122+
".Select(lambda __fa_tmp_0: __fa_tmp_0[0] + __fa_tmp_0[1].pt())"
123+
)
124+
) == ast.dump(a_new)
100125

101126

102127
def test_resolve_no_async():

0 commit comments

Comments
 (0)