about summary refs log tree commit diff
path: root/compiler/rustc_smir/src/rustc_internal
diff options
context:
space:
mode:
authorCelina G. Val <celinval@amazon.com>2023-09-01 21:12:46 -0700
committerCelina G. Val <celinval@amazon.com>2023-09-05 09:19:56 -0700
commitd10d8290ac50b5ea8c0ca72bf216cf414b2ada24 (patch)
treea59b79c4369094fc43c707e79789de4c2ebfb1ff /compiler/rustc_smir/src/rustc_internal
parent1a8a5d0a296f24b5184a3997e6dbedb2d31c079a (diff)
Add tests and use ControlFlow
Diffstat (limited to 'compiler/rustc_smir/src/rustc_internal')
-rw-r--r--compiler/rustc_smir/src/rustc_internal/mod.rs59
1 files changed, 27 insertions, 32 deletions
diff --git a/compiler/rustc_smir/src/rustc_internal/mod.rs b/compiler/rustc_smir/src/rustc_internal/mod.rs
index 26127c5eb85..73b8e060dd8 100644
--- a/compiler/rustc_smir/src/rustc_internal/mod.rs
+++ b/compiler/rustc_smir/src/rustc_internal/mod.rs
@@ -4,7 +4,7 @@
 //! until stable MIR is complete.
 
 use std::fmt::Debug;
-use std::ops::Index;
+use std::ops::{ControlFlow, Index};
 
 use crate::rustc_internal;
 use crate::stable_mir::CompilerError;
@@ -190,52 +190,44 @@ pub(crate) fn opaque<T: Debug>(value: &T) -> Opaque {
     Opaque(format!("{value:?}"))
 }
 
-pub struct StableMir<T: Send>
+pub struct StableMir<B = (), C = ()>
 where
-    T: Send,
+    B: Send,
+    C: Send,
 {
     args: Vec<String>,
-    callback: fn(TyCtxt<'_>) -> T,
-    after_analysis: Compilation,
-    result: Option<T>,
+    callback: fn(TyCtxt<'_>) -> ControlFlow<B, C>,
+    result: Option<ControlFlow<B, C>>,
 }
 
-impl<T> StableMir<T>
+impl<B, C> StableMir<B, C>
 where
-    T: Send,
+    B: Send,
+    C: Send,
 {
     /// Creates a new `StableMir` instance, with given test_function and arguments.
-    pub fn new(args: Vec<String>, callback: fn(TyCtxt<'_>) -> T) -> Self {
-        StableMir { args, callback, result: None, after_analysis: Compilation::Stop }
-    }
-
-    /// Configure object to stop compilation after callback is called.
-    pub fn stop_compilation(&mut self) -> &mut Self {
-        self.after_analysis = Compilation::Stop;
-        self
-    }
-
-    /// Configure object to continue compilation after callback is called.
-    pub fn continue_compilation(&mut self) -> &mut Self {
-        self.after_analysis = Compilation::Continue;
-        self
+    pub fn new(args: Vec<String>, callback: fn(TyCtxt<'_>) -> ControlFlow<B, C>) -> Self {
+        StableMir { args, callback, result: None }
     }
 
     /// Runs the compiler against given target and tests it with `test_function`
-    pub fn run(&mut self) -> Result<T, CompilerError> {
+    pub fn run(&mut self) -> Result<C, CompilerError<B>> {
         let compiler_result =
             rustc_driver::catch_fatal_errors(|| RunCompiler::new(&self.args.clone(), self).run());
-        match compiler_result {
-            Ok(Ok(())) => Ok(self.result.take().unwrap()),
-            Ok(Err(_)) => Err(CompilerError::CompilationFailed),
-            Err(_) => Err(CompilerError::ICE),
+        match (compiler_result, self.result.take()) {
+            (Ok(Ok(())), Some(ControlFlow::Continue(value))) => Ok(value),
+            (Ok(Ok(())), Some(ControlFlow::Break(value))) => Err(CompilerError::Interrupted(value)),
+            (Ok(Ok(_)), None) => Err(CompilerError::Skipped),
+            (Ok(Err(_)), _) => Err(CompilerError::CompilationFailed),
+            (Err(_), _) => Err(CompilerError::ICE),
         }
     }
 }
 
-impl<T> Callbacks for StableMir<T>
+impl<B, C> Callbacks for StableMir<B, C>
 where
-    T: Send,
+    B: Send,
+    C: Send,
 {
     /// Called after analysis. Return value instructs the compiler whether to
     /// continue the compilation afterwards (defaults to `Compilation::Continue`)
@@ -249,8 +241,11 @@ where
             rustc_internal::run(tcx, || {
                 self.result = Some((self.callback)(tcx));
             });
-        });
-        // Let users define if they want to stop compilation.
-        self.after_analysis
+            if self.result.as_ref().is_some_and(|val| val.is_continue()) {
+                Compilation::Continue
+            } else {
+                Compilation::Stop
+            }
+        })
     }
 }