use std::collections::BTreeMap;
use std::env;
use std::path::{Path, PathBuf};
fn primitive_rust(ty: &str) -> Option<&'static str> {
Some(match ty {
"uint8" => "u8",
"uint16" => "u16",
"uint32" => "u32",
"uint64" => "u64",
"int8" => "i8",
"int16" => "i16",
"int32" => "i32",
"int64" => "i64",
"float32" => "f32",
"float64" => "f64",
"bool" => "bool",
"char" => "u8",
"byte" => "i8",
"string" => "std::string::String",
_ => return None,
})
}
enum Dim { Scalar, Fixed(usize), Variable }
fn parse_field(line: &str) -> Option<(String, String, bool)> {
let line = line.trim();
if line.is_empty() || line.starts_with('#') { return None; }
let line = line.split('#').next().unwrap_or("").trim();
let mut it = line.split_whitespace();
let ty = it.next()?;
let name = it.next()?;
let name = name.trim();
let (base, dim) = if let Some(stripped) = ty.strip_suffix("[]") {
(stripped.to_string(), Dim::Variable)
} else if let Some(open) = ty.find('[') {
let base = ty[..open].to_string();
let inner = ty[open..].trim_start_matches('[').trim_end_matches(']');
let n: usize = inner.parse().ok()?;
(base, Dim::Fixed(n))
} else {
(ty.to_string(), Dim::Scalar)
};
let big = matches!(dim, Dim::Fixed(n) if n > 32);
let rust = if let Some(p) = primitive_rust(&base) {
match dim {
Dim::Scalar => p.to_string(),
Dim::Fixed(n) => format!("[{}; {n}]", p),
Dim::Variable => format!("Vec<{p}>"),
}
} else {
match dim {
Dim::Scalar => base,
Dim::Fixed(n) => format!("[{base}; {n}]"),
Dim::Variable => format!("Vec<{base}>"),
}
};
Some((rust, name.to_string(), big))
}
fn gen_struct(name: &str, lines: &[String], indent: &str) -> String {
let mut s = format!("{indent}#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]\n{indent}pub struct {name} {{\n");
for line in lines {
let Some((ty, field, big)) = parse_field(line) else { continue };
let field_ident = if is_keyword(&field) {
format!("r#{}", field)
} else {
field
};
let big_attr = if big {
&format!("{indent} #[serde(with = \"serde_big_array::BigArray\")]\n")
} else { "" };
s.push_str(&format!("{big_attr}{indent} pub {field_ident}: {ty},\n"));
}
s.push_str(&format!("{indent}}}\n\n"));
s
}
fn is_keyword(s: &str) -> bool {
matches!(s, "fn" | "type" | "loop" | "match" | "move" | "ref" | "as" | "in"
| "if" | "else" | "for" | "while" | "return" | "use" | "mod" | "struct"
| "enum" | "trait" | "impl" | "let" | "mut" | "const" | "static" | "unsafe")
}
struct Package {
name: String, cpp_ns: String, idl_dir: String, tag: String, types: BTreeMap<String, Vec<String>>,
}
fn pkg_tag(package: &str) -> String {
match package {
"unitree_hg" => "hg".into(),
"unitree_go" => "go".into(),
"std_msgs" => "std".into(),
_ => package.to_string(),
}
}
fn idl_dir(package: &str) -> Option<&str> {
match package {
"unitree_hg" => Some("hg"),
"unitree_go" => Some("go2"),
"std_msgs" => Some("ros2"),
_ => None, }
}
fn parse_packages(msg_root: &Path) -> Vec<Package> {
let mut pkgs = Vec::new();
if !msg_root.exists() { return pkgs; }
let mut dirs: Vec<_> = std::fs::read_dir(msg_root).unwrap()
.filter_map(|e| e.ok())
.filter(|e| e.path().join("msg").is_dir())
.collect();
dirs.sort_by_key(|e| e.file_name());
for e in dirs {
let name = e.file_name().to_string_lossy().to_string();
let Some(idl) = idl_dir(&name) else { continue };
let msg_dir = e.path().join("msg");
let mut types = BTreeMap::new();
let mut entries: Vec<_> = std::fs::read_dir(&msg_dir).unwrap()
.filter_map(|e| e.ok())
.filter(|e| e.path().extension().map_or(false, |x| x == "msg"))
.collect();
entries.sort_by_key(|e| e.file_name());
for m in entries {
let fname = m.file_name().to_string_lossy().to_string();
let tname = fname.trim_end_matches(".msg").to_string();
let lines = std::fs::read_to_string(m.path()).unwrap_or_default()
.lines().map(|s| s.to_string()).collect();
types.insert(tname, lines);
}
pkgs.push(Package {
name: name.clone(),
cpp_ns: format!("{name}::msg::dds_"),
idl_dir: idl.to_string(),
tag: pkg_tag(&name),
types,
});
}
pkgs
}
fn snake_case(name: &str) -> String {
let b = name.as_bytes();
let mut out = String::new();
for (i, &c) in b.iter().enumerate() {
if c.is_ascii_uppercase() {
let prev_lower = i > 0 && b[i - 1].is_ascii_lowercase();
let next_lower = i + 1 < b.len() && b[i + 1].is_ascii_lowercase();
if i > 0 && (prev_lower || next_lower) { out.push('_'); }
out.push((c as char).to_ascii_lowercase());
} else {
out.push(c as char);
}
}
out
}
fn gen_rust_types(pkgs: &[Package]) -> String {
let mut out = String::new();
out.push_str("// ⚡ 自动生成(msg → Rust),改 msg 或 build.rs 后重新编译。\n\n");
for pkg in pkgs {
let parts: Vec<&str> = pkg.cpp_ns.split("::").collect();
let depth = parts.len();
for (i, part) in parts.iter().enumerate() {
out.push_str(&format!("{}pub mod {part} {{\n", " ".repeat(i)));
}
let body_indent = " ".repeat(depth);
for (name, lines) in &pkg.types {
out.push_str(&gen_struct(name, lines, &body_indent));
out.push_str(&format!(
"{body_indent}impl crate::DdsType for {name} {{\n\
{body_indent} fn dds_name() -> &'static str {{ \"{cpp_ns}::{name}_\" }}\n\
{body_indent}}}\n\n",
cpp_ns = pkg.cpp_ns
));
}
for i in (0..depth).rev() {
out.push_str(&format!("{}}}\n", " ".repeat(i)));
}
}
out
}
fn gen_rust_ffi(pkgs: &[Package]) -> String {
let mut out = String::new();
out.push_str("// ⚡ 自动生成(msg → Rust FFI),改 msg 或 build.rs 后重新编译。\n\n");
out.push_str("pub struct ByteHandler {\n");
out.push_str(" cb: Box<dyn Fn(Vec<u8>) + Send>,\n");
out.push_str("}\n");
out.push_str("impl ByteHandler {\n");
out.push_str(" pub fn new(cb: impl Fn(Vec<u8>) + Send + 'static) -> Self { Self { cb: Box::new(cb) } }\n");
out.push_str(" pub(crate) fn on_bytes(&self, bytes: Vec<u8>) { (self.cb)(bytes); }\n");
out.push_str("}\n\n");
out.push_str("#[cxx::bridge(namespace = \"unitree\")]\n");
out.push_str("pub mod ffi {\n");
out.push_str(" extern \"Rust\" {\n");
out.push_str(" type ByteHandler;\n");
out.push_str(" fn on_bytes(self: &ByteHandler, bytes: Vec<u8>);\n");
out.push_str(" }\n");
out.push_str(" unsafe extern \"C++\" {\n");
out.push_str(" include!(\"dds_bridge.h\");\n");
out.push_str(" fn boot_dds(domain_id: i32, network_interface: &str, config_file: &str);\n");
out.push_str(" fn subscribe_any(topic: &str, type_name: &str, handler: Box<ByteHandler>) -> i32;\n");
out.push_str(" fn unsubscribe(id: i32);\n");
for pkg in pkgs {
for name in pkg.types.keys() {
let low = snake_case(name);
out.push_str(&format!(
" type {tag}{name}Publisher;\n\
\x20 fn new_{tag}_{low}_publisher(topic: &str) -> UniquePtr<{tag}{name}Publisher>;\n\
\x20 fn publish_{tag}_{low}_bytes(p: &{tag}{name}Publisher, bytes: Vec<u8>);\n",
tag = pkg.tag
));
}
}
out.push_str(" }\n");
out.push_str("}\n");
out
}
fn gen_rust_topic_impl(pkgs: &[Package]) -> String {
let mut out = String::new();
out.push_str("// ⚡ 自动生成(msg → Rust topic_publish),改 msg 或 build.rs 后重新编译。\n\n");
for pkg in pkgs {
for name in pkg.types.keys() {
let low = snake_case(name);
out.push_str(&format!(
"topic_publish!(crate::gen_types::{cpp_ns}::{name}, {tag}{name}Publisher, new_{tag}_{low}_publisher, publish_{tag}_{low}_bytes);\n",
cpp_ns = pkg.cpp_ns, tag = pkg.tag
));
}
}
out
}
fn gen_topics_h(pkgs: &[Package]) -> String {
let mut out = String::new();
out.push_str("// ⚡ 自动生成(msg → C++),改 msg 或 build.rs 后重新编译。\n#pragma once\n\n");
for pkg in pkgs {
for name in pkg.types.keys() {
out.push_str(&format!("#include <unitree/idl/{idl}/{name}_.hpp>\n", idl = pkg.idl_dir));
}
}
out.push_str("\nnamespace unitree {\n\n");
for pkg in pkgs {
for name in pkg.types.keys() {
let low = snake_case(name);
out.push_str(&format!(
"using {tag}{name}Publisher = Publisher<{cpp_ns}::{name}_>;\n\
std::unique_ptr<{tag}{name}Publisher> new_{tag}_{low}_publisher(rust::Str topic);\n\
void publish_{tag}_{low}_bytes(const {tag}{name}Publisher& pub, rust::Vec<uint8_t> bytes);\n\n",
tag = pkg.tag, cpp_ns = pkg.cpp_ns
));
}
}
out.push_str("std::unique_ptr<SubscriberBase> make_subscriber(\n");
out.push_str(" const std::string& topic, const std::string& type_name,\n");
out.push_str(" std::function<void(rust::Vec<uint8_t>)> cb);\n\n");
out.push_str("} // namespace unitree\n");
out
}
fn gen_topics_inc(pkgs: &[Package]) -> String {
let mut out = String::new();
out.push_str("// ⚡ 自动生成(msg → C++),改 msg 或 build.rs 后重新编译。\n\n");
for pkg in pkgs {
for name in pkg.types.keys() {
let low = snake_case(name);
out.push_str(&format!(
"template class Subscriber<{cpp_ns}::{name}_>;\n\
template class Publisher<{cpp_ns}::{name}_>;\n\
std::unique_ptr<{tag}{name}Publisher> new_{tag}_{low}_publisher(rust::Str t) {{\n\
\x20 return std::make_unique<{tag}{name}Publisher>(std::string(t));\n\
}}\n\
void publish_{tag}_{low}_bytes(const {tag}{name}Publisher& pub, rust::Vec<uint8_t> bytes) {{\n\
\x20 auto msg = bytes_to_dds<{cpp_ns}::{name}_>(to_std_vec(bytes));\n\
\x20 pub.write(msg);\n\
}}\n\n",
cpp_ns = pkg.cpp_ns, tag = pkg.tag
));
}
}
out.push_str("std::unique_ptr<SubscriberBase> make_subscriber(\n");
out.push_str(" const std::string& topic, const std::string& type_name,\n");
out.push_str(" std::function<void(rust::Vec<uint8_t>)> cb) {\n");
for pkg in pkgs {
for name in pkg.types.keys() {
out.push_str(&format!(
" if (type_name == \"{cpp_ns}::{name}_\")\n\
\x20\x20 return std::make_unique<Subscriber<{cpp_ns}::{name}_>>(topic, cb);\n",
cpp_ns = pkg.cpp_ns
));
}
}
out.push_str(" return nullptr;\n}\n");
out
}
fn write_if_changed(path: &Path, content: &str) {
if let Some(parent) = path.parent() {
let _ = std::fs::create_dir_all(parent);
}
if std::fs::read_to_string(path).map(|s| s != content).unwrap_or(true) {
std::fs::write(path, content).unwrap();
println!("cargo:warning=regenerated {}", path.display());
}
}
const SDK2_DEFAULT_URL: &str = "https://github.com/unitreerobotics/unitree_sdk2.git";
const SDK2_INTERNAL_URL: &str = "https://git.unitree.com/unitree/rd/base-environment/unitree_sdk2.git";
fn sdk_cfg(manifest_dir: &Path) -> (String, PathBuf) {
let url = if let Ok(u) = env::var("UNITREE_SDK2_URL") {
u } else {
let source = env::var("UNITREE_SDK2_SOURCE")
.or_else(|_| env::var("CARGO_PKG_METADATA_SDK2_SOURCE"))
.unwrap_or_else(|_| {
if std::env::var_os("CARGO_FEATURE_INTERNAL").is_some() { "internal".to_string() }
else if std::env::var_os("CARGO_FEATURE_CUSTOM").is_some() { "custom".to_string() }
else { "default".to_string() }
});
match source.as_str() {
"internal" => SDK2_INTERNAL_URL.to_string(),
"custom" => env::var("CARGO_PKG_METADATA_SDK2_URL")
.ok()
.filter(|u| !u.is_empty())
.unwrap_or_else(|| panic!("sdk2-source = \"custom\" 但 sdk2-url 未设置,或用 env UNITREE_SDK2_URL 指定")),
_ => SDK2_DEFAULT_URL.to_string(), }
};
let dir = env::var("UNITREE_SDK2_DIR")
.or_else(|_| env::var("CARGO_PKG_METADATA_SDK2_DIR"))
.unwrap_or_else(|_| "tmp_sdk2".to_string());
let dir = PathBuf::from(dir);
let dir = if dir.is_absolute() {
dir
} else if manifest_dir.join("Cargo.toml.orig").exists() {
PathBuf::from(env::var("OUT_DIR").expect("OUT_DIR")).join(dir)
} else {
manifest_dir.join(dir)
};
(url, dir)
}
fn ensure_sdk(sdk_dir: &Path, url: &str, target_arch: &str) {
if sdk_dir.join("include").exists() && sdk_dir.join("lib").exists() {
if sdk_dir.join("lib").join(target_arch).join("libunitree_sdk2.a").exists() {
return; }
println!("cargo:warning=unitree_sdk2 缺少 {target_arch} 预编译库,尝试 cmake 编译...");
let status = std::process::Command::new("bash").arg("-lc").arg(format!(
"cd {} && mkdir -p build && cd build && cmake .. && make -j$(nproc)", sdk_dir.display()
)).status().unwrap_or_else(|e| panic!("无法执行 cmake: {e}"));
if !status.success() {
panic!("cmake 编译 unitree_sdk2 失败。请检查依赖(cmake/g++/libyaml-cpp-dev 等)或用 UNITREE_SDK2_URL 换带预编译库的版本。");
}
return;
}
println!("cargo:warning=克隆 unitree_sdk2: {url}");
if sdk_dir.exists() {
std::fs::remove_dir_all(sdk_dir).expect("清理 SDK 目录失败");
}
let status = std::process::Command::new("git")
.args(["clone", "--depth", "1", url])
.arg(sdk_dir)
.status()
.unwrap_or_else(|e| panic!("无法执行 git clone(需要 git 和网络): {e}"));
if !status.success() {
panic!("git clone 失败,无法获取 unitree_sdk2。请检查网络,或用环境变量 UNITREE_SDK2_URL 指定可访问的仓库。");
}
ensure_sdk(sdk_dir, url, target_arch);
}
fn main() {
let manifest_dir = PathBuf::from(env::var("CARGO_MANIFEST_DIR").unwrap());
let msg_root = PathBuf::from(env::var("UNITREE_MSG_ROOT").unwrap_or_else(|_| {
manifest_dir.join("msgs").to_string_lossy().to_string()
}));
let pkgs = parse_packages(&msg_root);
if pkgs.is_empty() {
panic!("msgs 目录下没有可生成的包(需含 unitree_hg / unitree_go / std_msgs 且 SDK 有对应 IDL): {}", msg_root.display());
}
println!("cargo:warning=生成 {} 个包: {:?}", pkgs.len(), pkgs.iter().map(|p| p.name.as_str()).collect::<Vec<_>>());
write_if_changed(&manifest_dir.join("gen/gen_types.rs"), &gen_rust_types(&pkgs));
write_if_changed(&manifest_dir.join("gen/gen_ffi.rs"), &gen_rust_ffi(&pkgs));
write_if_changed(&manifest_dir.join("gen/gen_topic_impl.rs"), &gen_rust_topic_impl(&pkgs));
write_if_changed(&manifest_dir.join("gen/gen_topics.h"), &gen_topics_h(&pkgs));
write_if_changed(&manifest_dir.join("gen/gen_topics.inc"), &gen_topics_inc(&pkgs));
println!("cargo:rerun-if-changed={}", msg_root.display());
let target_os = env::var("CARGO_CFG_TARGET_OS").unwrap_or_default();
if target_os == "macos" {
println!("cargo:warning=跳过 C++ 编译(macOS 仅 type-check,真机需在 Docker/Linux 编译)");
} else {
let (sdk2_url, sdk_dir) = sdk_cfg(&manifest_dir);
let target_arch = env::var("CARGO_CFG_TARGET_ARCH").unwrap_or_default();
let sdk_arch = match target_arch.as_str() {
"aarch64" | "x86_64" => target_arch,
o => panic!("不支持的目标架构: {o}"),
};
ensure_sdk(&sdk_dir, &sdk2_url, &sdk_arch);
let thirdparty = sdk_dir.join("thirdparty");
let sdk_lib_dir = sdk_dir.join("lib").join(&sdk_arch);
let thirdparty_lib_dir = thirdparty.join("lib").join(&sdk_arch);
assert!(sdk_lib_dir.join("libunitree_sdk2.a").exists(), "libunitree_sdk2.a 未找到");
cxx_build::bridges(&["src/lib.rs", "gen/gen_ffi.rs"])
.file("cpp/dds_bridge.cpp")
.file("cpp/rpc_bridge.cpp")
.include("cpp")
.include(sdk_dir.join("include"))
.include(thirdparty.join("include"))
.include(thirdparty.join("include/ddscxx"))
.include(thirdparty.join("include/iceoryx/v2.0.2"))
.std("c++17")
.compile("unitree_sdk2_rs_bridge");
println!("cargo:rustc-link-search=native={}", sdk_lib_dir.display());
println!("cargo:rustc-link-lib=static=unitree_sdk2");
println!("cargo:rustc-link-search=native={}", thirdparty_lib_dir.display());
println!("cargo:rustc-link-lib=dylib=ddsc");
println!("cargo:rustc-link-lib=dylib=ddscxx");
println!("cargo:rustc-link-arg=-Wl,--disable-new-dtags");
println!("cargo:rustc-link-arg=-Wl,-rpath,{}", thirdparty_lib_dir.display());
println!("cargo:rustc-link-lib=pthread");
}
println!("cargo:rerun-if-changed=src/lib.rs");
println!("cargo:rerun-if-changed=gen/gen_types.rs");
println!("cargo:rerun-if-changed=gen/gen_ffi.rs");
println!("cargo:rerun-if-changed=gen/gen_topic_impl.rs");
println!("cargo:rerun-if-changed=gen/gen_topics.h");
println!("cargo:rerun-if-changed=gen/gen_topics.inc");
println!("cargo:rerun-if-changed=cpp/dds_bridge.h");
println!("cargo:rerun-if-changed=cpp/dds_bridge.cpp");
println!("cargo:rerun-if-changed=cpp/cdr_util.h");
println!("cargo:rerun-if-changed=cpp/rpc_bridge.h");
println!("cargo:rerun-if-changed=cpp/rpc_bridge.cpp");
}