Skip to main content

robonix_codegen/codegen/
proto_gen.rs

1// SPDX-License-Identifier: MulanPSL-2.0
2// Protobuf (.proto) code generator — ROS IDL -> proto3
3//
4// Emits **one `{package}.proto` per ROS package** that has content:
5//   - Every `.msg` -> a `message` (always).
6//   - `*_Request` / `*_Response` from `.srv` **only** when `--contracts` is passed and that srv is listed in a contract `[io.srv].srv`.
7// Does **not** emit per-package `service FooService { rpc ... }`; gRPC facades are **only** in `robonix_contracts.proto`.
8
9use anyhow::{Context, Result};
10use std::collections::BTreeSet;
11use std::fmt::Write as FmtWrite;
12use std::fs;
13use std::path::Path;
14
15use super::msg_parser::{
16    MsgConstant, MsgField, MsgResolver, MsgSpec, MsgTypeRef, RosPrimitive, SrvSpec,
17};
18
19fn proto_primitive_type(p: &str) -> &'static str {
20    // Single source of truth for primitive → proto3 mapping is
21    // RosPrimitive::proto_type. Panic on a primitive the parser
22    // wouldn't have accepted — that's a codegen bug, not a runtime
23    // condition we want to silently swallow as `bytes` (which used to
24    // be the catch-all and produced wrong wire formats).
25    let prim = RosPrimitive::parse(p).unwrap_or_else(|| {
26        panic!(
27            "[robonix-codegen] proto: primitive '{p}' is not in RosPrimitive::parse — \
28             every primitive accepted by the parser must have a proto mapping; \
29             extend RosPrimitive if this is a new ROS primitive."
30        )
31    });
32    prim.proto_type()
33}
34
35fn proto_field_type(field: &MsgField, current_package: &str) -> String {
36    let base = match &field.type_ref {
37        MsgTypeRef::Primitive(p) => {
38            // `uint8[]` (unsigned 8-bit, unbounded) is the canonical
39            // raw-byte buffer and becomes proto `bytes`. `uint8[N]`
40            // (fixed-size) does NOT collapse to `bytes` because proto3
41            // bytes has no length constraint, so a fixed-size array
42            // would lose its size invariant on the wire — keep it as
43            // `repeated uint32` and let the consumer enforce length.
44            // `byte[]` is signed 8-bit and stays `repeated int32`.
45            let prim = RosPrimitive::parse(p).unwrap_or_else(|| {
46                panic!("[robonix-codegen] proto: primitive '{p}' is not in RosPrimitive::parse")
47            });
48            if field.is_array && field.array_size.is_none() && prim.is_blob_element() {
49                return "bytes".to_string();
50            }
51            proto_primitive_type(p).to_string()
52        }
53        MsgTypeRef::Named { package, name } => {
54            if package == current_package {
55                name.clone()
56            } else {
57                format!("{}.{}", proto_package_name(package), name)
58            }
59        }
60    };
61    if field.is_array {
62        format!("repeated {}", base)
63    } else {
64        base
65    }
66}
67
68/// ROS package name (`prm_base`, `sensor_msgs`) → protobuf package (`robonix.prm_base`, …).
69pub fn proto_package_name(ros_package: &str) -> String {
70    format!("robonix.{}", ros_package)
71}
72
73fn emit_message(out: &mut String, spec: &MsgSpec) -> Result<()> {
74    let _ = writeln!(out, "message {} {{", spec.name);
75    emit_enum(out, spec)?;
76    for (i, field) in spec.fields.iter().enumerate() {
77        let proto_type = proto_field_type(field, &spec.package);
78        let _ = writeln!(out, "  {} {} = {};", proto_type, field.name, i + 1);
79    }
80    let _ = writeln!(out, "}}");
81    Ok(())
82}
83
84/// Emit one nested `MsgNameEnum` for integer constants declared by a `.msg`.
85///
86/// Regular fields keep their original proto scalar/message types; the enum is
87/// a generated namespace for constants, not a replacement for field types.
88fn emit_enum(out: &mut String, spec: &MsgSpec) -> Result<()> {
89    let Some(constants) = proto_enum_constants(&spec.constants) else {
90        return Ok(());
91    };
92    let _ = writeln!(out, "  enum {}Enum {{", spec.name);
93    for constant in constants {
94        let _ = writeln!(out, "    {} = {};", constant.name, constant.value);
95    }
96    let _ = writeln!(out, "  }}");
97    let _ = writeln!(out);
98    Ok(())
99}
100
101/// Return constants in protobuf-valid enum order.
102///
103/// Proto3 requires the first enum value to be zero, while ROS messages may put
104/// a zero-valued constant later in the file. Messages without a zero-valued
105/// constant cannot be emitted as a proto3 enum and are skipped. Duplicate
106/// values are skipped instead of using `allow_alias`.
107fn proto_enum_constants(constants: &[MsgConstant]) -> Option<Vec<&MsgConstant>> {
108    if constants.is_empty()
109        || !constants.iter().any(|c| c.value == 0)
110        || has_duplicate_values(constants)
111    {
112        return None;
113    }
114    let mut ordered = Vec::with_capacity(constants.len());
115    if let Some(first_zero) = constants.iter().find(|c| c.value == 0) {
116        ordered.push(first_zero);
117    }
118    for constant in constants {
119        if constant.value != 0 || !std::ptr::eq(*ordered.first()?, constant) {
120            ordered.push(constant);
121        }
122    }
123    Some(ordered)
124}
125
126fn has_duplicate_values(constants: &[MsgConstant]) -> bool {
127    let mut seen = BTreeSet::new();
128    constants
129        .iter()
130        .any(|constant| !seen.insert(constant.value))
131}
132
133fn emit_srv_messages(out: &mut String, srv: &SrvSpec) -> Result<()> {
134    emit_message(out, &srv.request)?;
135    let _ = writeln!(out);
136    emit_message(out, &srv.response)?;
137    Ok(())
138}
139
140fn import_named_type(imports: &mut BTreeSet<String>, current_package: &str, tr: &MsgTypeRef) {
141    if let MsgTypeRef::Named { package, .. } = tr
142        && package != current_package
143    {
144        imports.insert(package.clone());
145    }
146}
147
148fn collect_imports(
149    specs: &[&MsgSpec],
150    srvs: &[&SrvSpec],
151    current_package: &str,
152) -> BTreeSet<String> {
153    let mut imports = BTreeSet::new();
154    for spec in specs {
155        for field in &spec.fields {
156            import_named_type(&mut imports, current_package, &field.type_ref);
157        }
158    }
159    for srv in srvs {
160        for field in srv.request.fields.iter().chain(srv.response.fields.iter()) {
161            import_named_type(&mut imports, current_package, &field.type_ref);
162        }
163    }
164    imports
165}
166
167pub fn generate(
168    resolver: &MsgResolver,
169    out_dir: &Path,
170    contract_srvs: Option<&BTreeSet<(String, String)>>,
171    verbose: bool,
172) -> Result<()> {
173    fs::create_dir_all(out_dir)?;
174
175    let mut all_packages = BTreeSet::new();
176    for spec in resolver.ordered_specs() {
177        all_packages.insert(spec.package.clone());
178    }
179    if let Some(set) = contract_srvs {
180        for (pkg, _) in set.iter() {
181            all_packages.insert(pkg.clone());
182        }
183    }
184
185    for package in &all_packages {
186        let specs: Vec<_> = resolver
187            .ordered_specs()
188            .into_iter()
189            .filter(|s| &s.package == package)
190            .collect();
191        let srvs: Vec<_> = match contract_srvs {
192            Some(set) => resolver
193                .ordered_srvs()
194                .into_iter()
195                .filter(|s| {
196                    &s.package == package && set.contains(&(s.package.clone(), s.name.clone()))
197                })
198                .collect(),
199            None => Vec::new(),
200        };
201
202        if specs.is_empty() && srvs.is_empty() {
203            continue;
204        }
205
206        let mut out = String::new();
207        let _ = writeln!(out, "// @generated by robonix-codegen --lang proto");
208        let _ = writeln!(out, "// source: ROS IDL package '{}'", package);
209        let _ = writeln!(out, "syntax = \"proto3\";");
210        let _ = writeln!(out);
211        let _ = writeln!(out, "package {};", proto_package_name(package));
212        let _ = writeln!(out);
213
214        let imports = collect_imports(&specs, &srvs, package);
215        for imp in &imports {
216            let _ = writeln!(out, "import \"{}.proto\";", imp);
217        }
218        if !imports.is_empty() {
219            let _ = writeln!(out);
220        }
221
222        for spec in &specs {
223            emit_message(&mut out, spec)?;
224            let _ = writeln!(out);
225        }
226
227        for srv in &srvs {
228            emit_srv_messages(&mut out, srv)?;
229            let _ = writeln!(out);
230        }
231
232        let filename = format!("{}.proto", package);
233        let filepath = out_dir.join(&filename);
234        fs::write(&filepath, &out)
235            .with_context(|| format!("failed to write proto file to '{}'", filepath.display()))?;
236        if verbose {
237            eprintln!(
238                "[robonix-codegen] generated proto for '{}' ({} msgs, {} contract srvs) -> {}",
239                package,
240                specs.len(),
241                srvs.len(),
242                filepath.display()
243            );
244        }
245    }
246
247    Ok(())
248}