Skip to main content

melior_macro/dialect/
operation.rs

1mod attribute;
2mod builder;
3mod operand;
4mod operation_element;
5mod operation_field;
6mod region;
7mod result;
8mod successor;
9mod variadic_kind;
10
11pub use self::{
12    attribute::Attribute, builder::OperationBuilder, operand::Operand,
13    operation_element::OperationElement, region::Region, result::OperationResult,
14    successor::Successor, variadic_kind::VariadicKind,
15};
16use super::utility::{sanitize_documentation, sanitize_snake_case_identifier};
17use crate::dialect::{
18    error::{Error, OdsError},
19    r#trait::Trait,
20    r#type::Type,
21    utility::capitalize_string,
22};
23pub use operation_field::OperationField;
24use std::collections::HashSet;
25use syn::Ident;
26use tblgen::{error::WithLocation, record::Record, TypedInit};
27
28// spell-checker: disable-next-line
29const VOWELS: &str = "aeiou";
30
31#[derive(Debug)]
32pub struct Operation<'a> {
33    name: String,
34    short_dialect_name: &'a str,
35    dialect_name: &'a str,
36    operation_name: &'a str,
37    constructor_identifier: Ident,
38    summary: &'a str,
39    description: String,
40    can_infer_type: bool,
41    results: Vec<OperationResult<'a>>,
42    operands: Vec<Operand<'a>>,
43    regions: Vec<Region<'a>>,
44    successors: Vec<Successor<'a>>,
45    attributes: Vec<Attribute<'a>>,
46    derived_attributes: Vec<Attribute<'a>>,
47}
48
49impl<'a> Operation<'a> {
50    pub fn new(definition: Record<'a>) -> Result<Self, Error> {
51        let operation_name = definition.str_value("opName")?;
52        let traits = Self::collect_traits(definition)?;
53        let trait_names = traits
54            .iter()
55            .flat_map(|r#trait| r#trait.name())
56            .collect::<HashSet<_>>();
57
58        let arguments = Self::dag_constraints(definition, "arguments")?;
59        let regions = Self::collect_regions(definition)?;
60        let (results, unfixed_result_count) = Self::collect_results(
61            definition,
62            trait_names.contains("::mlir::OpTrait::SameVariadicResultSize"),
63            trait_names.contains("::mlir::OpTrait::AttrSizedResultSegments"),
64        )?;
65
66        Ok(Self {
67            name: Self::build_name(definition)?,
68            dialect_name: definition.def_value("opDialect")?.name()?,
69            short_dialect_name: definition.def_value("opDialect")?.str_value("name")?,
70            operation_name,
71            constructor_identifier: sanitize_snake_case_identifier(operation_name)?,
72            summary: definition.str_value("summary")?,
73            description: sanitize_documentation(definition.str_value("description")?)?,
74            can_infer_type: traits.iter().any(|r#trait| {
75                (r#trait.name() == Some("::mlir::OpTrait::FirstAttrDerivedResultType")
76                    || r#trait.name() == Some("::mlir::OpTrait::SameOperandsAndResultType"))
77                    && unfixed_result_count == 0
78                    || r#trait.name() == Some("::mlir::InferTypeOpInterface::Trait")
79                        && regions.is_empty()
80            }),
81            results,
82            operands: Self::collect_operands(
83                &arguments,
84                trait_names.contains("::mlir::OpTrait::SameVariadicOperandSize"),
85                trait_names.contains("::mlir::OpTrait::AttrSizedOperandSegments"),
86            )?,
87            regions,
88            successors: Self::collect_successors(definition)?,
89            attributes: Self::collect_attributes(&arguments)?,
90            derived_attributes: Self::collect_derived_attributes(definition)?,
91        })
92    }
93
94    fn build_name(definition: Record) -> Result<String, Error> {
95        let name = definition.name()?;
96
97        Ok(if let Some((_, name)) = name.split_once('_') {
98            name
99        } else {
100            name
101        }
102        .trim_end_matches("Op")
103        .to_owned()
104            + "Operation")
105    }
106
107    pub fn name(&self) -> &str {
108        &self.name
109    }
110
111    pub const fn can_infer_type(&self) -> bool {
112        self.can_infer_type
113    }
114
115    pub const fn dialect_name(&self) -> &str {
116        self.dialect_name
117    }
118
119    pub const fn operation_name(&self) -> &str {
120        self.operation_name
121    }
122
123    pub fn full_operation_name(&self) -> String {
124        format!("{}.{}", self.short_dialect_name, self.operation_name)
125    }
126
127    pub fn documentation_name(&self) -> String {
128        format!(
129            "{} [`{}`]({}) operation",
130            if VOWELS.contains(&self.operation_name()[..1]) {
131                "an"
132            } else {
133                "a"
134            },
135            self.operation_name,
136            &self.name
137        )
138    }
139
140    pub fn summary(&self) -> String {
141        format!(
142            "{}. {}",
143            capitalize_string(&self.documentation_name()),
144            if self.summary.is_empty() {
145                Default::default()
146            } else {
147                capitalize_string(self.summary) + "."
148            },
149        )
150    }
151
152    pub fn description(&self) -> &str {
153        &self.description
154    }
155
156    pub const fn constructor_identifier(&self) -> &Ident {
157        &self.constructor_identifier
158    }
159
160    pub fn results(&self) -> impl Iterator<Item = &OperationResult<'a>> + Clone {
161        self.results.iter()
162    }
163
164    pub fn result_len(&self) -> usize {
165        self.results.len()
166    }
167
168    pub fn operands(&self) -> impl Iterator<Item = &Operand<'a>> + Clone {
169        self.operands.iter()
170    }
171
172    pub fn operand_len(&self) -> usize {
173        self.operands.len()
174    }
175
176    pub fn regions(&self) -> impl Iterator<Item = &Region<'a>> {
177        self.regions.iter()
178    }
179
180    pub fn successors(&self) -> impl Iterator<Item = &Successor<'a>> {
181        self.successors.iter()
182    }
183
184    pub fn attributes(&self) -> impl Iterator<Item = &Attribute<'a>> {
185        self.attributes.iter()
186    }
187
188    pub fn all_attributes(&self) -> impl Iterator<Item = &Attribute<'a>> {
189        self.attributes().chain(&self.derived_attributes)
190    }
191
192    pub fn required_results(&self) -> impl Iterator<Item = &OperationResult> {
193        if self.can_infer_type {
194            Default::default()
195        } else {
196            self.results.iter()
197        }
198        .filter(|field| !field.is_optional())
199    }
200
201    pub fn required_operands(&self) -> impl Iterator<Item = &Operand> {
202        self.operands.iter().filter(|field| !field.is_optional())
203    }
204
205    pub fn required_regions(&self) -> impl Iterator<Item = &Region> {
206        self.regions.iter().filter(|field| !field.is_optional())
207    }
208
209    pub fn required_successors(&self) -> impl Iterator<Item = &Successor> {
210        self.successors.iter().filter(|field| !field.is_optional())
211    }
212
213    pub fn required_attributes(&self) -> impl Iterator<Item = &Attribute> {
214        self.attributes.iter().filter(|field| !field.is_optional())
215    }
216
217    pub fn required_fields(&self) -> impl Iterator<Item = &dyn OperationField> {
218        fn convert(field: &impl OperationField) -> &dyn OperationField {
219            field
220        }
221
222        self.required_results()
223            .map(convert)
224            .chain(self.required_operands().map(convert))
225            .chain(self.required_regions().map(convert))
226            .chain(self.required_successors().map(convert))
227            .chain(self.required_attributes().map(convert))
228    }
229
230    fn collect_successors(definition: Record<'a>) -> Result<Vec<Successor<'a>>, Error> {
231        definition
232            .dag_value("successors")?
233            .args()
234            .map(|(name, value)| {
235                Successor::new(
236                    name,
237                    Record::try_from(value)
238                        .map_err(|error| error.set_location(definition))?
239                        .subclass_of("VariadicSuccessor"),
240                )
241            })
242            .collect()
243    }
244
245    fn collect_regions(definition: Record<'a>) -> Result<Vec<Region<'a>>, Error> {
246        definition
247            .dag_value("regions")?
248            .args()
249            .map(|(name, value)| {
250                Region::new(
251                    name,
252                    Record::try_from(value)
253                        .map_err(|error| error.set_location(definition))?
254                        .subclass_of("VariadicRegion"),
255                )
256            })
257            .collect()
258    }
259
260    fn collect_traits(definition: Record<'a>) -> Result<Vec<Trait>, Error> {
261        let mut trait_lists = vec![definition.list_value("traits")?];
262        let mut traits = vec![];
263
264        while let Some(trait_list) = trait_lists.pop() {
265            for value in trait_list.iter() {
266                let definition =
267                    Record::try_from(value).map_err(|error| error.set_location(definition))?;
268
269                if definition.subclass_of("TraitList") {
270                    trait_lists.push(definition.list_value("traits")?);
271                } else {
272                    if definition.subclass_of("Interface") {
273                        trait_lists.push(definition.list_value("baseInterfaces")?);
274                    }
275                    traits.push(Trait::new(definition)?)
276                }
277            }
278        }
279
280        Ok(traits)
281    }
282
283    fn dag_constraints(
284        definition: Record<'a>,
285        name: &str,
286    ) -> Result<Vec<(&'a str, Record<'a>)>, Error> {
287        definition
288            .dag_value(name)?
289            .args()
290            .map(|(name, argument)| {
291                let definition =
292                    Record::try_from(argument).map_err(|error| error.set_location(definition))?;
293
294                Ok((
295                    name,
296                    if definition.subclass_of("OpVariable") {
297                        definition.def_value("constraint")?
298                    } else {
299                        definition
300                    },
301                ))
302            })
303            .collect()
304    }
305
306    fn collect_results(
307        definition: Record<'a>,
308        same_size: bool,
309        attribute_sized: bool,
310    ) -> Result<(Vec<OperationResult<'a>>, usize), Error> {
311        Self::collect_elements(
312            &Self::dag_constraints(definition, "results")?
313                .into_iter()
314                .map(|(name, constraint)| (name, Type::new(constraint)))
315                .collect::<Vec<_>>(),
316            OperationResult::new,
317            same_size,
318            attribute_sized,
319        )
320    }
321
322    fn collect_operands(
323        arguments: &[(&'a str, Record<'a>)],
324        same_size: bool,
325        attribute_sized: bool,
326    ) -> Result<Vec<Operand<'a>>, Error> {
327        Ok(Self::collect_elements(
328            &arguments
329                .iter()
330                .filter(|(_, definition)| definition.subclass_of("TypeConstraint"))
331                .map(|(name, definition)| (*name, Type::new(*definition)))
332                .collect::<Vec<_>>(),
333            Operand::new,
334            same_size,
335            attribute_sized,
336        )?
337        .0)
338    }
339
340    fn collect_elements<T>(
341        elements: &[(&'a str, Type)],
342        create: impl Fn(&'a str, Type, VariadicKind) -> Result<T, Error>,
343        same_size: bool,
344        attribute_sized: bool,
345    ) -> Result<(Vec<T>, usize), Error> {
346        let unfixed_count = elements
347            .iter()
348            .filter(|(_, r#type)| r#type.is_unfixed())
349            .count();
350        let mut variadic_kind = VariadicKind::new(unfixed_count, same_size, attribute_sized);
351        let mut fields = vec![];
352
353        for (name, r#type) in elements {
354            fields.push(create(name, *r#type, variadic_kind.clone())?);
355
356            match &mut variadic_kind {
357                VariadicKind::Simple { unfixed_seen } => {
358                    if r#type.is_unfixed() {
359                        *unfixed_seen = true;
360                    }
361                }
362                VariadicKind::SameSize {
363                    preceding_simple_count,
364                    preceding_variadic_count,
365                    ..
366                } => {
367                    if r#type.is_unfixed() {
368                        *preceding_variadic_count += 1;
369                    } else {
370                        *preceding_simple_count += 1;
371                    }
372                }
373                VariadicKind::AttributeSized => {}
374            }
375        }
376
377        Ok((fields, unfixed_count))
378    }
379
380    fn collect_attributes(
381        arguments: &[(&'a str, Record<'a>)],
382    ) -> Result<Vec<Attribute<'a>>, Error> {
383        arguments
384            .iter()
385            .filter(|(_, definition)| definition.subclass_of("Attr"))
386            .map(|(name, definition)| {
387                if definition.subclass_of("DerivedAttr") {
388                    Err(OdsError::UnexpectedSuperClass("DerivedAttr")
389                        .with_location(*definition)
390                        .into())
391                } else {
392                    Attribute::new(name, *definition)
393                }
394            })
395            .collect()
396    }
397
398    fn collect_derived_attributes(definition: Record<'a>) -> Result<Vec<Attribute<'a>>, Error> {
399        definition
400            .values()
401            .filter(|value| matches!(value.init, TypedInit::Def(_)))
402            .map(Record::try_from)
403            .collect::<Result<Vec<_>, _>>()?
404            .into_iter()
405            .filter(|definition| definition.subclass_of("Attr"))
406            .map(|definition| {
407                if definition.subclass_of("DerivedAttr") {
408                    Attribute::new(definition.name()?, definition)
409                } else {
410                    Err(OdsError::ExpectedSuperClass("DerivedAttr")
411                        .with_location(definition)
412                        .into())
413                }
414            })
415            .collect()
416    }
417}