Skip to main content

macho_header_syntax/
render.rs

1//! Deterministic rendering from typed header syntax.
2
3use std::fmt::Write;
4
5use crate::{
6    Access, BuiltinType, Decl, FunctionQualifiers, Language, Linkage, MethodKind, NamedTypeTag,
7    ObjectiveCForwardKind, Parameter, ParameterState, RecordKind, ReferenceKind, StorageClass,
8    TemplateArgument, TranslationUnit, Type, TypeQualifiers,
9};
10
11/// Rendering failure for a semantically incompatible AST.
12#[derive(Debug, thiserror::Error, PartialEq, Eq)]
13pub enum RenderError {
14    /// A declaration cannot be expressed in the selected language.
15    #[error("{construct} cannot be rendered as {language:?}")]
16    LanguageMismatch {
17        /// Selected language.
18        language: Language,
19        /// Incompatible construct.
20        construct: &'static str,
21    },
22    /// A function declaration did not contain a function type.
23    #[error("function declaration has a non-function signature")]
24    InvalidFunctionSignature,
25    /// An Objective-C selector and parameter list disagree.
26    #[error("Objective-C selector arity does not match its parameter list")]
27    SelectorArity,
28}
29
30/// Renders a complete translation unit deterministically.
31pub fn render(unit: &TranslationUnit) -> Result<String, RenderError> {
32    let mut output = String::new();
33    for declaration in &unit.declarations {
34        render_decl(declaration, unit.language, 0, &mut output)?;
35    }
36    Ok(output)
37}
38
39fn render_decl(
40    declaration: &Decl,
41    language: Language,
42    indent: usize,
43    output: &mut String,
44) -> Result<(), RenderError> {
45    let prefix = "    ".repeat(indent);
46    match declaration {
47        Decl::Function {
48            name,
49            signature,
50            storage,
51            linkage,
52        } => {
53            ensure_linkage(language, *linkage)?;
54            let Type::Function {
55                return_type,
56                parameters,
57                parameter_state,
58                variadic,
59                qualifiers,
60                ..
61            } = signature
62            else {
63                return Err(RenderError::InvalidFunctionSignature);
64            };
65            write!(output, "{prefix}{}", render_storage(*storage)).unwrap();
66            render_type(return_type, language, output)?;
67            write!(output, " {name}(").unwrap();
68            render_parameters(parameters, *parameter_state, *variadic, language, output)?;
69            writeln!(output, "){};", render_function_qualifiers(*qualifiers)).unwrap();
70        }
71        Decl::Variable {
72            name,
73            ty,
74            storage,
75            linkage,
76        } => {
77            ensure_linkage(language, *linkage)?;
78            write!(output, "{prefix}{}", render_storage(*storage)).unwrap();
79            render_type(ty, language, output)?;
80            writeln!(output, " {name};").unwrap();
81        }
82        Decl::Record {
83            kind,
84            path,
85            bases,
86            fields,
87            members,
88        } => {
89            if *kind == RecordKind::Class && language != Language::Cpp {
90                return Err(RenderError::LanguageMismatch {
91                    language,
92                    construct: "class",
93                });
94            }
95            write!(
96                output,
97                "{prefix}{} {}",
98                render_record_kind(*kind),
99                render_path(path)
100            )
101            .unwrap();
102            if !bases.is_empty() {
103                output.push_str(" : ");
104                for (index, base) in bases.iter().enumerate() {
105                    if index != 0 {
106                        output.push_str(", ");
107                    }
108                    if base.is_virtual {
109                        output.push_str("virtual ");
110                    }
111                    write!(output, "{} ", render_access(base.access)).unwrap();
112                    render_type(&base.ty, language, output)?;
113                }
114            }
115            output.push_str(" {\n");
116            for field in fields {
117                write!(output, "{prefix}    ").unwrap();
118                render_type(&field.ty, language, output)?;
119                write!(output, " {}", field.name).unwrap();
120                if let Some(width) = field.bit_width {
121                    write!(output, " : {width}").unwrap();
122                }
123                output.push_str(";\n");
124            }
125            for member in members {
126                render_decl(member, language, indent + 1, output)?;
127            }
128            writeln!(output, "{prefix}}};").unwrap();
129        }
130        Decl::Forward { kind, path } => {
131            writeln!(
132                output,
133                "{prefix}{} {};",
134                render_record_kind(*kind),
135                render_path(path)
136            )
137            .unwrap();
138        }
139        Decl::Alias { path, target } => {
140            if language == Language::Cpp {
141                write!(output, "{prefix}using {} = ", render_path(path)).unwrap();
142                render_type(target, language, output)?;
143                output.push_str(";\n");
144            } else {
145                write!(output, "{prefix}typedef ").unwrap();
146                render_type(target, language, output)?;
147                writeln!(output, " {};", render_path(path)).unwrap();
148            }
149        }
150        Decl::ObjectiveCInterface {
151            name,
152            superclass,
153            protocols,
154            ivars,
155            methods,
156            properties,
157        } => {
158            ensure_objc(language, "Objective-C interface")?;
159            write!(output, "{prefix}@interface {name}").unwrap();
160            if let Some(superclass) = superclass {
161                write!(output, " : {superclass}").unwrap();
162            }
163            render_protocol_list(protocols, output);
164            output.push('\n');
165            if !ivars.is_empty() {
166                output.push_str("{\n");
167                let mut access = None;
168                for ivar in ivars {
169                    if access != Some(ivar.access) {
170                        writeln!(output, "{}", render_objc_access(ivar.access)).unwrap();
171                        access = Some(ivar.access);
172                    }
173                    output.push_str("    ");
174                    render_type(&ivar.ty, Language::ObjectiveC, output)?;
175                    writeln!(output, " {};", ivar.name).unwrap();
176                }
177                output.push_str("}\n");
178            }
179            for property in properties {
180                render_objc_property(property, output)?;
181            }
182            for method in methods {
183                render_objc_method(method, output)?;
184            }
185            output.push_str("@end\n");
186        }
187        Decl::ObjectiveCCategory {
188            name,
189            extended_class,
190            protocols,
191            methods,
192            properties,
193        } => {
194            ensure_objc(language, "Objective-C category")?;
195            write!(output, "{prefix}@interface {extended_class} ({name})").unwrap();
196            render_protocol_list(protocols, output);
197            output.push('\n');
198            for property in properties {
199                render_objc_property(property, output)?;
200            }
201            for method in methods {
202                render_objc_method(method, output)?;
203            }
204            output.push_str("@end\n");
205        }
206        Decl::ObjectiveCProtocol {
207            name,
208            protocols,
209            methods,
210            properties,
211        } => {
212            ensure_objc(language, "Objective-C protocol")?;
213            write!(output, "{prefix}@protocol {name}").unwrap();
214            render_protocol_list(protocols, output);
215            output.push('\n');
216            for property in properties {
217                render_objc_property(property, output)?;
218            }
219            for method in methods {
220                render_objc_method(method, output)?;
221            }
222            output.push_str("@end\n");
223        }
224        Decl::ObjectiveCForward { kind, names } => {
225            ensure_objc(language, "Objective-C forward declaration")?;
226            let keyword = match kind {
227                ObjectiveCForwardKind::Class => "@class",
228                ObjectiveCForwardKind::Protocol => "@protocol",
229            };
230            write!(output, "{prefix}{keyword} ").unwrap();
231            for (index, name) in names.iter().enumerate() {
232                if index != 0 {
233                    output.push_str(", ");
234                }
235                write!(output, "{name}").unwrap();
236            }
237            output.push_str(";\n");
238        }
239    }
240    Ok(())
241}
242
243fn render_type(ty: &Type, language: Language, output: &mut String) -> Result<(), RenderError> {
244    match ty {
245        Type::Builtin(builtin) => output.push_str(render_builtin(*builtin)),
246        Type::Named {
247            tag,
248            path,
249            template_arguments,
250        } => {
251            if !matches!(tag, NamedTypeTag::Typedef) {
252                output.push_str(render_named_tag(*tag));
253                output.push(' ');
254            }
255            output.push_str(&render_path(path));
256            if !template_arguments.is_empty() {
257                ensure_cpp(language, "template argument")?;
258                output.push('<');
259                for (index, argument) in template_arguments.iter().enumerate() {
260                    if index != 0 {
261                        output.push_str(", ");
262                    }
263                    match argument {
264                        TemplateArgument::Type(ty) => render_type(ty, language, output)?,
265                        TemplateArgument::Integer(value) => write!(output, "{value}").unwrap(),
266                        TemplateArgument::Identifier(path) => output.push_str(&render_path(path)),
267                    }
268                }
269                output.push('>');
270            }
271        }
272        Type::Pointer {
273            pointee,
274            qualifiers,
275        } => {
276            render_type(pointee, language, output)?;
277            output.push_str(" *");
278            render_qualifiers(*qualifiers, output);
279        }
280        Type::Reference { target, kind } => {
281            ensure_cpp(language, "reference")?;
282            render_type(target, language, output)?;
283            output.push_str(match kind {
284                ReferenceKind::Lvalue => " &",
285                ReferenceKind::Rvalue => " &&",
286            });
287        }
288        Type::Array { element, count } => {
289            render_type(element, language, output)?;
290            match count {
291                Some(count) => write!(output, "[{count}]").unwrap(),
292                None => output.push_str("[]"),
293            }
294        }
295        Type::Function {
296            return_type,
297            parameters,
298            parameter_state,
299            variadic,
300            qualifiers,
301            ..
302        } => {
303            render_type(return_type, language, output)?;
304            output.push_str(" (");
305            render_parameters(parameters, *parameter_state, *variadic, language, output)?;
306            output.push(')');
307            output.push_str(&render_function_qualifiers(*qualifiers));
308        }
309        Type::ObjectiveCObject {
310            name,
311            protocols,
312            qualifiers,
313        } => {
314            ensure_objc(language, "Objective-C object")?;
315            if let Some(name) = name {
316                write!(output, "{name}").unwrap();
317            } else {
318                output.push_str("id");
319            }
320            render_protocol_list(protocols, output);
321            if name.is_some() {
322                output.push_str(" *");
323            }
324            render_qualifiers(*qualifiers, output);
325        }
326        Type::ObjectiveCBlock(signature) => {
327            ensure_objc(language, "Objective-C block")?;
328            render_type(signature, language, output)?;
329        }
330    }
331    Ok(())
332}
333
334fn render_parameters(
335    parameters: &[Parameter],
336    state: ParameterState,
337    variadic: bool,
338    language: Language,
339    output: &mut String,
340) -> Result<(), RenderError> {
341    if parameters.is_empty() && state == ParameterState::Known && language == Language::C {
342        output.push_str("void");
343    }
344    for (index, parameter) in parameters.iter().enumerate() {
345        if index != 0 {
346            output.push_str(", ");
347        }
348        render_type(&parameter.ty, language, output)?;
349        write!(output, " {}", parameter.name).unwrap();
350    }
351    if variadic {
352        if !parameters.is_empty() {
353            output.push_str(", ");
354        }
355        output.push_str("...");
356    }
357    Ok(())
358}
359
360fn render_objc_method(
361    method: &crate::ObjectiveCMethod,
362    output: &mut String,
363) -> Result<(), RenderError> {
364    output.push_str(match method.kind {
365        MethodKind::Instance => "- (",
366        MethodKind::Class => "+ (",
367    });
368    render_type(&method.return_type, Language::ObjectiveC, output)?;
369    output.push(')');
370    let pieces = method.selector.split_terminator(':').collect::<Vec<_>>();
371    if pieces.len() != method.parameters.len() {
372        if method.parameters.is_empty() && !method.selector.contains(':') {
373            output.push_str(&method.selector);
374            output.push_str(";\n");
375            return Ok(());
376        }
377        return Err(RenderError::SelectorArity);
378    }
379    for (piece, parameter) in pieces.into_iter().zip(&method.parameters) {
380        write!(output, "{piece}:(").unwrap();
381        render_type(&parameter.ty, Language::ObjectiveC, output)?;
382        write!(output, "){} ", parameter.name).unwrap();
383    }
384    if output.ends_with(' ') {
385        output.pop();
386    }
387    output.push_str(";\n");
388    Ok(())
389}
390
391fn render_objc_property(
392    property: &crate::ObjectiveCProperty,
393    output: &mut String,
394) -> Result<(), RenderError> {
395    output.push_str("@property");
396    if !property.attributes.is_empty() {
397        output.push_str(" (");
398        for (index, attribute) in property.attributes.iter().enumerate() {
399            if index != 0 {
400                output.push_str(", ");
401            }
402            output.push_str(render_objc_property_attribute(*attribute));
403        }
404        output.push(')');
405    }
406    output.push(' ');
407    render_type(&property.ty, Language::ObjectiveC, output)?;
408    writeln!(output, " {};", property.name).unwrap();
409    Ok(())
410}
411
412fn render_objc_access(access: crate::ObjectiveCAccess) -> &'static str {
413    match access {
414        crate::ObjectiveCAccess::Public => "@public",
415        crate::ObjectiveCAccess::Protected => "@protected",
416        crate::ObjectiveCAccess::Private => "@private",
417        crate::ObjectiveCAccess::Package => "@package",
418    }
419}
420
421fn render_objc_property_attribute(attribute: crate::ObjectiveCPropertyAttribute) -> &'static str {
422    use crate::ObjectiveCPropertyAttribute as Attribute;
423    match attribute {
424        Attribute::Readonly => "readonly",
425        Attribute::Readwrite => "readwrite",
426        Attribute::Copy => "copy",
427        Attribute::Retain => "retain",
428        Attribute::Strong => "strong",
429        Attribute::Weak => "weak",
430        Attribute::Assign => "assign",
431        Attribute::Atomic => "atomic",
432        Attribute::Nonatomic => "nonatomic",
433        Attribute::Dynamic => "dynamic",
434        Attribute::Class => "class",
435    }
436}
437
438fn ensure_linkage(language: Language, linkage: Linkage) -> Result<(), RenderError> {
439    let compatible = matches!(
440        (language, linkage),
441        (Language::C, Linkage::C)
442            | (Language::Cpp, Linkage::C | Linkage::Cpp)
443            | (Language::ObjectiveC, Linkage::C | Linkage::ObjectiveC)
444    );
445    if compatible {
446        Ok(())
447    } else {
448        Err(RenderError::LanguageMismatch {
449            language,
450            construct: "language linkage",
451        })
452    }
453}
454
455fn ensure_cpp(language: Language, construct: &'static str) -> Result<(), RenderError> {
456    if language == Language::Cpp {
457        Ok(())
458    } else {
459        Err(RenderError::LanguageMismatch {
460            language,
461            construct,
462        })
463    }
464}
465
466fn ensure_objc(language: Language, construct: &'static str) -> Result<(), RenderError> {
467    if language == Language::ObjectiveC {
468        Ok(())
469    } else {
470        Err(RenderError::LanguageMismatch {
471            language,
472            construct,
473        })
474    }
475}
476
477fn render_storage(storage: StorageClass) -> &'static str {
478    match storage {
479        StorageClass::None => "",
480        StorageClass::Extern => "extern ",
481        StorageClass::Static => "static ",
482        StorageClass::ThreadLocal => "thread_local ",
483    }
484}
485
486fn render_record_kind(kind: RecordKind) -> &'static str {
487    match kind {
488        RecordKind::Struct => "struct",
489        RecordKind::Union => "union",
490        RecordKind::Class => "class",
491        RecordKind::Enum => "enum",
492    }
493}
494
495fn render_named_tag(tag: NamedTypeTag) -> &'static str {
496    match tag {
497        NamedTypeTag::Typedef => "",
498        NamedTypeTag::Struct => "struct",
499        NamedTypeTag::Union => "union",
500        NamedTypeTag::Enum => "enum",
501        NamedTypeTag::Class => "class",
502        NamedTypeTag::Protocol => "protocol",
503    }
504}
505
506fn render_builtin(builtin: BuiltinType) -> &'static str {
507    match builtin {
508        BuiltinType::Void => "void",
509        BuiltinType::Bool => "bool",
510        BuiltinType::Char => "char",
511        BuiltinType::SignedChar => "signed char",
512        BuiltinType::UnsignedChar => "unsigned char",
513        BuiltinType::Short => "short",
514        BuiltinType::UnsignedShort => "unsigned short",
515        BuiltinType::Int => "int",
516        BuiltinType::UnsignedInt => "unsigned int",
517        BuiltinType::Long => "long",
518        BuiltinType::UnsignedLong => "unsigned long",
519        BuiltinType::LongLong => "long long",
520        BuiltinType::UnsignedLongLong => "unsigned long long",
521        BuiltinType::Int128 => "__int128",
522        BuiltinType::UnsignedInt128 => "unsigned __int128",
523        BuiltinType::Float => "float",
524        BuiltinType::Double => "double",
525        BuiltinType::LongDouble => "long double",
526    }
527}
528
529fn render_qualifiers(qualifiers: TypeQualifiers, output: &mut String) {
530    if qualifiers.is_const {
531        output.push_str(" const");
532    }
533    if qualifiers.is_volatile {
534        output.push_str(" volatile");
535    }
536    if qualifiers.is_restrict {
537        output.push_str(" restrict");
538    }
539}
540
541fn render_function_qualifiers(qualifiers: FunctionQualifiers) -> String {
542    let mut output = String::new();
543    if qualifiers.is_const {
544        output.push_str(" const");
545    }
546    if qualifiers.is_volatile {
547        output.push_str(" volatile");
548    }
549    if let Some(reference) = qualifiers.reference {
550        output.push_str(match reference {
551            ReferenceKind::Lvalue => " &",
552            ReferenceKind::Rvalue => " &&",
553        });
554    }
555    if let Some(noexcept) = qualifiers.noexcept {
556        output.push_str(if noexcept {
557            " noexcept"
558        } else {
559            " noexcept(false)"
560        });
561    }
562    output
563}
564
565fn render_protocol_list(protocols: &[crate::Identifier], output: &mut String) {
566    if protocols.is_empty() {
567        return;
568    }
569    output.push('<');
570    for (index, protocol) in protocols.iter().enumerate() {
571        if index != 0 {
572            output.push_str(", ");
573        }
574        write!(output, "{protocol}").unwrap();
575    }
576    output.push('>');
577}
578
579fn render_access(access: Access) -> &'static str {
580    match access {
581        Access::Public => "public",
582        Access::Protected => "protected",
583        Access::Private => "private",
584        Access::Unspecified => "public",
585    }
586}
587
588fn render_path(path: &crate::IdentifierPath) -> String {
589    path.components()
590        .iter()
591        .map(ToString::to_string)
592        .collect::<Vec<_>>()
593        .join("::")
594}
595
596#[cfg(test)]
597mod tests {
598    use crate::{HeaderParser, TreeSitterHeaderParser};
599
600    use super::*;
601
602    #[test]
603    fn rendered_c_reparses() {
604        let source = "struct Point { int x; int y; };\nint distance(struct Point *point);";
605        let unit = TreeSitterHeaderParser.parse(Language::C, source).unwrap();
606        let rendered = render(&unit).unwrap();
607        TreeSitterHeaderParser
608            .parse(Language::C, &rendered)
609            .unwrap();
610    }
611
612    #[test]
613    fn rendered_objective_c_reparses() {
614        let source = "@interface Widget : NSObject\n- (int)value;\n@end";
615        let unit = TreeSitterHeaderParser
616            .parse(Language::ObjectiveC, source)
617            .unwrap();
618        let rendered = render(&unit).unwrap();
619        TreeSitterHeaderParser
620            .parse(Language::ObjectiveC, &rendered)
621            .unwrap();
622    }
623}