1use heck::{ToPascalCase, ToSnakeCase};
4use std::fmt::Write;
5
6type Result<T = ()> = std::result::Result<T, std::fmt::Error>;
7use zlink::idl::{CustomEnum, CustomObject, CustomType, Field, Interface, Method, Type};
8
9pub struct CodeGenerator {
11 output: String,
12 indent_level: usize,
13}
14
15impl CodeGenerator {
16 pub fn new() -> Self {
18 Self {
19 output: String::new(),
20 indent_level: 0,
21 }
22 }
23
24 pub fn output(self) -> String {
26 self.output
27 }
28
29 pub fn write_module_header(&mut self) -> Result<()> {
31 writeln!(
32 &mut self.output,
33 "// Generated code from Varlink IDL files."
34 )?;
35 writeln!(&mut self.output)?;
36 writeln!(&mut self.output, "use serde::{{Deserialize, Serialize}};")?;
37 writeln!(&mut self.output, "use zlink::{{proxy, ReplyError}};")?;
38 writeln!(&mut self.output)?;
39 Ok(())
40 }
41
42 pub fn generate_interface(
44 &mut self,
45 interface: &Interface<'_>,
46 skip_module_header: bool,
47 ) -> Result<()> {
48 if skip_module_header {
49 self.write_interface_comment(interface)?;
50 } else {
51 self.write_header(interface)?;
52 self.writeln("use serde::{Deserialize, Serialize};")?;
53 self.writeln("use zlink::{proxy, ReplyError};")?;
55 self.writeln("")?;
56 }
57
58 self.generate_proxy_trait(interface)?;
60 self.writeln("")?;
61
62 self.generate_output_structs(interface)?;
64
65 for custom_type in interface.custom_types() {
67 self.generate_custom_type(custom_type)?;
68 self.writeln("")?;
69 }
70
71 if interface.errors().count() > 0 {
73 self.generate_errors(interface)?;
74 self.writeln("")?;
75 }
76
77 Ok(())
78 }
79
80 fn write_interface_comment(&mut self, interface: &Interface<'_>) -> Result<()> {
81 writeln!(
82 &mut self.output,
83 "// Generated code for Varlink interface `{}`.",
84 interface.name()
85 )?;
86 writeln!(&mut self.output)?;
87 Ok(())
88 }
89
90 fn write_header(&mut self, interface: &Interface<'_>) -> Result<()> {
91 writeln!(
92 &mut self.output,
93 "//! Generated code for Varlink interface `{}`.",
94 interface.name()
95 )?;
96 writeln!(&mut self.output, "//!",)?;
97 writeln!(
98 &mut self.output,
99 "//! This code was generated by `zlink-codegen` from Varlink IDL.",
100 )?;
101 writeln!(
102 &mut self.output,
103 "//! You may prefer to adapt it, instead of using it verbatim.",
104 )?;
105 writeln!(&mut self.output)?;
106
107 for comment in interface.comments() {
109 writeln!(&mut self.output, "//! {}", comment.text())?;
110 }
111 writeln!(&mut self.output)?;
112
113 Ok(())
114 }
115
116 fn generate_custom_type(&mut self, custom_type: &CustomType<'_>) -> Result<()> {
117 match custom_type {
118 CustomType::Object(obj) => self.generate_custom_object(obj),
119 CustomType::Enum(enum_type) => self.generate_custom_enum(enum_type),
120 }
121 }
122
123 fn generate_custom_object(&mut self, obj: &CustomObject<'_>) -> Result<()> {
124 for comment in obj.comments() {
126 self.writeln(&format!("/// {}", comment.text()))?;
127 }
128
129 self.writeln("#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]")?;
130 self.writeln(&format!("pub struct {} {{", obj.name().to_pascal_case()))?;
131 self.indent();
132
133 for field in obj.fields() {
134 self.generate_field(field)?;
135 }
136
137 self.dedent();
138 self.writeln("}")?;
139
140 Ok(())
141 }
142
143 fn generate_custom_enum(&mut self, enum_type: &CustomEnum<'_>) -> Result<()> {
144 for comment in enum_type.comments() {
146 self.writeln(&format!("/// {}", comment.text()))?;
147 }
148
149 self.writeln("#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]")?;
150 self.writeln("#[serde(rename_all = \"snake_case\")]")?;
151 self.writeln(&format!(
152 "pub enum {} {{",
153 enum_type.name().to_pascal_case()
154 ))?;
155 self.indent();
156
157 for variant in enum_type.variants() {
158 for comment in variant.comments() {
160 self.writeln(&format!("/// {}", comment.text()))?;
161 }
162
163 self.writeln(&format!("{},", variant.name().to_pascal_case()))?;
165 }
166
167 self.dedent();
168 self.writeln("}")?;
169
170 Ok(())
171 }
172
173 fn generate_field(&mut self, field: &Field<'_>) -> Result<()> {
174 for comment in field.comments() {
176 self.writeln(&format!("/// {}", comment.text()))?;
177 }
178
179 let field_name = field.name().to_snake_case();
180 let rust_type = self.type_to_rust(field.ty())?;
181
182 let rust_type = if matches!(field.ty(), Type::Optional(_)) {
184 rust_type
186 } else {
187 rust_type
188 };
189
190 let field_name_attr = if is_rust_keyword(&field_name) || field_name != field.name() {
192 format!("#[serde(rename = \"{}\")]", field.name())
193 } else {
194 String::new()
195 };
196
197 if !field_name_attr.is_empty() {
198 self.writeln(&field_name_attr)?;
199 }
200
201 let safe_field_name = if is_rust_keyword(&field_name) {
202 format!("r#{}", field_name)
203 } else {
204 field_name
205 };
206
207 self.writeln(&format!("pub {}: {},", safe_field_name, rust_type))?;
208
209 Ok(())
210 }
211
212 fn generate_errors(&mut self, interface: &Interface<'_>) -> Result<()> {
213 self.writeln("/// Errors that can occur in this interface.")?;
214 self.writeln("#[derive(Debug, Clone, PartialEq, ReplyError)]")?;
215 self.writeln(&format!("#[zlink(interface = \"{}\")]", interface.name()))?;
216 self.writeln(&format!(
217 "pub enum {}Error {{",
218 interface_name_to_rust(interface.name())
219 ))?;
220 self.indent();
221
222 for error in interface.errors() {
223 for comment in error.comments() {
225 self.writeln(&format!("/// {}", comment.text()))?;
226 }
227
228 let variant_name = error.name().to_pascal_case();
229 if error.fields().count() == 0 {
230 self.writeln(&format!("{},", variant_name))?;
231 } else {
232 self.writeln(&format!("{} {{", variant_name))?;
233 self.indent();
234 for field in error.fields() {
235 self.generate_error_field(field)?;
236 }
237 self.dedent();
238 self.writeln("},")?;
239 }
240 }
241
242 self.dedent();
243 self.writeln("}")?;
244
245 Ok(())
246 }
247
248 fn generate_output_structs(&mut self, interface: &Interface<'_>) -> Result<()> {
250 for method in interface.methods() {
251 if method.outputs().count() > 0 {
255 let struct_name = format!("{}Output", method.name().to_pascal_case());
256
257 self.writeln(&format!(
259 "/// Output parameters for the {} method.",
260 method.name()
261 ))?;
262
263 let needs_lifetime = method.outputs().any(|o| type_needs_lifetime(o.ty()));
265
266 self.writeln("#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]")?;
267 if needs_lifetime {
268 self.writeln(&format!("pub struct {}<'a> {{", struct_name))?;
269 } else {
270 self.writeln(&format!("pub struct {} {{", struct_name))?;
271 }
272 self.indent();
273
274 for output in method.outputs() {
275 let field_name = output.name().to_snake_case();
276 let rust_type = if needs_lifetime {
278 self.type_to_rust_output(output.ty())?
279 } else {
280 self.type_to_rust(output.ty())?
281 };
282
283 if needs_lifetime && type_needs_borrow(output.ty()) {
285 self.writeln("#[serde(borrow)]")?;
286 }
287
288 if field_name != output.name() {
289 self.writeln(&format!("#[serde(rename = \"{}\")]", output.name()))?;
290 }
291
292 let safe_field_name = if is_rust_keyword(&field_name) {
293 format!("r#{}", field_name)
294 } else {
295 field_name
296 };
297
298 self.writeln(&format!("pub {}: {},", safe_field_name, rust_type))?;
299 }
300
301 self.dedent();
302 self.writeln("}")?;
303 self.writeln("")?;
304 }
305 }
306
307 Ok(())
308 }
309
310 fn generate_proxy_trait(&mut self, interface: &Interface<'_>) -> Result<()> {
311 let trait_name = interface_name_to_rust(interface.name());
312
313 let error_type = if interface.errors().count() > 0 {
315 format!("{}Error", interface_name_to_rust(interface.name()))
316 } else {
317 let stub_error_name = format!("{}Error", interface_name_to_rust(interface.name()));
319
320 self.writeln("/// Stub error type for interface without errors.")?;
322 self.writeln("///")?;
323 self.writeln("/// This is an empty enum that can never be instantiated.")?;
324 self.writeln("/// It exists only to satisfy the proxy trait requirements.")?;
325 self.writeln("#[derive(Debug, Clone, PartialEq, ReplyError)]")?;
326 self.writeln(&format!("#[zlink(interface = \"{}\")]", interface.name()))?;
327 self.writeln(&format!("pub enum {} {{}}", stub_error_name))?;
328 self.writeln("")?;
329
330 stub_error_name
331 };
332
333 self.writeln("/// Proxy trait for calling methods on the interface.")?;
334 self.writeln(&format!("#[proxy(\"{}\")]", interface.name()))?;
335 self.writeln(&format!("pub trait {} {{", trait_name))?;
336 self.indent();
337
338 for method in interface.methods() {
339 self.generate_proxy_method_signature(method, &error_type)?;
340 }
341
342 self.dedent();
343 self.writeln("}")?;
344
345 Ok(())
346 }
347
348 fn generate_proxy_method_signature(
349 &mut self,
350 method: &Method<'_>,
351 error_type: &str,
352 ) -> Result<()> {
353 for comment in method.comments() {
355 self.writeln(&format!("/// {}", comment.text()))?;
356 }
357
358 let method_name = method.name().to_snake_case();
359 let safe_method_name = if is_rust_keyword(&method_name) {
360 format!("r#{}", method_name)
361 } else {
362 method_name
363 };
364
365 let mut signature = format!("async fn {}(&mut self", safe_method_name);
367
368 for param in method.inputs() {
370 let param_name = param.name().to_snake_case();
371 let safe_param_name = if is_rust_keyword(¶m_name) {
372 format!("r#{}", param_name)
373 } else {
374 param_name
375 };
376 let rust_type = self.type_to_rust_param(param.ty())?;
378
379 write!(&mut signature, ",")?;
380 if safe_param_name != param.name() {
382 write!(&mut signature, " #[zlink(rename = \"{}\")]", param.name(),)?;
383 }
384
385 write!(&mut signature, " {}: {}", safe_param_name, rust_type)?;
386 }
387
388 signature.push_str(") -> zlink::Result<Result<");
389
390 let output_count = method.outputs().count();
392 if output_count == 0 {
393 signature.push_str("()");
394 } else {
395 let struct_name = format!("{}Output", method.name().to_pascal_case());
399 let needs_lifetime = method.outputs().any(|o| type_needs_lifetime(o.ty()));
401 if needs_lifetime {
402 signature.push_str(&format!("{}<'_>", struct_name));
403 } else {
404 signature.push_str(&struct_name);
405 }
406 }
407
408 write!(&mut signature, ", {}>>", error_type)?;
409 signature.push(';');
410
411 self.writeln(&signature)?;
412
413 Ok(())
414 }
415
416 fn generate_error_field(&mut self, field: &Field<'_>) -> Result<()> {
417 for comment in field.comments() {
419 self.writeln(&format!("/// {}", comment.text()))?;
420 }
421
422 let field_name = field.name().to_snake_case();
423 let rust_type = self.type_to_rust(field.ty())?;
424
425 let field_name_attr = if is_rust_keyword(&field_name) || field_name != field.name() {
427 format!("#[zlink(rename = \"{}\")]", field.name())
428 } else {
429 String::new()
430 };
431
432 if !field_name_attr.is_empty() {
433 self.writeln(&field_name_attr)?;
434 }
435
436 let safe_field_name = if is_rust_keyword(&field_name) {
437 format!("r#{}", field_name)
438 } else {
439 field_name
440 };
441
442 self.writeln(&format!("{}: {},", safe_field_name, rust_type))?;
443
444 Ok(())
445 }
446
447 fn type_to_rust(&self, ty: &Type) -> Result<String> {
448 type_to_rust(ty)
449 }
450
451 fn type_to_rust_param(&self, ty: &Type) -> Result<String> {
452 type_to_rust_param(ty)
453 }
454
455 fn type_to_rust_output(&self, ty: &Type) -> Result<String> {
456 type_to_rust_output(ty)
457 }
458
459 fn writeln(&mut self, s: &str) -> Result<()> {
460 self.write(s)?;
461 writeln!(&mut self.output)?;
462 Ok(())
463 }
464
465 fn write(&mut self, s: &str) -> Result<()> {
466 for _ in 0..self.indent_level {
467 write!(&mut self.output, " ")?;
468 }
469 write!(&mut self.output, "{}", s)?;
470 Ok(())
471 }
472
473 fn indent(&mut self) {
474 self.indent_level += 1;
475 }
476
477 fn dedent(&mut self) {
478 if self.indent_level > 0 {
479 self.indent_level -= 1;
480 }
481 }
482}
483
484impl Default for CodeGenerator {
485 fn default() -> Self {
486 Self::new()
487 }
488}
489
490fn type_to_rust(ty: &Type) -> Result<String> {
491 Ok(match ty {
492 Type::Bool => "bool".to_string(),
493 Type::Int => "i64".to_string(),
494 Type::Float => "f64".to_string(),
495 Type::String => "String".to_string(),
496 Type::Object(_fields) => {
497 "serde_json::Value".to_string()
501 }
502 Type::Enum(_variants) => {
503 "String".to_string()
505 }
506 Type::Array(elem_type) => {
507 let elem_rust = type_to_rust(elem_type.inner())?;
508 format!("Vec<{}>", elem_rust)
509 }
510 Type::Map(value_type) => {
511 let value_rust = type_to_rust(value_type.inner())?;
512 format!("std::collections::HashMap<String, {}>", value_rust)
513 }
514 Type::ForeignObject => "serde_json::Value".to_string(),
515 Type::Optional(inner_type) => {
516 let inner_rust = type_to_rust(inner_type.inner())?;
517 format!("Option<{}>", inner_rust)
518 }
519 Type::Custom(name) => name.to_pascal_case(),
520 Type::Any => "serde_json::Value".to_string(),
521 })
522}
523
524fn type_to_rust_param(ty: &Type) -> Result<String> {
525 Ok(match ty {
526 Type::Bool => "bool".to_string(),
527 Type::Int => "i64".to_string(),
528 Type::Float => "f64".to_string(),
529 Type::String => "&str".to_string(),
530 Type::Object(_fields) => {
531 "&serde_json::Value".to_string()
533 }
534 Type::Enum(_variants) => {
535 "&str".to_string()
537 }
538 Type::Array(elem_type) => {
539 let elem_rust = type_to_rust_param_elem(elem_type.inner())?;
541 format!("&[{}]", elem_rust)
542 }
543 Type::Map(value_type) => {
544 let value_rust = type_to_rust_param_elem(value_type.inner())?;
546 format!("&std::collections::HashMap<&str, {}>", value_rust)
547 }
548 Type::ForeignObject => "&serde_json::Value".to_string(),
549 Type::Optional(inner_type) => {
550 let inner_rust = type_to_rust_param(inner_type.inner())?;
551 format!("Option<{}>", inner_rust)
553 }
554 Type::Custom(name) => format!("&{}", name.to_pascal_case()),
555 Type::Any => "&serde_json::Value".to_string(),
556 })
557}
558
559fn type_to_rust_param_elem(ty: &Type) -> Result<String> {
562 Ok(match ty {
563 Type::Bool => "bool".to_string(),
564 Type::Int => "i64".to_string(),
565 Type::Float => "f64".to_string(),
566 Type::String => "&str".to_string(),
567 Type::Object(_fields) => "serde_json::Value".to_string(),
568 Type::Enum(_variants) => "&str".to_string(),
569 Type::Array(elem_type) => {
570 let elem_rust = type_to_rust_param_elem(elem_type.inner())?;
571 format!("Vec<{}>", elem_rust)
572 }
573 Type::Map(value_type) => {
574 let value_rust = type_to_rust_param_elem(value_type.inner())?;
575 format!("std::collections::HashMap<&str, {}>", value_rust)
576 }
577 Type::ForeignObject => "serde_json::Value".to_string(),
578 Type::Any => "serde_json::Value".to_string(),
579 Type::Optional(inner_type) => {
580 let inner_rust = type_to_rust_param_elem(inner_type.inner())?;
581 format!("Option<{}>", inner_rust)
582 }
583 Type::Custom(name) => name.to_pascal_case(),
584 })
585}
586
587fn type_to_rust_output(ty: &Type) -> Result<String> {
588 Ok(match ty {
589 Type::Bool => "bool".to_string(),
590 Type::Int => "i64".to_string(),
591 Type::Float => "f64".to_string(),
592 Type::String => "&'a str".to_string(),
593 Type::Object(_fields) => {
594 "serde_json::Value".to_string()
596 }
597 Type::Enum(_variants) => {
598 "&'a str".to_string()
600 }
601 Type::Array(elem_type) => {
602 let elem_rust = match elem_type.inner() {
604 Type::String => "&'a str".to_string(),
605 Type::Enum(_) => "&'a str".to_string(),
606 _ => type_to_rust(elem_type.inner())?,
607 };
608 format!("Vec<{}>", elem_rust)
609 }
610 Type::Map(value_type) => {
611 let value_rust = match value_type.inner() {
613 Type::String => "&'a str".to_string(),
614 Type::Enum(_) => "&'a str".to_string(),
615 _ => type_to_rust(value_type.inner())?,
616 };
617 format!("std::collections::HashMap<&'a str, {}>", value_rust)
618 }
619 Type::ForeignObject => "serde_json::Value".to_string(),
620 Type::Any => "serde_json::Value".to_string(),
621 Type::Optional(inner_type) => {
622 let inner_rust = type_to_rust_output(inner_type.inner())?;
625 format!("Option<{}>", inner_rust)
626 }
627 Type::Custom(name) => name.to_pascal_case(),
628 })
629}
630
631fn interface_name_to_rust(name: &str) -> String {
632 name.split('.').next_back().unwrap_or(name).to_pascal_case()
634}
635
636fn type_needs_lifetime(ty: &Type) -> bool {
637 match ty {
638 Type::String => true,
639 Type::Enum(_) => true, Type::Array(inner) => type_needs_lifetime(inner.inner()),
641 Type::Map(_) => {
642 true
644 }
645 Type::Optional(inner) => type_needs_lifetime(inner.inner()),
646 _ => false,
647 }
648}
649
650fn type_needs_borrow(ty: &Type) -> bool {
651 match ty {
652 Type::String => true,
653 Type::Enum(_) => true, Type::Array(inner) => type_needs_borrow(inner.inner()),
655 Type::Map(_) => {
656 true
658 }
659 Type::Optional(inner) => type_needs_borrow(inner.inner()),
660 _ => false,
661 }
662}
663
664fn is_rust_keyword(s: &str) -> bool {
665 [
666 "as", "async", "await", "break", "const", "continue", "crate", "dyn", "else", "enum",
667 "extern", "false", "fn", "for", "if", "impl", "in", "let", "loop", "match", "mod", "move",
668 "mut", "pub", "ref", "return", "self", "Self", "static", "struct", "super", "trait",
669 "true", "type", "unsafe", "use", "where", "while",
670 ]
671 .contains(&s)
672}