Skip to main content

rustpython_codegen/
preprocess.rs

1use alloc::{boxed::Box, string::String, vec::Vec};
2
3use ruff_python_ast::{
4    self as ast, AtomicNodeIndex, ConversionFlag, Expr, ExprFString, FString, FStringFlags,
5    FStringValue, HasNodeIndex, InterpolatedElement, InterpolatedStringElement,
6    InterpolatedStringElements, InterpolatedStringFormatSpec, InterpolatedStringLiteralElement,
7    Operator,
8    visitor::transformer::{self, Transformer},
9};
10use ruff_text_size::{Ranged, TextRange};
11
12use crate::compile::FutureFeature;
13use rustpython_compiler_core::bytecode;
14
15const MAXDIGITS: usize = 3;
16const F_LJUST: u8 = 1;
17
18/// ast_preprocess.c ControlFlowInFinallyContext
19#[derive(Clone, Copy)]
20struct ControlFlowInFinallyContext {
21    in_finally: bool,
22    in_funcdef: bool,
23    in_loop: bool,
24}
25
26/// ast_preprocess.c before_return
27fn before_return<E>(
28    contexts: &[ControlFlowInFinallyContext],
29    range: TextRange,
30    warn: &mut impl FnMut(TextRange, String) -> Result<(), E>,
31) -> Result<(), E> {
32    if let Some(ctx) = contexts.last()
33        && ctx.in_finally
34        && !ctx.in_funcdef
35    {
36        warn(range, "'return' in a 'finally' block".to_owned())?;
37    }
38    Ok(())
39}
40
41/// ast_preprocess.c before_loop_exit
42fn before_loop_exit<E>(
43    contexts: &[ControlFlowInFinallyContext],
44    range: TextRange,
45    kw: &str,
46    warn: &mut impl FnMut(TextRange, String) -> Result<(), E>,
47) -> Result<(), E> {
48    if let Some(ctx) = contexts.last()
49        && ctx.in_finally
50        && !ctx.in_loop
51    {
52        warn(range, format!("'{kw}' in a 'finally' block"))?;
53    }
54    Ok(())
55}
56
57fn visit_body_with_control_flow_context<E>(
58    body: &[ast::Stmt],
59    contexts: &mut Vec<ControlFlowInFinallyContext>,
60    warn: &mut impl FnMut(TextRange, String) -> Result<(), E>,
61    in_finally: bool,
62    in_funcdef: bool,
63    in_loop: bool,
64) -> Result<(), E> {
65    contexts.push(ControlFlowInFinallyContext {
66        in_finally,
67        in_funcdef,
68        in_loop,
69    });
70    visit_body_for_control_flow_in_finally(body, contexts, warn)?;
71    contexts.pop();
72    Ok(())
73}
74
75fn visit_body_for_control_flow_in_finally<E>(
76    body: &[ast::Stmt],
77    contexts: &mut Vec<ControlFlowInFinallyContext>,
78    warn: &mut impl FnMut(TextRange, String) -> Result<(), E>,
79) -> Result<(), E> {
80    for stmt in body {
81        visit_stmt_for_control_flow_in_finally(stmt, contexts, warn)?;
82    }
83    Ok(())
84}
85
86/// ast_preprocess.c astfold_stmt control-flow warning traversal.
87fn visit_stmt_for_control_flow_in_finally<E>(
88    stmt: &ast::Stmt,
89    contexts: &mut Vec<ControlFlowInFinallyContext>,
90    warn: &mut impl FnMut(TextRange, String) -> Result<(), E>,
91) -> Result<(), E> {
92    match stmt {
93        ast::Stmt::FunctionDef(function) => {
94            visit_body_with_control_flow_context(
95                &function.body,
96                contexts,
97                warn,
98                false,
99                true,
100                false,
101            )?;
102        }
103        ast::Stmt::ClassDef(class) => {
104            visit_body_for_control_flow_in_finally(&class.body, contexts, warn)?;
105        }
106        ast::Stmt::Return(return_stmt) => {
107            before_return(contexts, return_stmt.range, warn)?;
108        }
109        ast::Stmt::For(for_stmt) => {
110            visit_body_with_control_flow_context(
111                &for_stmt.body,
112                contexts,
113                warn,
114                false,
115                false,
116                true,
117            )?;
118            visit_body_for_control_flow_in_finally(&for_stmt.orelse, contexts, warn)?;
119        }
120        ast::Stmt::While(while_stmt) => {
121            visit_body_with_control_flow_context(
122                &while_stmt.body,
123                contexts,
124                warn,
125                false,
126                false,
127                true,
128            )?;
129            visit_body_for_control_flow_in_finally(&while_stmt.orelse, contexts, warn)?;
130        }
131        ast::Stmt::If(if_stmt) => {
132            visit_body_for_control_flow_in_finally(&if_stmt.body, contexts, warn)?;
133            for clause in &if_stmt.elif_else_clauses {
134                visit_body_for_control_flow_in_finally(&clause.body, contexts, warn)?;
135            }
136        }
137        ast::Stmt::Try(try_stmt) => {
138            visit_body_for_control_flow_in_finally(&try_stmt.body, contexts, warn)?;
139            for handler in &try_stmt.handlers {
140                match handler {
141                    ast::ExceptHandler::ExceptHandler(handler) => {
142                        visit_body_for_control_flow_in_finally(&handler.body, contexts, warn)?;
143                    }
144                }
145            }
146            visit_body_for_control_flow_in_finally(&try_stmt.orelse, contexts, warn)?;
147            visit_body_with_control_flow_context(
148                &try_stmt.finalbody,
149                contexts,
150                warn,
151                true,
152                false,
153                false,
154            )?;
155        }
156        ast::Stmt::With(with_stmt) => {
157            visit_body_for_control_flow_in_finally(&with_stmt.body, contexts, warn)?;
158        }
159        ast::Stmt::Match(match_stmt) => {
160            for case in &match_stmt.cases {
161                visit_body_for_control_flow_in_finally(&case.body, contexts, warn)?;
162            }
163        }
164        ast::Stmt::Break(break_stmt) => {
165            before_loop_exit(contexts, break_stmt.range, "break", warn)?;
166        }
167        ast::Stmt::Continue(continue_stmt) => {
168            before_loop_exit(contexts, continue_stmt.range, "continue", warn)?;
169        }
170        _ => {}
171    }
172    Ok(())
173}
174
175/// ast_preprocess.c control_flow_in_finally_warning
176pub fn warn_control_flow_in_finally<E>(
177    module: &ast::Mod,
178    mut warn: impl FnMut(TextRange, String) -> Result<(), E>,
179) -> Result<(), E> {
180    let mut contexts = Vec::new();
181    match module {
182        ast::Mod::Module(module) => {
183            visit_body_for_control_flow_in_finally(&module.body, &mut contexts, &mut warn)?;
184        }
185        ast::Mod::Expression(_) => {}
186    }
187    Ok(())
188}
189
190pub fn has_future_annotations(module: &ast::Mod) -> bool {
191    future_features(module).contains(bytecode::CodeFlags::FUTURE_ANNOTATIONS)
192}
193
194pub fn future_features(module: &ast::Mod) -> bytecode::CodeFlags {
195    checked_future_features(module).unwrap_or_else(|err| err.features)
196}
197
198pub struct FutureFeatureError {
199    pub features: bytecode::CodeFlags,
200    pub range: TextRange,
201    pub kind: FutureFeatureErrorKind,
202}
203
204pub enum FutureFeatureErrorKind {
205    InvalidFeature(String),
206    InvalidBraces,
207}
208
209pub fn checked_future_features(
210    module: &ast::Mod,
211) -> Result<bytecode::CodeFlags, FutureFeatureError> {
212    let ast::Mod::Module(module) = module else {
213        return Ok(bytecode::CodeFlags::empty());
214    };
215    checked_future_features_in_body(&module.body)
216}
217
218pub fn checked_future_features_in_body(
219    body: &[ast::Stmt],
220) -> Result<bytecode::CodeFlags, FutureFeatureError> {
221    let mut future_features = bytecode::CodeFlags::empty();
222    let mut statements = body.iter();
223    if let Some(ast::Stmt::Expr(ast::StmtExpr { value, .. })) = statements.clone().next()
224        && string_literal_expr_value(value).is_some()
225    {
226        statements.next();
227    }
228    for statement in statements {
229        match statement {
230            ast::Stmt::ImportFrom(ast::StmtImportFrom {
231                module,
232                names,
233                level,
234                ..
235            }) if *level == 0 && module.as_ref().map(|id| id.as_str()) == Some("__future__") => {
236                for alias in names {
237                    let future_feature =
238                        alias
239                            .name
240                            .as_str()
241                            .try_into()
242                            .map_err(|name| FutureFeatureError {
243                                features: future_features,
244                                range: alias.range,
245                                kind: FutureFeatureErrorKind::InvalidFeature(name),
246                            })?;
247
248                    match future_feature {
249                        FutureFeature::Braces => {
250                            return Err(FutureFeatureError {
251                                features: future_features,
252                                range: alias.range,
253                                kind: FutureFeatureErrorKind::InvalidBraces,
254                            });
255                        }
256                        FutureFeature::Annotations => {
257                            future_features.insert(bytecode::CodeFlags::FUTURE_ANNOTATIONS)
258                        }
259                        FutureFeature::BarryAsFLUFL => {
260                            future_features.insert(bytecode::CodeFlags::FUTURE_BARRY_AS_BDFL)
261                        }
262                        FutureFeature::AbsoluteImport
263                        | FutureFeature::Division
264                        | FutureFeature::GeneratorStop
265                        | FutureFeature::Generators
266                        | FutureFeature::NestedScopes
267                        | FutureFeature::PrintFunction
268                        | FutureFeature::UnicodeLiterals
269                        | FutureFeature::WithStatement => {
270                            // Python 3 features. They are already implemented by default.
271                        }
272                    }
273                }
274            }
275            _ => return Ok(future_features),
276        }
277    }
278    Ok(future_features)
279}
280
281pub fn preprocess_statements(
282    body: &mut [ast::Stmt],
283    optimize: u8,
284    future_annotations: bool,
285    syntax_check_only: bool,
286) {
287    let preprocessor = AstPreprocessor {
288        optimize,
289        future_annotations,
290        constant_folding: !syntax_check_only,
291    };
292    for stmt in body {
293        preprocessor.visit_stmt(stmt);
294    }
295}
296
297pub fn preprocess_mod(
298    module: &mut ast::Mod,
299    optimize: u8,
300    future_annotations: bool,
301    syntax_check_only: bool,
302) {
303    let preprocessor = AstPreprocessor {
304        optimize,
305        future_annotations,
306        constant_folding: !syntax_check_only,
307    };
308    match module {
309        ast::Mod::Module(module) => preprocessor.visit_astfold_body(&mut module.body),
310        ast::Mod::Expression(expr) => preprocessor.visit_expr(&mut expr.body),
311    }
312}
313
314#[derive(Clone, Copy, Debug, Eq, PartialEq)]
315struct AstPreprocessor {
316    optimize: u8,
317    future_annotations: bool,
318    constant_folding: bool,
319}
320
321impl AstPreprocessor {
322    fn visit_astfold_body(self, body: &mut ast::Suite) {
323        let mut docstring = body_starts_with_docstring(body);
324        if docstring && self.optimize >= 2 {
325            remove_docstring_from_body(body);
326            docstring = false;
327        }
328
329        for stmt in body.iter_mut() {
330            self.visit_stmt(stmt);
331        }
332
333        if !docstring && body_starts_with_docstring(body) {
334            wrap_first_docstring_as_fstring(body);
335        }
336    }
337}
338
339impl Transformer for AstPreprocessor {
340    fn visit_stmt(&self, stmt: &mut ast::Stmt) {
341        match stmt {
342            ast::Stmt::FunctionDef(function) => {
343                if let Some(type_params) = &mut function.type_params {
344                    self.visit_type_params(type_params);
345                }
346                self.visit_parameters(&mut function.parameters);
347                self.visit_astfold_body(&mut function.body);
348                for decorator in &mut function.decorator_list {
349                    self.visit_decorator(decorator);
350                }
351                if let Some(returns) = &mut function.returns {
352                    self.visit_annotation(returns);
353                }
354            }
355            ast::Stmt::ClassDef(class) => {
356                if let Some(type_params) = &mut class.type_params {
357                    self.visit_type_params(type_params);
358                }
359                if let Some(arguments) = &mut class.arguments {
360                    self.visit_arguments(arguments);
361                }
362                self.visit_astfold_body(&mut class.body);
363                for decorator in &mut class.decorator_list {
364                    self.visit_decorator(decorator);
365                }
366            }
367            _ => transformer::walk_stmt(self, stmt),
368        }
369    }
370
371    fn visit_annotation(&self, expr: &mut Expr) {
372        if !self.future_annotations {
373            transformer::walk_annotation(self, expr);
374        }
375    }
376
377    fn visit_pattern(&self, pattern: &mut ast::Pattern) {
378        transformer::walk_pattern(self, pattern);
379        if !self.constant_folding {
380            return;
381        }
382        match pattern {
383            ast::Pattern::MatchValue(value) => fold_match_value_constant_expr(&mut value.value),
384            ast::Pattern::MatchMapping(mapping) => {
385                for key in &mut mapping.keys {
386                    fold_match_value_constant_expr(key);
387                }
388            }
389            _ => {}
390        }
391    }
392
393    fn visit_expr(&self, expr: &mut Expr) {
394        transformer::walk_expr(self, expr);
395        widen_implicit_call_generator_range(expr);
396        if self.constant_folding {
397            if let Some(optimized) = optimize_format(expr) {
398                *expr = optimized;
399            } else if let Some(optimized) = fold_debug_constant(expr, self.optimize) {
400                *expr = optimized;
401            }
402        }
403    }
404}
405
406/// Give a generator expression written straight into a call's parentheses the
407/// range of those parentheses.
408///
409/// `genexp` is a grammar rule of its own that consumes the parentheses it is
410/// written in, so every position taken from the node covers them. The parser
411/// here leaves the node spanning only the element through the last iterable.
412fn widen_implicit_call_generator_range(expr: &mut Expr) {
413    let Expr::Call(call) = expr else {
414        return;
415    };
416    let [Expr::Generator(generator)] = &mut *call.arguments.args else {
417        return;
418    };
419    if !generator.parenthesized {
420        generator.range = call.arguments.range;
421    }
422}
423
424fn fold_debug_constant(expr: &Expr, optimize: u8) -> Option<Expr> {
425    let Expr::Name(name) = expr else {
426        return None;
427    };
428    if !matches!(name.ctx, ast::ExprContext::Load) || name.id.as_str() != "__debug__" {
429        return None;
430    }
431
432    Some(Expr::BooleanLiteral(ast::ExprBooleanLiteral {
433        node_index: name.node_index.clone(),
434        range: name.range,
435        value: optimize == 0,
436    }))
437}
438
439fn optimize_format(expr: &Expr) -> Option<Expr> {
440    let Expr::BinOp(binop) = expr else {
441        return None;
442    };
443    if !matches!(binop.op, Operator::Mod) {
444        return None;
445    }
446    let (format, _) = string_literal_expr_value(&binop.left)?;
447    let Expr::Tuple(tuple) = binop.right.as_ref() else {
448        return None;
449    };
450    if tuple
451        .elts
452        .iter()
453        .any(|expr| matches!(expr, Expr::Starred(_)))
454    {
455        return None;
456    }
457
458    let elements = parse_format(format, &tuple.elts)?;
459    Some(Expr::FString(ExprFString {
460        node_index: binop.node_index.clone(),
461        range: binop.range,
462        value: FStringValue::single(FString {
463            range: binop.range,
464            node_index: binop.node_index.clone(),
465            elements: InterpolatedStringElements::from(elements),
466            flags: FStringFlags::empty(),
467        }),
468        runtime_joined_str: None,
469        runtime_values: None,
470    }))
471}
472
473fn parse_format(format: &str, args: &[Expr]) -> Option<Vec<InterpolatedStringElement>> {
474    let chars: Vec<char> = format.chars().collect();
475    let mut elements = Vec::with_capacity(args.len().saturating_mul(2).saturating_add(1));
476    let mut pos = 0;
477    let mut arg_idx = 0;
478
479    loop {
480        if let Some(literal) = parse_literal(&chars, &mut pos) {
481            elements.push(literal.into());
482        }
483        if pos >= chars.len() {
484            break;
485        }
486        if arg_idx >= args.len() {
487            return None;
488        }
489        debug_assert_eq!(chars[pos], '%');
490        pos += 1;
491        let formatted = parse_format_arg(&chars, &mut pos, args[arg_idx].clone())?;
492        elements.push(formatted.into());
493        arg_idx += 1;
494    }
495
496    (arg_idx == args.len()).then_some(elements)
497}
498
499fn parse_literal(chars: &[char], pos: &mut usize) -> Option<InterpolatedStringLiteralElement> {
500    let start = *pos;
501    let mut has_percents = false;
502    while *pos < chars.len() {
503        if chars[*pos] != '%' {
504            *pos += 1;
505        } else if *pos + 1 < chars.len() && chars[*pos + 1] == '%' {
506            has_percents = true;
507            *pos += 2;
508        } else {
509            break;
510        }
511    }
512    if *pos == start {
513        return None;
514    }
515
516    let mut value = String::new();
517    let mut i = start;
518    while i < *pos {
519        if has_percents && chars[i] == '%' && i + 1 < *pos && chars[i + 1] == '%' {
520            value.push('%');
521            i += 2;
522        } else {
523            value.push(chars[i]);
524            i += 1;
525        }
526    }
527
528    Some(generated_literal(value))
529}
530
531fn parse_format_arg(chars: &[char], pos: &mut usize, arg: Expr) -> Option<InterpolatedElement> {
532    let (spec, flags, width, precision) = simple_format_arg_parse(chars, pos)?;
533    let conversion = match spec {
534        's' => ConversionFlag::Str,
535        'r' => ConversionFlag::Repr,
536        'a' => ConversionFlag::Ascii,
537        _ => return None,
538    };
539
540    let mut format_spec = String::new();
541    if flags & F_LJUST == 0
542        && let Some(width) = width
543        && width > 0
544    {
545        format_spec.push('>');
546    }
547    if let Some(width) = width {
548        format_spec.push_str(&width.to_string());
549    }
550    if let Some(precision) = precision {
551        format_spec.push('.');
552        format_spec.push_str(&precision.to_string());
553    }
554
555    let range = arg.range();
556    let format_spec = (!format_spec.is_empty()).then(|| {
557        Box::new(InterpolatedStringFormatSpec {
558            range: TextRange::default(),
559            node_index: AtomicNodeIndex::NONE,
560            elements: InterpolatedStringElements::from(vec![generated_literal(format_spec).into()]),
561        })
562    });
563
564    Some(InterpolatedElement {
565        range,
566        node_index: arg.node_index().clone(),
567        expression: Box::new(arg),
568        debug_text: None,
569        conversion,
570        format_spec,
571        runtime_str: None,
572        runtime_interpolation_format_spec: None,
573        runtime_formatted_value_format_spec: None,
574    })
575}
576
577fn simple_format_arg_parse(
578    chars: &[char],
579    pos: &mut usize,
580) -> Option<(char, u8, Option<u16>, Option<u16>)> {
581    let mut flags = 0;
582    let mut ch = next_char(chars, pos)?;
583    loop {
584        match ch {
585            '-' => flags |= F_LJUST,
586            '+' | ' ' | '#' | '0' => {}
587            _ => break,
588        }
589        ch = next_char(chars, pos)?;
590    }
591
592    let width = parse_digits(chars, pos, &mut ch)?;
593    let precision = if ch == '.' {
594        ch = next_char(chars, pos)?;
595        Some(parse_digits(chars, pos, &mut ch)?.unwrap_or(0))
596    } else {
597        None
598    };
599
600    Some((ch, flags, width, precision))
601}
602
603fn parse_digits(chars: &[char], pos: &mut usize, ch: &mut char) -> Option<Option<u16>> {
604    if !ch.is_ascii_digit() {
605        return Some(None);
606    }
607
608    let mut value = 0u16;
609    let mut digits = 0usize;
610    while ch.is_ascii_digit() {
611        value = value * 10 + (*ch as u16 - b'0' as u16);
612        *ch = next_char(chars, pos)?;
613        digits += 1;
614        if digits >= MAXDIGITS {
615            return None;
616        }
617    }
618    Some(Some(value))
619}
620
621fn next_char(chars: &[char], pos: &mut usize) -> Option<char> {
622    let ch = chars.get(*pos).copied()?;
623    *pos += 1;
624    Some(ch)
625}
626
627fn generated_literal(value: String) -> InterpolatedStringLiteralElement {
628    InterpolatedStringLiteralElement {
629        range: TextRange::default(),
630        node_index: AtomicNodeIndex::NONE,
631        value: value.into_boxed_str(),
632    }
633}
634
635fn remove_docstring_from_body(body: &mut ast::Suite) {
636    if let Some(range) = take_docstring(body) {
637        if !body.is_empty() {
638            return;
639        }
640        let start = range.start();
641        let pass_range = TextRange::new(start, start + ruff_text_size::TextSize::from(4));
642        body.push(ast::Stmt::Pass(ast::StmtPass {
643            node_index: Default::default(),
644            range: pass_range,
645        }));
646    }
647}
648
649fn take_docstring(body: &mut ast::Suite) -> Option<TextRange> {
650    let ast::Stmt::Expr(expr_stmt) = body.first()? else {
651        return None;
652    };
653    if let Some((_, range)) = string_literal_expr_value(&expr_stmt.value) {
654        body.remove(0);
655        return Some(range);
656    }
657    None
658}
659
660fn body_starts_with_docstring(body: &[ast::Stmt]) -> bool {
661    let Some(ast::Stmt::Expr(expr_stmt)) = body.first() else {
662        return false;
663    };
664    string_literal_expr_value(&expr_stmt.value).is_some()
665}
666
667fn wrap_first_docstring_as_fstring(body: &mut [ast::Stmt]) {
668    let Some(ast::Stmt::Expr(expr_stmt)) = body.first_mut() else {
669        return;
670    };
671    let Some((value, range)) = string_literal_expr_value(&expr_stmt.value) else {
672        return;
673    };
674    let value = value.to_string();
675    *expr_stmt.value = ast::Expr::FString(ast::ExprFString {
676        node_index: AtomicNodeIndex::NONE,
677        range,
678        value: FStringValue::single(FString {
679            range,
680            node_index: AtomicNodeIndex::NONE,
681            elements: InterpolatedStringElements::from(vec![InterpolatedStringElement::Literal(
682                InterpolatedStringLiteralElement {
683                    range,
684                    node_index: AtomicNodeIndex::NONE,
685                    value: value.into_boxed_str(),
686                },
687            )]),
688            flags: FStringFlags::empty(),
689        }),
690        runtime_joined_str: None,
691        runtime_values: None,
692    });
693}
694
695fn string_literal_expr_value(expr: &Expr) -> Option<(&str, TextRange)> {
696    match expr {
697        Expr::StringLiteral(string) => Some((string.value.to_str(), expr.range())),
698        Expr::Constant(ast::ExprConstant {
699            value: ast::ConstantValue::Str(value),
700            ..
701        }) => Some((value.as_ref(), expr.range())),
702        _ => None,
703    }
704}
705
706fn fold_match_value_constant_expr(expr: &mut ast::Expr) {
707    match expr {
708        ast::Expr::UnaryOp(unary)
709            if matches!(unary.op, ast::UnaryOp::USub)
710                && matches!(unary.operand.as_ref(), ast::Expr::NumberLiteral(_)) =>
711        {
712            if let Some(number) = negate_match_number(&unary.operand) {
713                *expr = ast::Expr::NumberLiteral(ast::ExprNumberLiteral {
714                    node_index: unary.node_index.clone(),
715                    range: unary.range,
716                    value: number,
717                });
718            }
719        }
720        ast::Expr::BinOp(binop) if matches!(binop.op, ast::Operator::Add | ast::Operator::Sub) => {
721            fold_match_value_constant_expr(&mut binop.left);
722            if let Some(number) = fold_match_number_binop(&binop.left, binop.op, &binop.right) {
723                *expr = ast::Expr::NumberLiteral(ast::ExprNumberLiteral {
724                    node_index: binop.node_index.clone(),
725                    range: binop.range,
726                    value: number,
727                });
728            }
729        }
730        _ => {}
731    }
732}
733
734fn negate_match_number(expr: &ast::Expr) -> Option<ast::Number> {
735    let ast::Expr::NumberLiteral(number) = expr else {
736        return None;
737    };
738    Some(match &number.value {
739        ast::Number::Int(value) => {
740            if *value == ast::Int::ZERO {
741                ast::Number::Int(ast::Int::ZERO)
742            } else {
743                return None;
744            }
745        }
746        ast::Number::Float(value) => ast::Number::Float(-value),
747        ast::Number::Complex { real, imag } => ast::Number::Complex {
748            real: -real,
749            imag: -imag,
750        },
751    })
752}
753
754fn fold_match_number_binop(
755    left: &ast::Expr,
756    op: ast::Operator,
757    right: &ast::Expr,
758) -> Option<ast::Number> {
759    let ast::Expr::NumberLiteral(left) = left else {
760        return None;
761    };
762    let ast::Expr::NumberLiteral(right) = right else {
763        return None;
764    };
765    let right = match right.value {
766        ast::Number::Complex { real, imag } => (real, imag),
767        _ => return None,
768    };
769    enum MatchNumberLeft {
770        Real(f64),
771        Complex { real: f64, imag: f64 },
772    }
773    let left = match &left.value {
774        ast::Number::Int(value) => MatchNumberLeft::Real(value.as_i64()? as f64),
775        ast::Number::Float(value) => MatchNumberLeft::Real(*value),
776        ast::Number::Complex { real, imag } => MatchNumberLeft::Complex {
777            real: *real,
778            imag: *imag,
779        },
780    };
781    let (real, imag) = match (left, op) {
782        (MatchNumberLeft::Real(left), ast::Operator::Add) => (left + right.0, right.1),
783        (MatchNumberLeft::Real(left), ast::Operator::Sub) => (left - right.0, -right.1),
784        (MatchNumberLeft::Complex { real, imag }, ast::Operator::Add) => {
785            (real + right.0, imag + right.1)
786        }
787        (MatchNumberLeft::Complex { real, imag }, ast::Operator::Sub) => {
788            (real - right.0, imag - right.1)
789        }
790        _ => return None,
791    };
792    Some(ast::Number::Complex { real, imag })
793}
794
795#[cfg(test)]
796mod tests {
797    use super::*;
798
799    fn first_match_value(source: &str) -> ast::Expr {
800        let parsed = ruff_python_parser::parse(source, ruff_python_parser::Mode::Module.into())
801            .unwrap()
802            .into_syntax();
803        let mut module = parsed;
804        let future_annotations = has_future_annotations(&module);
805        preprocess_mod(&mut module, 0, future_annotations, false);
806        let ast::Mod::Module(module) = module else {
807            panic!("expected module");
808        };
809        let [ast::Stmt::Match(match_stmt)] = &module.body[..] else {
810            panic!("expected a single match statement");
811        };
812        let ast::Pattern::MatchValue(value) = &match_stmt.cases[0].pattern else {
813            panic!("expected a value pattern");
814        };
815        *value.value.clone()
816    }
817
818    fn preprocess_source(source: &str) -> ast::Mod {
819        let mut module = ruff_python_parser::parse(source, ruff_python_parser::Mode::Module.into())
820            .unwrap()
821            .into_syntax();
822        let future_annotations = has_future_annotations(&module);
823        preprocess_mod(&mut module, 0, future_annotations, false);
824        module
825    }
826
827    fn preprocess_source_with_optimize(source: &str, optimize: u8) -> ast::Mod {
828        let mut module = ruff_python_parser::parse(source, ruff_python_parser::Mode::Module.into())
829            .unwrap()
830            .into_syntax();
831        let future_annotations = has_future_annotations(&module);
832        preprocess_mod(&mut module, optimize, future_annotations, false);
833        module
834    }
835
836    fn preprocess_source_syntax_check_only(source: &str, optimize: u8) -> ast::Mod {
837        let mut module = ruff_python_parser::parse(source, ruff_python_parser::Mode::Module.into())
838            .unwrap()
839            .into_syntax();
840        let future_annotations = has_future_annotations(&module);
841        preprocess_mod(&mut module, optimize, future_annotations, true);
842        module
843    }
844
845    #[test]
846    fn folds_match_value_negative_float_in_preprocess() {
847        let value = first_match_value(
848            "\
849match value:
850    case -1.5:
851        pass
852",
853        );
854        let ast::Expr::NumberLiteral(number) = value else {
855            panic!("expected folded number literal, got {value:?}");
856        };
857        assert!(matches!(number.value, ast::Number::Float(value) if value == -1.5));
858    }
859
860    #[test]
861    fn folds_match_value_complex_binop_in_preprocess() {
862        let value = first_match_value(
863            "\
864match value:
865    case 1 + 2j:
866        pass
867",
868        );
869        let ast::Expr::NumberLiteral(number) = value else {
870            panic!("expected folded number literal, got {value:?}");
871        };
872        assert!(
873            matches!(number.value, ast::Number::Complex { real, imag } if real == 1.0 && imag == 2.0)
874        );
875    }
876
877    #[test]
878    fn folds_match_value_complex_complex_binop_in_preprocess() {
879        let left = ast::Expr::NumberLiteral(ast::ExprNumberLiteral {
880            node_index: AtomicNodeIndex::NONE,
881            range: TextRange::default(),
882            value: ast::Number::Complex {
883                real: 0.0,
884                imag: 1.0,
885            },
886        });
887        let right = ast::Expr::NumberLiteral(ast::ExprNumberLiteral {
888            node_index: AtomicNodeIndex::NONE,
889            range: TextRange::default(),
890            value: ast::Number::Complex {
891                real: 0.0,
892                imag: 2.0,
893            },
894        });
895        let number = fold_match_number_binop(&left, ast::Operator::Add, &right)
896            .expect("CPython fold_const_match_patterns() uses PyNumber_Add");
897        assert!(
898            matches!(number, ast::Number::Complex { real, imag } if real == 0.0 && imag == 3.0)
899        );
900    }
901
902    #[test]
903    fn folds_match_value_real_minus_zero_complex_preserves_negative_zero_in_preprocess() {
904        let value = first_match_value(
905            "\
906match value:
907    case 0 - 0j:
908        pass
909",
910        );
911        let ast::Expr::NumberLiteral(number) = value else {
912            panic!("expected folded number literal, got {value:?}");
913        };
914        assert!(matches!(number.value, ast::Number::Complex { real, imag }
915                if real == 0.0 && imag == 0.0 && imag.is_sign_negative()));
916    }
917
918    #[test]
919    fn future_annotations_skip_annotation_preprocess_like_cpython() {
920        let module = preprocess_source(
921            "\
922from __future__ import annotations
923def f(x: __debug__) -> __debug__:
924    pass
925y: __debug__
926z = __debug__
927",
928        );
929        let ast::Mod::Module(module) = module else {
930            panic!("expected module");
931        };
932        let ast::Stmt::FunctionDef(function) = &module.body[1] else {
933            panic!("expected function");
934        };
935        let annotation = function.parameters.args[0]
936            .parameter
937            .annotation
938            .as_deref()
939            .expect("missing parameter annotation");
940        assert!(
941            matches!(annotation, ast::Expr::Name(name) if name.id.as_str() == "__debug__"),
942            "future annotations should skip parameter annotation folding, got {annotation:?}"
943        );
944        let returns = function
945            .returns
946            .as_deref()
947            .expect("missing return annotation");
948        assert!(
949            matches!(returns, ast::Expr::Name(name) if name.id.as_str() == "__debug__"),
950            "future annotations should skip return annotation folding, got {returns:?}"
951        );
952        let ast::Stmt::AnnAssign(ann_assign) = &module.body[2] else {
953            panic!("expected annotated assignment");
954        };
955        assert!(
956            matches!(ann_assign.annotation.as_ref(), ast::Expr::Name(name) if name.id.as_str() == "__debug__"),
957            "future annotations should skip annotated assignment annotation folding, got {:?}",
958            ann_assign.annotation
959        );
960        let ast::Stmt::Assign(assign) = &module.body[3] else {
961            panic!("expected assignment");
962        };
963        assert!(
964            matches!(assign.value.as_ref(), ast::Expr::BooleanLiteral(boolean) if boolean.value),
965            "non-annotation expression should still fold __debug__, got {:?}",
966            assign.value
967        );
968    }
969
970    #[test]
971    fn late_future_annotations_do_not_affect_preprocess_like_cpython() {
972        let module = preprocess_source(
973            "\
974x = 1
975from __future__ import annotations
976y: __debug__
977",
978        );
979        let ast::Mod::Module(module) = module else {
980            panic!("expected module");
981        };
982        let ast::Stmt::AnnAssign(ann_assign) = &module.body[2] else {
983            panic!("expected annotated assignment");
984        };
985        assert!(
986            matches!(ann_assign.annotation.as_ref(), ast::Expr::BooleanLiteral(boolean) if boolean.value),
987            "late future import should not disable annotation folding, got {:?}",
988            ann_assign.annotation
989        );
990    }
991
992    #[test]
993    fn optimize_two_wraps_new_docstring_after_removing_original() {
994        let module = preprocess_source_with_optimize("\"first\"\n\"second\"\n", 2);
995        let ast::Mod::Module(module) = module else {
996            panic!("expected module");
997        };
998        let [ast::Stmt::Expr(expr)] = &module.body[..] else {
999            panic!("expected only the second statement to remain");
1000        };
1001        assert!(
1002            matches!(expr.value.as_ref(), ast::Expr::FString(_)),
1003            "CPython wraps the new leading string as JoinedStr so it is not a docstring"
1004        );
1005    }
1006
1007    #[test]
1008    fn syntax_check_only_disables_constant_folding_but_keeps_docstring_strip() {
1009        let module = preprocess_source_syntax_check_only("\"doc\"\nvalue = __debug__\n", 2);
1010        let ast::Mod::Module(module) = module else {
1011            panic!("expected module");
1012        };
1013        assert!(
1014            matches!(module.body[0], ast::Stmt::Assign(_)),
1015            "optimize=2 should still strip docstrings in syntax_check_only mode"
1016        );
1017        let ast::Stmt::Assign(assign) = &module.body[0] else {
1018            panic!("expected assignment");
1019        };
1020        assert!(
1021            matches!(assign.value.as_ref(), ast::Expr::Name(name) if name.id.as_str() == "__debug__"),
1022            "syntax_check_only should skip __debug__ folding, got {:?}",
1023            assign.value
1024        );
1025    }
1026}