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
28const 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}