1use 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#[derive(Debug, thiserror::Error, PartialEq, Eq)]
13pub enum RenderError {
14 #[error("{construct} cannot be rendered as {language:?}")]
16 LanguageMismatch {
17 language: Language,
19 construct: &'static str,
21 },
22 #[error("function declaration has a non-function signature")]
24 InvalidFunctionSignature,
25 #[error("Objective-C selector arity does not match its parameter list")]
27 SelectorArity,
28}
29
30pub 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(¶meter.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(¶meter.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}