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#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
13pub struct ListComp {
14 pub elt: Box<ExprType>,
16 pub generators: Vec<Comprehension>,
18 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#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
27pub struct SetComp {
28 pub elt: Box<ExprType>,
30 pub generators: Vec<Comprehension>,
32 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#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
41pub struct GeneratorExp {
42 pub elt: Box<ExprType>,
44 pub generators: Vec<Comprehension>,
46 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#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
55pub struct DictComp {
56 pub key: Box<ExprType>,
58 pub value: Box<ExprType>,
60 pub generators: Vec<Comprehension>,
62 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#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
71pub struct Comprehension {
72 pub target: ExprType,
74 pub iter: ExprType,
76 pub ifs: Vec<ExprType>,
78 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 let elt = ob.extract_attr_with_context("elt", "list comprehension element")?;
87 let elt: ExprType = elt.extract()?;
88
89 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 let elt = ob.extract_attr_with_context("elt", "set comprehension element")?;
108 let elt: ExprType = elt.extract()?;
109
110 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 let elt = ob.extract_attr_with_context("elt", "generator expression element")?;
129 let elt: ExprType = elt.extract()?;
130
131 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 let key = ob.extract_attr_with_context("key", "dict comprehension key")?;
150 let key: ExprType = key.extract()?;
151
152 let value = ob.extract_attr_with_context("value", "dict comprehension value")?;
154 let value: ExprType = value.extract()?;
155
156 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 let target = ob.extract_attr_with_context("target", "comprehension target")?;
176 let target: ExprType = target.extract()?;
177
178 let iter = ob.extract_attr_with_context("iter", "comprehension iter")?;
180 let iter: ExprType = iter.extract()?;
181
182 let ifs: Vec<ExprType> = extract_list(&ob, "ifs", "comprehension conditions").unwrap_or_default();
184
185 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
225fn 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 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 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 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 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 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 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 }