Skip to main content

python_ast/ast/tree/
list_comp.rs

1use proc_macro2::TokenStream;
2use pyo3::{Borrowed, FromPyObject, PyAny, PyResult, prelude::PyAnyMethods};
3use quote::quote;
4use serde::{Deserialize, Serialize};
5
6use crate::{
7    CodeGen, CodeGenContext, ExprType, Node, PythonOptions, SymbolTableScopes,
8    PyAttributeExtractor, extract_list,
9};
10
11/// List comprehension (e.g., [x ** 2 for x in range(10) if x % 2 == 0])
12#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
13pub struct ListComp {
14    /// The element expression being computed
15    pub elt: Box<ExprType>,
16    /// The generators (for clauses)
17    pub generators: Vec<Comprehension>,
18    /// Position information
19    pub lineno: Option<usize>,
20    pub col_offset: Option<usize>,
21    pub end_lineno: Option<usize>,
22    pub end_col_offset: Option<usize>,
23}
24
25/// Set comprehension (e.g., {x for x in range(10) if x % 2 == 0})
26#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
27pub struct SetComp {
28    /// The element expression being computed
29    pub elt: Box<ExprType>,
30    /// The generators (for clauses)
31    pub generators: Vec<Comprehension>,
32    /// Position information
33    pub lineno: Option<usize>,
34    pub col_offset: Option<usize>,
35    pub end_lineno: Option<usize>,
36    pub end_col_offset: Option<usize>,
37}
38
39/// Generator expression (e.g., (x for x in range(10) if x % 2 == 0))
40#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
41pub struct GeneratorExp {
42    /// The element expression being computed
43    pub elt: Box<ExprType>,
44    /// The generators (for clauses)
45    pub generators: Vec<Comprehension>,
46    /// Position information
47    pub lineno: Option<usize>,
48    pub col_offset: Option<usize>,
49    pub end_lineno: Option<usize>,
50    pub end_col_offset: Option<usize>,
51}
52
53/// Dictionary comprehension (e.g., {k: v for k, v in items.items()})
54#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
55pub struct DictComp {
56    /// The key expression being computed
57    pub key: Box<ExprType>,
58    /// The value expression being computed
59    pub value: Box<ExprType>,
60    /// The generators (for clauses)
61    pub generators: Vec<Comprehension>,
62    /// Position information
63    pub lineno: Option<usize>,
64    pub col_offset: Option<usize>,
65    pub end_lineno: Option<usize>,
66    pub end_col_offset: Option<usize>,
67}
68
69/// A comprehension generator (for x in iter if condition)
70#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
71pub struct Comprehension {
72    /// The target variable(s) (e.g., x in "for x in range(10)")
73    pub target: ExprType,
74    /// The iterable expression (e.g., range(10) in "for x in range(10)")
75    pub iter: ExprType,
76    /// The conditions (if clauses)
77    pub ifs: Vec<ExprType>,
78    /// Whether this is an async comprehension
79    pub is_async: bool,
80}
81
82impl<'a, 'py> FromPyObject<'a, 'py> for ListComp {
83    type Error = pyo3::PyErr;
84    fn extract(ob: Borrowed<'a, 'py, PyAny>) -> PyResult<Self> {
85        // Extract the element expression
86        let elt = ob.extract_attr_with_context("elt", "list comprehension element")?;
87        let elt: ExprType = elt.extract()?;
88        
89        // Extract generators
90        let generators: Vec<Comprehension> = extract_list(&ob, "generators", "list comprehension generators")?;
91        
92        Ok(ListComp {
93            elt: Box::new(elt),
94            generators,
95            lineno: ob.lineno(),
96            col_offset: ob.col_offset(),
97            end_lineno: ob.end_lineno(),
98            end_col_offset: ob.end_col_offset(),
99        })
100    }
101}
102
103impl<'a, 'py> FromPyObject<'a, 'py> for SetComp {
104    type Error = pyo3::PyErr;
105    fn extract(ob: Borrowed<'a, 'py, PyAny>) -> PyResult<Self> {
106        // Extract the element expression
107        let elt = ob.extract_attr_with_context("elt", "set comprehension element")?;
108        let elt: ExprType = elt.extract()?;
109        
110        // Extract generators
111        let generators: Vec<Comprehension> = extract_list(&ob, "generators", "set comprehension generators")?;
112        
113        Ok(SetComp {
114            elt: Box::new(elt),
115            generators,
116            lineno: ob.lineno(),
117            col_offset: ob.col_offset(),
118            end_lineno: ob.end_lineno(),
119            end_col_offset: ob.end_col_offset(),
120        })
121    }
122}
123
124impl<'a, 'py> FromPyObject<'a, 'py> for GeneratorExp {
125    type Error = pyo3::PyErr;
126    fn extract(ob: Borrowed<'a, 'py, PyAny>) -> PyResult<Self> {
127        // Extract the element expression
128        let elt = ob.extract_attr_with_context("elt", "generator expression element")?;
129        let elt: ExprType = elt.extract()?;
130        
131        // Extract generators
132        let generators: Vec<Comprehension> = extract_list(&ob, "generators", "generator expression generators")?;
133        
134        Ok(GeneratorExp {
135            elt: Box::new(elt),
136            generators,
137            lineno: ob.lineno(),
138            col_offset: ob.col_offset(),
139            end_lineno: ob.end_lineno(),
140            end_col_offset: ob.end_col_offset(),
141        })
142    }
143}
144
145impl<'a, 'py> FromPyObject<'a, 'py> for DictComp {
146    type Error = pyo3::PyErr;
147    fn extract(ob: Borrowed<'a, 'py, PyAny>) -> PyResult<Self> {
148        // Extract the key expression
149        let key = ob.extract_attr_with_context("key", "dict comprehension key")?;
150        let key: ExprType = key.extract()?;
151        
152        // Extract the value expression
153        let value = ob.extract_attr_with_context("value", "dict comprehension value")?;
154        let value: ExprType = value.extract()?;
155        
156        // Extract generators
157        let generators: Vec<Comprehension> = extract_list(&ob, "generators", "dict comprehension generators")?;
158        
159        Ok(DictComp {
160            key: Box::new(key),
161            value: Box::new(value),
162            generators,
163            lineno: ob.lineno(),
164            col_offset: ob.col_offset(),
165            end_lineno: ob.end_lineno(),
166            end_col_offset: ob.end_col_offset(),
167        })
168    }
169}
170
171impl<'a, 'py> FromPyObject<'a, 'py> for Comprehension {
172    type Error = pyo3::PyErr;
173    fn extract(ob: Borrowed<'a, 'py, PyAny>) -> PyResult<Self> {
174        // Extract target
175        let target = ob.extract_attr_with_context("target", "comprehension target")?;
176        let target: ExprType = target.extract()?;
177        
178        // Extract iter
179        let iter = ob.extract_attr_with_context("iter", "comprehension iter")?;
180        let iter: ExprType = iter.extract()?;
181        
182        // Extract ifs (list of conditions)
183        let ifs: Vec<ExprType> = extract_list(&ob, "ifs", "comprehension conditions").unwrap_or_default();
184        
185        // Extract is_async
186        let is_async: bool = ob.getattr("is_async")?.extract().unwrap_or(false);
187        
188        Ok(Comprehension {
189            target,
190            iter,
191            ifs,
192            is_async,
193        })
194    }
195}
196
197impl Node for ListComp {
198    fn lineno(&self) -> Option<usize> { self.lineno }
199    fn col_offset(&self) -> Option<usize> { self.col_offset }
200    fn end_lineno(&self) -> Option<usize> { self.end_lineno }
201    fn end_col_offset(&self) -> Option<usize> { self.end_col_offset }
202}
203
204impl Node for SetComp {
205    fn lineno(&self) -> Option<usize> { self.lineno }
206    fn col_offset(&self) -> Option<usize> { self.col_offset }
207    fn end_lineno(&self) -> Option<usize> { self.end_lineno }
208    fn end_col_offset(&self) -> Option<usize> { self.end_col_offset }
209}
210
211impl Node for GeneratorExp {
212    fn lineno(&self) -> Option<usize> { self.lineno }
213    fn col_offset(&self) -> Option<usize> { self.col_offset }
214    fn end_lineno(&self) -> Option<usize> { self.end_lineno }
215    fn end_col_offset(&self) -> Option<usize> { self.end_col_offset }
216}
217
218impl Node for DictComp {
219    fn lineno(&self) -> Option<usize> { self.lineno }
220    fn col_offset(&self) -> Option<usize> { self.col_offset }
221    fn end_lineno(&self) -> Option<usize> { self.end_lineno }
222    fn end_col_offset(&self) -> Option<usize> { self.end_col_offset }
223}
224
225/// Lower a comprehension's generator clauses into nested `for` loops around
226/// `inner`, binding each generator's real target name (so the element and
227/// condition expressions can reference it) and applying `if` guards with
228/// `continue`. Generators nest left-to-right, matching Python's evaluation
229/// order, and later generators may reference earlier targets.
230fn build_comprehension_loops(
231    generators: &[Comprehension],
232    inner: TokenStream,
233    ctx: &CodeGenContext,
234    options: &PythonOptions,
235    symbols: &SymbolTableScopes,
236) -> Result<TokenStream, Box<dyn std::error::Error>> {
237    let mut acc = inner;
238    for generator in generators.iter().rev() {
239        let target = generator
240            .target
241            .clone()
242            .to_rust(ctx.clone(), options.clone(), symbols.clone())?;
243        let iter_expr = generator
244            .iter
245            .clone()
246            .to_rust(ctx.clone(), options.clone(), symbols.clone())?;
247        let conditions: Result<Vec<_>, _> = generator
248            .ifs
249            .iter()
250            .map(|if_expr| {
251                if_expr
252                    .clone()
253                    .to_rust(ctx.clone(), options.clone(), symbols.clone())
254            })
255            .collect();
256        let conditions = conditions?;
257        let guard = if conditions.is_empty() {
258            quote!()
259        } else {
260            quote! { if !( #((#conditions))&&* ) { continue; } }
261        };
262        acc = quote! {
263            for #target in #iter_expr {
264                #guard
265                #acc
266            }
267        };
268    }
269    Ok(acc)
270}
271
272impl CodeGen for ListComp {
273    type Context = CodeGenContext;
274    type Options = PythonOptions;
275    type SymbolTable = SymbolTableScopes;
276
277    fn find_symbols(self, symbols: Self::SymbolTable) -> Self::SymbolTable {
278        // Process the element and generators
279        let symbols = (*self.elt).clone().find_symbols(symbols);
280        self.generators.into_iter().fold(symbols, |acc, generator| {
281            let acc = generator.target.find_symbols(acc);
282            let acc = generator.iter.find_symbols(acc);
283            generator.ifs.into_iter().fold(acc, |acc, if_expr| if_expr.find_symbols(acc))
284        })
285    }
286
287    fn to_rust(
288        self,
289        ctx: Self::Context,
290        options: Self::Options,
291        symbols: Self::SymbolTable,
292    ) -> Result<TokenStream, Box<dyn std::error::Error>> {
293        let elt = (*self.elt).clone().to_rust(ctx.clone(), options.clone(), symbols.clone())?;
294        let loops = build_comprehension_loops(
295            &self.generators,
296            quote! { __rython_comp.push(#elt); },
297            &ctx,
298            &options,
299            &symbols,
300        )?;
301        Ok(quote! {
302            {
303                let mut __rython_comp = Vec::new();
304                #loops
305                __rython_comp
306            }
307        })
308    }
309}
310
311impl CodeGen for SetComp {
312    type Context = CodeGenContext;
313    type Options = PythonOptions;
314    type SymbolTable = SymbolTableScopes;
315
316    fn find_symbols(self, symbols: Self::SymbolTable) -> Self::SymbolTable {
317        // Process the element and generators
318        let symbols = (*self.elt).clone().find_symbols(symbols);
319        self.generators.into_iter().fold(symbols, |acc, generator| {
320            let acc = generator.target.find_symbols(acc);
321            let acc = generator.iter.find_symbols(acc);
322            generator.ifs.into_iter().fold(acc, |acc, if_expr| if_expr.find_symbols(acc))
323        })
324    }
325
326    fn to_rust(
327        self,
328        ctx: Self::Context,
329        options: Self::Options,
330        symbols: Self::SymbolTable,
331    ) -> Result<TokenStream, Box<dyn std::error::Error>> {
332        let elt = (*self.elt).clone().to_rust(ctx.clone(), options.clone(), symbols.clone())?;
333        let loops = build_comprehension_loops(
334            &self.generators,
335            quote! { __rython_comp.insert(#elt); },
336            &ctx,
337            &options,
338            &symbols,
339        )?;
340        Ok(quote! {
341            {
342                let mut __rython_comp = std::collections::HashSet::new();
343                #loops
344                __rython_comp
345            }
346        })
347    }
348}
349
350impl CodeGen for GeneratorExp {
351    type Context = CodeGenContext;
352    type Options = PythonOptions;
353    type SymbolTable = SymbolTableScopes;
354
355    fn find_symbols(self, symbols: Self::SymbolTable) -> Self::SymbolTable {
356        // Process the element and generators
357        let symbols = (*self.elt).clone().find_symbols(symbols);
358        self.generators.into_iter().fold(symbols, |acc, generator| {
359            let acc = generator.target.find_symbols(acc);
360            let acc = generator.iter.find_symbols(acc);
361            generator.ifs.into_iter().fold(acc, |acc, if_expr| if_expr.find_symbols(acc))
362        })
363    }
364
365    fn to_rust(
366        self,
367        ctx: Self::Context,
368        options: Self::Options,
369        symbols: Self::SymbolTable,
370    ) -> Result<TokenStream, Box<dyn std::error::Error>> {
371        // Generator expressions are lowered eagerly (like a list
372        // comprehension) and then turned back into an iterator; Python's lazy
373        // evaluation is not modeled yet.
374        let elt = (*self.elt).clone().to_rust(ctx.clone(), options.clone(), symbols.clone())?;
375        let loops = build_comprehension_loops(
376            &self.generators,
377            quote! { __rython_comp.push(#elt); },
378            &ctx,
379            &options,
380            &symbols,
381        )?;
382        Ok(quote! {
383            {
384                let mut __rython_comp = Vec::new();
385                #loops
386                __rython_comp.into_iter()
387            }
388        })
389    }
390}
391
392impl CodeGen for DictComp {
393    type Context = CodeGenContext;
394    type Options = PythonOptions;
395    type SymbolTable = SymbolTableScopes;
396
397    fn find_symbols(self, symbols: Self::SymbolTable) -> Self::SymbolTable {
398        // Process the key, value and generators
399        let symbols = (*self.key).clone().find_symbols(symbols);
400        let symbols = (*self.value).clone().find_symbols(symbols);
401        self.generators.into_iter().fold(symbols, |acc, generator| {
402            let acc = generator.target.find_symbols(acc);
403            let acc = generator.iter.find_symbols(acc);
404            generator.ifs.into_iter().fold(acc, |acc, if_expr| if_expr.find_symbols(acc))
405        })
406    }
407
408    fn to_rust(
409        self,
410        ctx: Self::Context,
411        options: Self::Options,
412        symbols: Self::SymbolTable,
413    ) -> Result<TokenStream, Box<dyn std::error::Error>> {
414        let key = (*self.key).clone().to_rust(ctx.clone(), options.clone(), symbols.clone())?;
415        let value = (*self.value).clone().to_rust(ctx.clone(), options.clone(), symbols.clone())?;
416        let loops = build_comprehension_loops(
417            &self.generators,
418            quote! { __rython_comp.insert(#key, #value); },
419            &ctx,
420            &options,
421            &symbols,
422        )?;
423        // PyDict, like dict literals: comprehension-built dicts preserve
424        // insertion order too.
425        Ok(quote! {
426            {
427                let mut __rython_comp = PyDict::new();
428                #loops
429                __rython_comp
430            }
431        })
432    }
433}
434
435#[cfg(test)]
436mod tests {
437    // Note: These tests might need additional AST node implementations
438    // create_parse_test!(test_simple_listcomp, "[x for x in range(5)]", "test.py");
439    // create_parse_test!(test_listcomp_with_condition, "[x for x in range(10) if x % 2 == 0]", "test.py");
440}