Skip to main content

melior/pass/
operation_manager.rs

1use super::PassManager;
2use crate::{pass::Pass, string_ref::StringRef};
3use mlir_sys::{
4    mlirOpPassManagerAddOwnedPass, mlirOpPassManagerGetNestedUnder, mlirPrintPassPipeline,
5    MlirOpPassManager, MlirStringRef,
6};
7use std::{
8    ffi::c_void,
9    fmt::{self, Display, Formatter},
10    marker::PhantomData,
11};
12
13/// An operation pass manager.
14#[derive(Clone, Copy, Debug)]
15pub struct OperationPassManager<'c, 'a> {
16    raw: MlirOpPassManager,
17    _parent: PhantomData<&'a PassManager<'c>>,
18}
19
20impl OperationPassManager<'_, '_> {
21    /// Returns an operation pass manager for nested operations corresponding to
22    /// a given name.
23    pub fn nested_under(&self, name: &str) -> Self {
24        let name = StringRef::new(name);
25
26        unsafe { Self::from_raw(mlirOpPassManagerGetNestedUnder(self.raw, name.to_raw())) }
27    }
28
29    /// Adds a pass.
30    pub fn add_pass(&self, pass: Pass) {
31        unsafe { mlirOpPassManagerAddOwnedPass(self.raw, pass.to_raw()) }
32    }
33
34    /// Converts an operation pass manager into a raw object.
35    pub const fn to_raw(self) -> MlirOpPassManager {
36        self.raw
37    }
38
39    /// Creates an operation pass manager from a raw object.
40    ///
41    /// # Safety
42    ///
43    /// A raw object must be valid.
44    pub unsafe fn from_raw(raw: MlirOpPassManager) -> Self {
45        Self {
46            raw,
47            _parent: Default::default(),
48        }
49    }
50}
51
52impl Display for OperationPassManager<'_, '_> {
53    fn fmt(&self, formatter: &mut Formatter) -> fmt::Result {
54        let mut data = (formatter, Ok(()));
55
56        unsafe extern "C" fn callback(string: MlirStringRef, data: *mut c_void) {
57            let data = &mut *(data as *mut (&mut Formatter, fmt::Result));
58            let result = (|| -> fmt::Result {
59                write!(
60                    data.0,
61                    "{}",
62                    StringRef::from_raw(string)
63                        .as_str()
64                        .map_err(|_| fmt::Error)?
65                )
66            })();
67
68            if data.1.is_ok() {
69                data.1 = result;
70            }
71        }
72
73        unsafe {
74            mlirPrintPassPipeline(self.raw, Some(callback), &mut data as *mut _ as *mut c_void);
75        }
76
77        data.1
78    }
79}