robonix_codegen/codegen/
proto_gen.rs1use 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 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 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
68pub 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
84fn 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
101fn 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}