Skip to main content

zlink_codegen/
codegen.rs

1//! Code generation implementation.
2
3use 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
9/// Code generator for Varlink interfaces.
10pub struct CodeGenerator {
11    output: String,
12    indent_level: usize,
13}
14
15impl CodeGenerator {
16    /// Create a new code generator.
17    pub fn new() -> Self {
18        Self {
19            output: String::new(),
20            indent_level: 0,
21        }
22    }
23
24    /// Get the generated output.
25    pub fn output(self) -> String {
26        self.output
27    }
28
29    /// Write module-level header for multiple interfaces.
30    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    /// Generate code for an interface.
43    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            // Always import ReplyError since we generate a stub error type when there are no errors
54            self.writeln("use zlink::{proxy, ReplyError};")?;
55            self.writeln("")?;
56        }
57
58        // Generate proxy trait using the proxy macro.
59        self.generate_proxy_trait(interface)?;
60        self.writeln("")?;
61
62        // Generate output structs for methods.
63        self.generate_output_structs(interface)?;
64
65        // Generate custom types.
66        for custom_type in interface.custom_types() {
67            self.generate_custom_type(custom_type)?;
68            self.writeln("")?;
69        }
70
71        // Generate errors.
72        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        // Add interface comments if any.
108        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        // Add comments.
125        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        // Add comments.
145        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            // Add variant comments.
159            for comment in variant.comments() {
160                self.writeln(&format!("/// {}", comment.text()))?;
161            }
162
163            // Varlink enum variants don't have explicit values, just names.
164            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        // Add field comments.
175        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        // Check if the field type is optional.
183        let rust_type = if matches!(field.ty(), Type::Optional(_)) {
184            // The type_to_rust will already wrap in Option
185            rust_type
186        } else {
187            rust_type
188        };
189
190        // Handle field name if it's a Rust keyword.
191        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            // Add error comments.
224            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    /// Generate output structs for all methods in the `interface`.
249    fn generate_output_structs(&mut self, interface: &Interface<'_>) -> Result<()> {
250        for method in interface.methods() {
251            // Generate output struct for any method with at least one output parameter.
252            // Varlink output parameters are always named, so we need a struct even for single
253            // outputs.
254            if method.outputs().count() > 0 {
255                let struct_name = format!("{}Output", method.name().to_pascal_case());
256
257                // Add method comments if available
258                self.writeln(&format!(
259                    "/// Output parameters for the {} method.",
260                    method.name()
261                ))?;
262
263                // Add lifetime parameter for output structs that need it
264                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                    // Use reference types for output parameters where appropriate
277                    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                    // Add #[serde(borrow)] for fields that need it
284                    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        // Generate a stub error type if there are no errors in the interface
314        let error_type = if interface.errors().count() > 0 {
315            format!("{}Error", interface_name_to_rust(interface.name()))
316        } else {
317            // Generate a stub error type for interfaces without errors
318            let stub_error_name = format!("{}Error", interface_name_to_rust(interface.name()));
319
320            // Generate the stub error type before the proxy trait
321            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        // Add method comments.
354        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        // Generate method signature.
366        let mut signature = format!("async fn {}(&mut self", safe_method_name);
367
368        // Add input parameters.
369        for param in method.inputs() {
370            let param_name = param.name().to_snake_case();
371            let safe_param_name = if is_rust_keyword(&param_name) {
372                format!("r#{}", param_name)
373            } else {
374                param_name
375            };
376            // Use references for parameters that can be borrowed
377            let rust_type = self.type_to_rust_param(param.ty())?;
378
379            write!(&mut signature, ",")?;
380            // Add parameter with potential rename attribute.
381            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        // Handle output parameters.
391        let output_count = method.outputs().count();
392        if output_count == 0 {
393            signature.push_str("()");
394        } else {
395            // Always use the generated output struct for any outputs.
396            // Varlink output parameters are always named, so we need a struct even for single
397            // outputs.
398            let struct_name = format!("{}Output", method.name().to_pascal_case());
399            // Add lifetime parameter if the struct needs one
400            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        // Add field comments.
418        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        // Handle field name if it's a Rust keyword.
426        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            // Anonymous struct - generate inline.
498            // For now, use serde_json::Value for anonymous objects.
499            // In the future, we could generate anonymous structs.
500            "serde_json::Value".to_string()
501        }
502        Type::Enum(_variants) => {
503            // Anonymous enum - use String for now.
504            "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            // For parameters, use reference to avoid clone
532            "&serde_json::Value".to_string()
533        }
534        Type::Enum(_variants) => {
535            // Anonymous enum - use &str for parameters
536            "&str".to_string()
537        }
538        Type::Array(elem_type) => {
539            // Use slice for array parameters with proper string handling
540            let elem_rust = type_to_rust_param_elem(elem_type.inner())?;
541            format!("&[{}]", elem_rust)
542        }
543        Type::Map(value_type) => {
544            // Use reference for map parameters with proper string handling
545            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            // For optional parameters, always wrap in Option
552            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
559// Helper function to get the proper type for collection elements in parameters.
560// Ensures strings always use &str instead of String.
561fn 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            // Use owned type for objects - serde can't deserialize to &Value
595            "serde_json::Value".to_string()
596        }
597        Type::Enum(_variants) => {
598            // Anonymous enum - use &str for outputs
599            "&'a str".to_string()
600        }
601        Type::Array(elem_type) => {
602            // Use Vec for array outputs with owned inner types (except strings stay as &'a str)
603            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            // Use HashMap for map outputs with borrowed types for efficiency
612            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            // For optional outputs, recursively apply type_to_rust_output to maintain
623            // correct reference types for strings within collections
624            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    // Convert interface name like "org.example.Interface" to "Interface".
633    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, // Anonymous enums use &'a str
640        Type::Array(inner) => type_needs_lifetime(inner.inner()),
641        Type::Map(_) => {
642            // Maps always need lifetime because keys are &'a str
643            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, // Anonymous enums use &'a str
654        Type::Array(inner) => type_needs_borrow(inner.inner()),
655        Type::Map(_) => {
656            // Maps always need borrow because keys are &'a str
657            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}