Skip to main content

portalis_transpiler/
reference_optimizer.rs

1//! Reference Optimization
2//!
3//! Optimizes generated Rust code to use references efficiently, avoiding
4//! unnecessary clones and memory allocations while maintaining safety.
5
6use std::collections::{HashMap, HashSet};
7
8/// Reference usage pattern
9#[derive(Debug, Clone, PartialEq, Eq)]
10pub enum ReferencePattern {
11    /// Value is only read, can use &T
12    ReadOnly,
13    /// Value is modified, needs &mut T
14    Mutable,
15    /// Value is moved/consumed, needs T
16    Owned,
17    /// Value is shared across threads, needs Arc<T>
18    Shared,
19    /// Value is mutated across threads, needs Arc<Mutex<T>>
20    SharedMutable,
21}
22
23/// Optimization suggestion
24#[derive(Debug, Clone)]
25pub struct Optimization {
26    pub location: String,
27    pub pattern: OptimizationPattern,
28    pub original: String,
29    pub optimized: String,
30    pub explanation: String,
31    pub safety_note: Option<String>,
32}
33
34/// Type of optimization
35#[derive(Debug, Clone, PartialEq, Eq, Hash)]
36pub enum OptimizationPattern {
37    /// Remove unnecessary clone()
38    CloneElimination,
39    /// Use &str instead of String
40    StringSlice,
41    /// Use iterator instead of collecting
42    IteratorChain,
43    /// Use reference in function parameter
44    BorrowParameter,
45    /// Use slice instead of Vec
46    SliceParameter,
47    /// Return reference instead of owned value
48    ReturnReference,
49    /// Use Cow for conditional ownership
50    CowOptimization,
51    /// Avoid intermediate allocations
52    AllocationReduction,
53    /// Use smart pointer efficiently
54    SmartPointer,
55}
56
57/// Variable usage tracking
58#[derive(Debug, Clone)]
59struct VariableUsage {
60    name: String,
61    reads: usize,
62    writes: usize,
63    moved: bool,
64    borrowed: bool,
65    borrowed_mut: bool,
66    escapes_scope: bool,
67}
68
69impl VariableUsage {
70    fn new(name: impl Into<String>) -> Self {
71        Self {
72            name: name.into(),
73            reads: 0,
74            writes: 0,
75            moved: false,
76            borrowed: false,
77            borrowed_mut: false,
78            escapes_scope: false,
79        }
80    }
81
82    fn infer_pattern(&self) -> ReferencePattern {
83        if self.moved || self.escapes_scope {
84            ReferencePattern::Owned
85        } else if self.writes > 0 || self.borrowed_mut {
86            ReferencePattern::Mutable
87        } else {
88            ReferencePattern::ReadOnly
89        }
90    }
91}
92
93/// Reference optimizer
94pub struct ReferenceOptimizer {
95    /// Variable usage tracking
96    usage: HashMap<String, VariableUsage>,
97    /// Detected optimizations
98    optimizations: Vec<Optimization>,
99    /// String literal detection
100    string_literals: HashSet<String>,
101    /// Function return types
102    return_types: HashMap<String, String>,
103}
104
105impl ReferenceOptimizer {
106    pub fn new() -> Self {
107        Self {
108            usage: HashMap::new(),
109            optimizations: Vec::new(),
110            string_literals: HashSet::new(),
111            return_types: HashMap::new(),
112        }
113    }
114
115    /// Track variable read
116    pub fn track_read(&mut self, var: &str) {
117        self.usage.entry(var.to_string())
118            .or_insert_with(|| VariableUsage::new(var))
119            .reads += 1;
120    }
121
122    /// Track variable write
123    pub fn track_write(&mut self, var: &str) {
124        self.usage.entry(var.to_string())
125            .or_insert_with(|| VariableUsage::new(var))
126            .writes += 1;
127    }
128
129    /// Track variable move
130    pub fn track_move(&mut self, var: &str) {
131        self.usage.entry(var.to_string())
132            .or_insert_with(|| VariableUsage::new(var))
133            .moved = true;
134    }
135
136    /// Track variable borrow
137    pub fn track_borrow(&mut self, var: &str, mutable: bool) {
138        let usage = self.usage.entry(var.to_string())
139            .or_insert_with(|| VariableUsage::new(var));
140
141        if mutable {
142            usage.borrowed_mut = true;
143        } else {
144            usage.borrowed = true;
145        }
146    }
147
148    /// Analyze function parameter and suggest optimization
149    pub fn optimize_parameter(&mut self, param_name: &str, param_type: &str, usage_in_body: &str) -> String {
150        // Check if parameter is only read
151        if !usage_in_body.contains(&format!("{} =", param_name))
152            && !usage_in_body.contains(&format!("&mut {}", param_name)) {
153
154            // Suggest reference for non-Copy types
155            if param_type == "String" {
156                self.add_optimization(
157                    param_name,
158                    OptimizationPattern::BorrowParameter,
159                    format!("{}: String", param_name),
160                    format!("{}: &str", param_name),
161                    "Parameter is only read, use &str to avoid unnecessary allocation".to_string(),
162                );
163                return format!("{}: &str", param_name);
164            } else if param_type.starts_with("Vec<") {
165                self.add_optimization(
166                    param_name,
167                    OptimizationPattern::SliceParameter,
168                    format!("{}: {}", param_name, param_type),
169                    format!("{}: &[{}]", param_name, param_type.trim_start_matches("Vec<").trim_end_matches('>')),
170                    "Parameter is only read, use slice to avoid cloning".to_string(),
171                );
172                let inner = param_type.trim_start_matches("Vec<").trim_end_matches('>');
173                return format!("{}: &[{}]", param_name, inner);
174            }
175        }
176
177        format!("{}: {}", param_name, param_type)
178    }
179
180    /// Optimize string usage
181    pub fn optimize_string(&mut self, var_name: &str, value: &str, is_literal: bool) -> String {
182        if is_literal {
183            self.string_literals.insert(var_name.to_string());
184
185            // String literal can use &str
186            self.add_optimization(
187                var_name,
188                OptimizationPattern::StringSlice,
189                format!("let {}: String = \"{}\".to_string()", var_name, value),
190                format!("let {}: &str = \"{}\"", var_name, value),
191                "String literal doesn't need allocation, use &str".to_string(),
192            );
193
194            format!("let {}: &str = \"{}\"", var_name, value)
195        } else {
196            format!("let {}: String = {}", var_name, value)
197        }
198    }
199
200    /// Detect and eliminate unnecessary clones
201    pub fn eliminate_clone(&mut self, var_name: &str, expr: &str) -> String {
202        // Pattern: variable.clone() when variable is last use
203        if expr.ends_with(".clone()") {
204            let base_var = expr.trim_end_matches(".clone()");
205
206            if let Some(usage) = self.usage.get(base_var) {
207                if usage.reads == 1 && !usage.borrowed && !usage.escapes_scope {
208                    self.add_optimization(
209                        var_name,
210                        OptimizationPattern::CloneElimination,
211                        expr.to_string(),
212                        base_var.to_string(),
213                        format!("Last use of {}, clone unnecessary", base_var),
214                    );
215                    return base_var.to_string();
216                }
217            }
218        }
219
220        expr.to_string()
221    }
222
223    /// Optimize iterator chains to avoid intermediate collections
224    pub fn optimize_iterator(&mut self, location: &str, code: &str) -> String {
225        // Pattern: .collect::<Vec<_>>().iter()
226        if code.contains(".collect::<Vec<_>>().iter()") {
227            let optimized = code.replace(".collect::<Vec<_>>().iter()", "");
228
229            self.add_optimization(
230                location,
231                OptimizationPattern::IteratorChain,
232                code.to_string(),
233                optimized.clone(),
234                "Avoid intermediate Vec allocation by chaining iterators".to_string(),
235            );
236
237            return optimized;
238        }
239
240        // Pattern: .collect() followed by .into_iter()
241        if code.contains(".collect::<Vec<_>>().into_iter()") {
242            let optimized = code.replace(".collect::<Vec<_>>().into_iter()", "");
243
244            self.add_optimization(
245                location,
246                OptimizationPattern::IteratorChain,
247                code.to_string(),
248                optimized.clone(),
249                "Chain iterators instead of collecting intermediate vector".to_string(),
250            );
251
252            return optimized;
253        }
254
255        code.to_string()
256    }
257
258    /// Optimize return type to use reference when possible
259    pub fn optimize_return(&mut self, func_name: &str, return_type: &str, returns_local: bool) -> String {
260        if returns_local {
261            // Can't return reference to local variable
262            return return_type.to_string();
263        }
264
265        // If returning String that comes from parameter, can return &str
266        if return_type == "String" {
267            self.add_optimization(
268                func_name,
269                OptimizationPattern::ReturnReference,
270                format!("-> {}", return_type),
271                "-> &str".to_string(),
272                "Return reference to avoid unnecessary allocation".to_string(),
273            );
274            // Note: This optimization requires parameter analysis
275            self.optimizations.last_mut().unwrap().safety_note = Some(
276                "Only valid if returned value has lifetime tied to parameter".to_string()
277            );
278        }
279
280        return_type.to_string()
281    }
282
283    /// Suggest Cow for conditional ownership
284    pub fn suggest_cow(&mut self, var_name: &str, sometimes_owned: bool, sometimes_borrowed: bool) -> String {
285        if sometimes_owned && sometimes_borrowed {
286            self.add_optimization(
287                var_name,
288                OptimizationPattern::CowOptimization,
289                "String".to_string(),
290                "Cow<'_, str>".to_string(),
291                "Use Cow for conditional ownership - avoids cloning when possible".to_string(),
292            );
293            "Cow<'_, str>".to_string()
294        } else if sometimes_borrowed {
295            "&str".to_string()
296        } else {
297            "String".to_string()
298        }
299    }
300
301    /// Optimize smart pointer usage
302    pub fn optimize_smart_pointer(&mut self, var_name: &str, is_shared: bool, is_mutable: bool, is_threadsafe: bool) -> String {
303        let suggested = if is_threadsafe {
304            if is_mutable {
305                "Arc<Mutex<T>>".to_string()
306            } else {
307                "Arc<T>".to_string()
308            }
309        } else if is_shared {
310            if is_mutable {
311                "RefCell<T>".to_string()
312            } else {
313                "Rc<T>".to_string()
314            }
315        } else if is_mutable {
316            "Box<T>".to_string()
317        } else {
318            "T".to_string()
319        };
320
321        self.add_optimization(
322            var_name,
323            OptimizationPattern::SmartPointer,
324            "Box<T>".to_string(),
325            suggested.clone(),
326            self.smart_pointer_rationale(is_shared, is_mutable, is_threadsafe),
327        );
328
329        suggested
330    }
331
332    fn smart_pointer_rationale(&self, shared: bool, mutable: bool, threadsafe: bool) -> String {
333        match (shared, mutable, threadsafe) {
334            (false, false, _) => "No sharing needed, use T directly".to_string(),
335            (false, true, _) => "Single owner with mutation, use Box<T>".to_string(),
336            (true, false, false) => "Shared ownership, immutable, use Rc<T>".to_string(),
337            (true, true, false) => "Shared ownership, mutable, use Rc<RefCell<T>>".to_string(),
338            (true, false, true) => "Thread-safe shared ownership, use Arc<T>".to_string(),
339            (true, true, true) => "Thread-safe shared mutation, use Arc<Mutex<T>>".to_string(),
340        }
341    }
342
343    /// Add optimization suggestion
344    fn add_optimization(
345        &mut self,
346        location: impl Into<String>,
347        pattern: OptimizationPattern,
348        original: impl Into<String>,
349        optimized: impl Into<String>,
350        explanation: impl Into<String>,
351    ) {
352        self.optimizations.push(Optimization {
353            location: location.into(),
354            pattern,
355            original: original.into(),
356            optimized: optimized.into(),
357            explanation: explanation.into(),
358            safety_note: None,
359        });
360    }
361
362    /// Get all optimizations found
363    pub fn get_optimizations(&self) -> &[Optimization] {
364        &self.optimizations
365    }
366
367    /// Generate optimization report
368    pub fn report(&self) -> String {
369        let mut report = String::new();
370        report.push_str("=== Reference Optimization Report ===\n\n");
371
372        if self.optimizations.is_empty() {
373            report.push_str("No optimizations found - code is already efficient!\n");
374            return report;
375        }
376
377        report.push_str(&format!("Found {} optimization opportunities:\n\n", self.optimizations.len()));
378
379        let mut by_pattern: HashMap<OptimizationPattern, Vec<&Optimization>> = HashMap::new();
380        for opt in &self.optimizations {
381            by_pattern.entry(opt.pattern.clone())
382                .or_default()
383                .push(opt);
384        }
385
386        for (pattern, opts) in by_pattern {
387            report.push_str(&format!("## {:?} ({} occurrences)\n\n", pattern, opts.len()));
388
389            for opt in opts {
390                report.push_str(&format!("Location: {}\n", opt.location));
391                report.push_str(&format!("  Original:  {}\n", opt.original));
392                report.push_str(&format!("  Optimized: {}\n", opt.optimized));
393                report.push_str(&format!("  Reason: {}\n", opt.explanation));
394                if let Some(note) = &opt.safety_note {
395                    report.push_str(&format!("  ⚠️  Safety: {}\n", note));
396                }
397                report.push('\n');
398            }
399        }
400
401        report
402    }
403
404    /// Clear all tracked data
405    pub fn clear(&mut self) {
406        self.usage.clear();
407        self.optimizations.clear();
408        self.string_literals.clear();
409        self.return_types.clear();
410    }
411}
412
413impl Default for ReferenceOptimizer {
414    fn default() -> Self {
415        Self::new()
416    }
417}
418
419/// Common optimization patterns
420pub struct OptimizationPatterns;
421
422impl OptimizationPatterns {
423    /// Example: Avoiding unnecessary clones
424    pub fn clone_elimination_example() -> (&'static str, &'static str) {
425        let before = r#"
426fn process(data: Vec<i32>) -> Vec<i32> {
427    let copy = data.clone();  // Unnecessary!
428    copy
429}
430"#;
431
432        let after = r#"
433fn process(data: Vec<i32>) -> Vec<i32> {
434    data  // No clone needed, data is consumed
435}
436"#;
437
438        (before, after)
439    }
440
441    /// Example: String slice optimization
442    pub fn string_slice_example() -> (&'static str, &'static str) {
443        let before = r#"
444fn greet(name: String) -> String {
445    format!("Hello, {}", name)
446}
447
448let msg = greet("Alice".to_string());  // Unnecessary allocation
449"#;
450
451        let after = r#"
452fn greet(name: &str) -> String {
453    format!("Hello, {}", name)
454}
455
456let msg = greet("Alice");  // No allocation needed
457"#;
458
459        (before, after)
460    }
461
462    /// Example: Iterator chain optimization
463    pub fn iterator_chain_example() -> (&'static str, &'static str) {
464        let before = r#"
465let result = data
466    .iter()
467    .map(|x| x * 2)
468    .collect::<Vec<_>>()  // Unnecessary intermediate Vec
469    .iter()
470    .filter(|x| **x > 10)
471    .collect();
472"#;
473
474        let after = r#"
475let result = data
476    .iter()
477    .map(|x| x * 2)
478    .filter(|x| *x > 10)  // Chain directly
479    .collect();
480"#;
481
482        (before, after)
483    }
484
485    /// Example: Slice parameter
486    pub fn slice_parameter_example() -> (&'static str, &'static str) {
487        let before = r#"
488fn sum(numbers: Vec<i32>) -> i32 {
489    numbers.iter().sum()
490}
491
492// Caller must clone or lose ownership
493let total = sum(vec.clone());
494"#;
495
496        let after = r#"
497fn sum(numbers: &[i32]) -> i32 {
498    numbers.iter().sum()
499}
500
501// Caller can pass reference
502let total = sum(&vec);
503"#;
504
505        (before, after)
506    }
507
508    /// Example: Cow optimization
509    pub fn cow_example() -> (&'static str, &'static str) {
510        let before = r#"
511fn process(s: String, uppercase: bool) -> String {
512    if uppercase {
513        s.to_uppercase()  // Must clone even for owned String
514    } else {
515        s
516    }
517}
518"#;
519
520        let after = r#"
521use std::borrow::Cow;
522
523fn process(s: &str, uppercase: bool) -> Cow<str> {
524    if uppercase {
525        Cow::Owned(s.to_uppercase())
526    } else {
527        Cow::Borrowed(s)
528    }
529}
530"#;
531
532        (before, after)
533    }
534
535    /// Example: Smart pointer selection
536    pub fn smart_pointer_example() -> &'static str {
537        r#"
538// Single ownership
539let data = Box::new(value);
540
541// Shared ownership (single thread)
542let data = Rc::new(value);
543let shared = Rc::clone(&data);
544
545// Shared ownership (multi-thread)
546let data = Arc::new(value);
547let shared = Arc::clone(&data);
548
549// Shared mutation (single thread)
550let data = Rc::new(RefCell::new(value));
551data.borrow_mut().update();
552
553// Shared mutation (multi-thread)
554let data = Arc::new(Mutex::new(value));
555data.lock().unwrap().update();
556"#
557    }
558}
559
560#[cfg(test)]
561mod tests {
562    use super::*;
563
564    #[test]
565    fn test_parameter_optimization() {
566        let mut optimizer = ReferenceOptimizer::new();
567
568        let optimized = optimizer.optimize_parameter("name", "String", "println!(\"{}\" name)");
569        assert_eq!(optimized, "name: &str");
570
571        assert_eq!(optimizer.optimizations.len(), 1);
572        assert_eq!(optimizer.optimizations[0].pattern, OptimizationPattern::BorrowParameter);
573    }
574
575    #[test]
576    fn test_clone_elimination() {
577        let mut optimizer = ReferenceOptimizer::new();
578
579        optimizer.usage.insert("data".to_string(), {
580            let mut usage = VariableUsage::new("data");
581            usage.reads = 1;
582            usage
583        });
584
585        let result = optimizer.eliminate_clone("result", "data.clone()");
586        assert_eq!(result, "data");
587        assert_eq!(optimizer.optimizations.len(), 1);
588    }
589
590    #[test]
591    fn test_iterator_optimization() {
592        let mut optimizer = ReferenceOptimizer::new();
593
594        let code = "data.iter().map(|x| x * 2).collect::<Vec<_>>().iter()";
595        let optimized = optimizer.optimize_iterator("chain", code);
596
597        assert!(!optimized.contains(".collect::<Vec<_>>().iter()"));
598        assert_eq!(optimizer.optimizations.len(), 1);
599        assert_eq!(optimizer.optimizations[0].pattern, OptimizationPattern::IteratorChain);
600    }
601
602    #[test]
603    fn test_smart_pointer_selection() {
604        let mut optimizer = ReferenceOptimizer::new();
605
606        let result = optimizer.optimize_smart_pointer("data", true, true, true);
607        assert_eq!(result, "Arc<Mutex<T>>");
608
609        let result = optimizer.optimize_smart_pointer("data", true, false, false);
610        assert_eq!(result, "Rc<T>");
611    }
612}