use proc_macro::{Delimiter, TokenStream, TokenTree};
#[derive(Default)]
struct TestOpts {
name: Option<String>,
timeout_time: Option<u64>,
timeout_unit: Option<String>,
skip: bool,
expect_fail: bool,
expect_error: Option<String>,
}
fn strip_quotes(s: &str) -> String {
s.trim_matches('"').to_string()
}
fn parse_test_opts(attr: TokenStream) -> Result<TestOpts, String> {
let mut opts = TestOpts::default();
let mut iter = attr.into_iter().peekable();
while let Some(tt) = iter.next() {
let key = match &tt {
TokenTree::Ident(i) => i.to_string(),
TokenTree::Punct(p) if p.as_char() == ',' => continue,
other => return Err(format!("unexpected token in #[rustdv::test(...)]: {other}")),
};
let mut value: Option<String> = None;
if let Some(TokenTree::Punct(p)) = iter.peek() {
if p.as_char() == '=' {
iter.next(); match iter.next() {
Some(TokenTree::Literal(l)) => value = Some(l.to_string()),
Some(TokenTree::Ident(i)) => value = Some(i.to_string()),
other => return Err(format!("expected value after '{key} =', got {other:?}")),
}
}
}
match key.as_str() {
"name" => opts.name = value.map(|v| strip_quotes(&v)),
"timeout_time" => {
let v = value.ok_or("timeout_time needs a value")?;
opts.timeout_time =
Some(v.parse::<u64>().map_err(|_| format!("bad timeout_time '{v}'"))?);
}
"timeout_unit" => opts.timeout_unit = value.map(|v| strip_quotes(&v)),
"skip" => opts.skip = value.map(|v| v == "true").unwrap_or(true),
"expect_fail" => opts.expect_fail = value.map(|v| v == "true").unwrap_or(true),
"expect_error" => {
let v = value.ok_or("expect_error needs a value, e.g. expect_error = \"config_not_found\"")?;
opts.expect_error = Some(strip_quotes(&v));
}
other => return Err(format!("unknown #[rustdv::test] option '{other}'")),
}
}
Ok(opts)
}
#[derive(Copy, Clone, PartialEq, Eq)]
enum TestForm {
Function,
Component,
}
fn find_item(item: &TokenStream) -> Option<(TestForm, String)> {
let mut form: Option<TestForm> = None;
for tt in item.clone() {
if let TokenTree::Ident(i) = tt {
let s = i.to_string();
if let Some(f) = form {
return Some((f, s));
}
form = match s.as_str() {
"fn" => Some(TestForm::Function),
"struct" | "type" => Some(TestForm::Component),
_ => None,
};
}
}
None
}
fn compile_error(msg: &str) -> TokenStream {
format!("compile_error!({msg:?});").parse().unwrap()
}
#[proc_macro_attribute]
pub fn test(attr: TokenStream, item: TokenStream) -> TokenStream {
let opts = match parse_test_opts(attr) {
Ok(o) => o,
Err(e) => return compile_error(&e),
};
let Some((form, item_name)) = find_item(&item) else {
return compile_error(
"#[rustdv::test] must be applied to an async fn, a struct, or a type alias",
);
};
let fn_name = item_name;
let test_name = opts.name.unwrap_or_else(|| fn_name.clone());
let timeout = match (opts.timeout_time, opts.timeout_unit) {
(Some(t), Some(u)) => format!("::core::option::Option::Some(({t}u64, \"{u}\"))"),
(Some(t), None) => format!("::core::option::Option::Some(({t}u64, \"ns\"))"),
_ => "::core::option::Option::None".to_string(),
};
let skip = opts.skip;
let expect_fail = opts.expect_fail;
let expect_error = match &opts.expect_error {
Some(k) => format!("::core::option::Option::Some(\"{k}\")"),
None => "::core::option::Option::None".to_string(),
};
let body = match form {
TestForm::Function => format!("::std::boxed::Box::pin({fn_name}(ctx))"),
TestForm::Component => format!(
r#"::std::boxed::Box::pin(async move {{
let mut __ctx = ctx;
let mut __test = <{fn_name} as ::core::default::Default>::default();
::rustdv::run_component_test(&mut __test, &mut __ctx).await
}})"#
),
};
let reg = format!(
r#"
const _: () = {{
fn __rustdv_shim(
ctx: ::rustdv::RustdvCtx,
) -> ::std::pin::Pin<::std::boxed::Box<
dyn ::std::future::Future<Output = ::core::result::Result<(), ::rustdv::TestError>>,
>> {{
{body}
}}
#[used]
#[cfg_attr(not(target_vendor = "apple"), link_section = "rustdv_tests")]
#[cfg_attr(target_vendor = "apple", link_section = "__DATA,rustdv_tests")]
static __RUSTDV_TEST_REG: &'static ::rustdv::TestRegistration = &::rustdv::TestRegistration {{
name: "{test_name}",
module: ::core::module_path!(),
file: ::core::file!(),
line: ::core::line!(),
run: __rustdv_shim,
timeout: {timeout},
skip: {skip},
expect_fail: {expect_fail},
expect_error: {expect_error},
}};
}};
"#
);
let mut out = item;
out.extend(reg.parse::<TokenStream>().expect("rustdv-macros: generated code failed to parse"));
out
}
struct Field {
name: String,
ty: String,
is_child: bool,
port: Option<String>,
}
fn parse_struct(input: TokenStream) -> Result<(String, String, String, Vec<Field>), String> {
let mut iter = input.into_iter().peekable();
let mut struct_name: Option<String> = None;
while let Some(tt) = iter.next() {
if let TokenTree::Ident(i) = &tt {
if i.to_string() == "struct" {
match iter.next() {
Some(TokenTree::Ident(n)) => {
struct_name = Some(n.to_string());
break;
}
_ => return Err("expected struct name".into()),
}
}
}
}
let name = struct_name.ok_or("#[derive(Component)] supports only structs")?;
let mut fields_group = None;
let mut unit_struct = false;
let mut generics_tokens: Vec<TokenTree> = Vec::new();
let mut depth = 0i32;
for tt in iter {
match &tt {
TokenTree::Group(g) if g.delimiter() == Delimiter::Brace && depth == 0 => {
fields_group = Some(g.clone());
break;
}
TokenTree::Punct(p) if p.as_char() == ';' && depth == 0 => {
unit_struct = true;
break;
}
TokenTree::Punct(p) if p.as_char() == '<' => {
depth += 1;
generics_tokens.push(tt.clone());
continue;
}
TokenTree::Punct(p) if p.as_char() == '>' => {
depth -= 1;
generics_tokens.push(tt.clone());
continue;
}
_ => {}
}
if depth > 0 {
generics_tokens.push(tt.clone());
}
}
let impl_generics: String = {
let mut out = String::new();
for t in &generics_tokens {
let text = t.to_string();
if !out.is_empty() && !out.ends_with('\'') {
out.push(' ');
}
out.push_str(&text);
}
out
};
let type_params = {
let mut params: Vec<String> = Vec::new();
let mut d = 0i32;
let mut take_next_ident = true;
let mut lifetime = false;
for t in &generics_tokens {
match t {
TokenTree::Punct(p) if p.as_char() == '<' => d += 1,
TokenTree::Punct(p) if p.as_char() == '>' => d -= 1,
TokenTree::Punct(p) if p.as_char() == ',' && d == 1 => take_next_ident = true,
TokenTree::Punct(p) if p.as_char() == '\'' && d == 1 && take_next_ident => {
lifetime = true
}
TokenTree::Ident(i) if d == 1 && take_next_ident => {
let word = i.to_string();
if word == "const" {
continue; }
params.push(if lifetime { format!("'{word}") } else { word });
take_next_ident = false;
lifetime = false;
}
_ => {}
}
}
if params.is_empty() { String::new() } else { format!("< {} >", params.join(" , ")) }
};
if unit_struct {
return Ok((name, impl_generics, type_params, Vec::new()));
}
let group = fields_group
.ok_or("#[derive(Component)] requires named fields, or a unit struct")?;
let mut fields = Vec::new();
let mut pending_child = false;
let mut pending_port: Option<String> = None;
let mut current: Vec<TokenTree> = Vec::new();
let mut toks = group.stream().into_iter().peekable();
let mut angle_depth = 0i32;
while let Some(tt) = toks.next() {
match &tt {
TokenTree::Punct(p) if p.as_char() == '<' => angle_depth += 1,
TokenTree::Punct(p) if p.as_char() == '>' => angle_depth -= 1,
TokenTree::Punct(p) if p.as_char() == '#' => {
if let Some(TokenTree::Group(g)) = toks.peek() {
if g.delimiter() == Delimiter::Bracket {
let text = g.stream().to_string();
if text.starts_with("component") {
pending_child = true;
}
if text.starts_with("port") {
pending_port = attr_arg(&text);
}
toks.next(); continue;
}
}
}
TokenTree::Punct(p) if p.as_char() == ',' && angle_depth == 0 => {
if !current.is_empty() {
fields.push(make_field(¤t, pending_child, pending_port.take())?);
current.clear();
pending_child = false;
}
continue;
}
_ => {}
}
current.push(tt);
}
if !current.is_empty() {
fields.push(make_field(¤t, pending_child, pending_port.take())?);
}
Ok((name, impl_generics, type_params, fields))
}
fn attr_arg(text: &str) -> Option<String> {
let open = text.find('(')?;
let close = text.rfind(')')?;
let arg = text[open + 1..close].trim();
if arg.is_empty() {
None
} else {
Some(arg.to_string())
}
}
fn make_field(
tokens: &[TokenTree],
is_child: bool,
port: Option<String>,
) -> Result<Field, String> {
let mut name = None;
let mut colon_at = None;
for (i, tt) in tokens.iter().enumerate() {
if let TokenTree::Punct(p) = tt {
if p.as_char() == ':' && colon_at.is_none() {
colon_at = Some(i);
break;
}
}
}
let colon = colon_at.ok_or("field without ':' (tuple structs unsupported)")?;
for tt in tokens[..colon].iter().rev() {
if let TokenTree::Ident(i) = tt {
name = Some(i.to_string());
break;
}
}
let name = name.ok_or("could not find field name")?;
let ty: String = tokens[colon + 1..].iter().map(|t| t.to_string()).collect::<Vec<_>>().join(" ");
Ok(Field { name, ty, is_child, port })
}
#[proc_macro_derive(Component, attributes(component, port))]
pub fn derive_component(input: TokenStream) -> TokenStream {
let (name, impl_generics, type_params, fields) = match parse_struct(input) {
Ok(v) => v,
Err(e) => return compile_error(&e),
};
let mut visits = String::new();
let mut resolves = String::new();
let mut takes = String::new();
let mut restores = String::new();
for f in fields.iter().filter(|f| f.is_child) {
let fname = &f.name;
let ty = f.ty.trim_start();
if ty.starts_with("RustdvComp") {
visits.push_str(&format!(
"if let ::core::option::Option::Some(__c) = self.{fname}.as_node_mut() {{ __out.push((::std::string::String::from(\"{fname}\"), __c)); }}\n"
));
resolves.push_str(&format!("self.{fname}.resolve(__ctx, \"{fname}\");\n"));
takes.push_str(&format!(
"if let ::core::option::Option::Some(__c) = self.{fname}.take_node() {{ __out.push((::std::string::String::from(\"{fname}\"), __c)); }}\n"
));
restores.push_str(&format!(
"if __name == \"{fname}\" {{ self.{fname}.put_node(__node); continue; }}\n"
));
} else if ty.starts_with("Option") {
visits.push_str(&format!(
"if let ::core::option::Option::Some(__c) = &mut self.{fname} {{ __out.push((::std::string::String::from(\"{fname}\"), __c as &mut (dyn ::rustdv::ComponentNode + 'static))); }}\n"
));
} else if ty.starts_with("Vec") {
visits.push_str(&format!(
"for (__i, __c) in self.{fname}.iter_mut().enumerate() {{ __out.push((::std::format!(\"{fname}[{{}}]\", __i), __c as &mut (dyn ::rustdv::ComponentNode + 'static))); }}\n"
));
} else {
visits.push_str(&format!(
"__out.push((::std::string::String::from(\"{fname}\"), &mut self.{fname} as &mut (dyn ::rustdv::ComponentNode + 'static)));\n"
));
}
}
let mut port_arms = String::new();
let mut port_items = String::new();
let mut port_consts = String::new();
for f in fields.iter() {
let Some(kind) = f.port.as_deref() else { continue };
if !matches!(kind, "put" | "get" | "peek" | "publish" | "subscribe" | "seq_item") {
return compile_error(&format!(
"#[port({kind})]: expected put, get, peek, peek, publish, subscribe or seq_item"
));
}
let required = !matches!(kind, "publish" | "subscribe");
let fname = &f.name;
let ty = f.ty.trim();
let konst = fname.to_uppercase();
port_arms.push_str(&format!(
"\"{fname}\" => ::core::option::Option::Some(::rustdv::PortField::slot_any(&self.{fname})),\n "
));
port_items.push_str(&format!(
"::rustdv::PortInfo {{ name: \"{fname}\", kind: \"{kind}\", required: {required}, connected: ::rustdv::PortField::bound(&self.{fname}) }},\n "
));
port_consts.push_str(&format!(
" /// The `{fname}` port, for `connect`.\n pub const {konst}: ::rustdv::PortName<<{ty} as ::rustdv::PortField>::Iface> = ::rustdv::PortName::new(\"{fname}\");\n"
));
}
let port_impl = if port_arms.is_empty() {
String::new()
} else {
format!(
" fn port_slot(&self, __name: &str) -> ::core::option::Option<::std::rc::Rc<dyn ::std::any::Any>> {{\n \
match __name {{\n {port_arms}_ => ::core::option::Option::None,\n }}\n }}\n\
\n fn port_infos(&self) -> ::std::vec::Vec<::rustdv::PortInfo> {{\n \
::std::vec![\n {port_items}]\n }}\n"
)
};
let owner_impl = format!(
r#"
impl {impl_generics} ::rustdv::PortOwner for {name} {type_params} {{
fn owner_port_slot(&self, __name: &str) -> ::core::option::Option<::std::rc::Rc<dyn ::std::any::Any>> {{
::rustdv::ComponentNode::port_slot(self, __name)
}}
fn owner_label(&self) -> &'static str {{ "{name}" }}
}}
"#
);
let const_impl = if port_consts.is_empty() {
String::new()
} else {
format!("\nimpl {impl_generics} {name} {type_params} {{\n{port_consts}}}\n")
};
let resolve_impl = if resolves.is_empty() {
String::new()
} else {
format!(
" fn resolve_children(&mut self, __ctx: &::rustdv::RustdvCtx) {{\n {resolves} }}\n"
)
};
let take_impl = if takes.is_empty() {
String::new()
} else {
format!(
" fn take_children(&mut self) -> ::std::vec::Vec<(::std::string::String, ::std::boxed::Box<dyn ::rustdv::ComponentNode>)> {{\n \
let mut __out: ::std::vec::Vec<(::std::string::String, ::std::boxed::Box<dyn ::rustdv::ComponentNode>)> = ::std::vec::Vec::new();\n \
{takes} __out\n }}\n\
\n fn restore_children(&mut self, __taken: ::std::vec::Vec<(::std::string::String, ::std::boxed::Box<dyn ::rustdv::ComponentNode>)>) {{\n \
for (__name, __node) in __taken {{\n \
let __name: &str = &__name;\n {restores} }}\n }}\n"
)
};
let registration = if type_params.is_empty() {
format!(
r#"
const _: () = {{
fn __rustdv_comp_name() -> &'static str {{ "{name}" }}
fn __rustdv_comp_make() -> ::std::boxed::Box<dyn ::rustdv::ComponentNode> {{
::std::boxed::Box::new(<{name} as ::core::default::Default>::default())
}}
#[used]
#[cfg_attr(not(target_vendor = "apple"), link_section = "rustdv_comps")]
#[cfg_attr(target_vendor = "apple", link_section = "__DATA,rustdv_comps")]
static __RUSTDV_COMP_REG: &::rustdv::ComponentReg = &::rustdv::ComponentReg {{
name: __rustdv_comp_name,
make: __rustdv_comp_make,
}};
}};
"#
)
} else {
String::new()
};
let out = format!(
r#"
impl {impl_generics} ::rustdv::ComponentNode for {name} {type_params} {{
fn node_name(&self) -> &'static str {{ "{name}" }}
fn children_mut(&mut self) -> ::std::vec::Vec<(::std::string::String, &mut (dyn ::rustdv::ComponentNode + 'static))> {{
let mut __out: ::std::vec::Vec<(::std::string::String, &mut (dyn ::rustdv::ComponentNode + 'static))> = ::std::vec::Vec::new();
{visits}
__out
}}
{port_impl}{resolve_impl}{take_impl}}}
{owner_impl}{const_impl}{registration}"#
);
out.parse().expect("rustdv-macros: generated ComponentNode impl failed to parse")
}