Skip to main content

miden_node_tracing_macro/
lib.rs

1use std::collections::BTreeSet;
2
3use proc_macro::TokenStream;
4use proc_macro2::{Delimiter, Group, TokenStream as TokenStream2, TokenTree};
5use quote::{ToTokens, quote};
6use syn::parse::{Parse, ParseStream};
7use syn::punctuated::Punctuated;
8use syn::token::Dot;
9use syn::visit::Visit;
10use syn::{Block, Expr, Ident, ItemFn, Macro, Result, Token, parse_macro_input, parse_quote};
11
12const ALLOWED_FIELD_NAMES: &[&str] = &[
13    "account.id",
14    "account.id.network_prefix",
15    "account.ids",
16    "account.ids.count",
17    "account.updated",
18    "batch.id",
19    "batch.account_updates.count",
20    "batch.expires_at",
21    "batch.expiration_height",
22    "batch.input_notes.count",
23    "batch.output_notes.count",
24    "batch.reference_block.commitment",
25    "batch.reference_block.number",
26    "block.batch.ids",
27    "block.batches.count",
28    "block.batches.output_notes.count",
29    "block.commitment",
30    "block.commitments.account",
31    "block.commitments.chain",
32    "block.commitments.kernel",
33    "block.commitments.note",
34    "block.commitments.nullifier",
35    "block.commitments.transaction",
36    "block.erased_note_proofs.count",
37    "block.erased_notes.count",
38    "block.from",
39    "block.nullifiers.count",
40    "block.number",
41    "block.output_notes.count",
42    "block.prev_block_commitment",
43    "block.protocol.version",
44    "block.size",
45    "block.sub_commitment",
46    "block.timestamp",
47    "block.transactions.ids",
48    "block.transactions.count",
49    "block.updated_accounts.count",
50    "block_range.from",
51    "block_range.to",
52    "current_client_block_height",
53    "cutoff_block",
54    "db.account_state_forest.size",
55    "db.account_tree.size",
56    "db.block_store.size",
57    "db.nullifier_tree.size",
58    "db.sqlite.size",
59    "db.sqlite.wal.size",
60    "dice_roll",
61    "failure_rate",
62    "finality_level",
63    "inputs_size",
64    "mempool.accounts",
65    "mempool.batches.proposed",
66    "mempool.batches.proven",
67    "mempool.nullifiers",
68    "mempool.output_notes",
69    "mempool.transactions.unbatched",
70    "mempool.transactions.uncommitted",
71    "note.id",
72    "notes.count",
73    "nullifiers",
74    "path",
75    "port",
76    "prefix_len",
77    "prefixes",
78    "proof_size",
79    "prover",
80    "prover.kind",
81    "reference_block.number",
82    "request.kind",
83    "script.root",
84    "transaction.id",
85    "transaction.expires_at",
86    "transaction.input_notes.count",
87    "transaction.output_notes.count",
88    "transaction.reference_block.commitment",
89    "transaction.reference_block.number",
90    "tip.number",
91    "transactions.count",
92    "transactions.ids",
93    "transactions.input_notes.count",
94    "transactions.output_notes.count",
95    "transactions.unauthenticated_notes.count",
96    "workers.active",
97    "workers.capacity",
98    "workers.count",
99];
100
101#[proc_macro_attribute]
102pub fn miden_instrument(attr: TokenStream, item: TokenStream) -> TokenStream {
103    let attr = TokenStream2::from(attr);
104    let mut function = parse_macro_input!(item as ItemFn);
105    let fields = collect_recorded_fields(&function);
106    let args = match merge_inferred_fields(attr, &fields) {
107        Ok(args) => args,
108        Err(error) => return error.into_compile_error().into(),
109    };
110    let statements = &function.block.stmts;
111    let block: Block = parse_quote! {{
112        #[allow(unused_macros)]
113        macro_rules! __miden_span_record_must_be_used_within_miden_instrument {
114            () => {};
115        }
116
117        #(#statements)*
118    }};
119    *function.block = block;
120
121    let expanded = quote! {
122        #[::tracing::instrument(#args)]
123        #function
124    };
125
126    expanded.into()
127}
128
129fn merge_inferred_fields(attr: TokenStream2, fields: &[FieldPath]) -> Result<TokenStream2> {
130    validate_explicit_fields(&attr)?;
131
132    if fields.is_empty() {
133        return Ok(attr);
134    }
135
136    let inferred_fields = quote! { #(#fields = ::tracing::field::Empty),* };
137    if attr.is_empty() {
138        return Ok(quote! { fields(#inferred_fields) });
139    }
140
141    let mut merged_existing_fields = false;
142    let args = split_top_level_args(attr)
143        .into_iter()
144        .map(|arg| {
145            if let Some(group) = fields_group(&arg) {
146                merged_existing_fields = true;
147                let existing_fields = group.stream();
148                let merged_fields = if existing_fields.is_empty() {
149                    inferred_fields.clone()
150                } else if ends_with_comma(&existing_fields) {
151                    quote! { #existing_fields #inferred_fields }
152                } else {
153                    quote! { #existing_fields, #inferred_fields }
154                };
155                let mut merged_group = Group::new(Delimiter::Parenthesis, merged_fields);
156                merged_group.set_span(group.span());
157                quote! { fields #merged_group }
158            } else {
159                arg
160            }
161        })
162        .collect::<Vec<_>>();
163
164    if merged_existing_fields {
165        Ok(quote! { #(#args),* })
166    } else {
167        Ok(quote! { #(#args,)* fields(#inferred_fields) })
168    }
169}
170
171fn validate_explicit_fields(attr: &TokenStream2) -> Result<()> {
172    for arg in split_top_level_args(attr.clone()) {
173        if let Some(group) = fields_group(&arg) {
174            syn::parse2::<InstrumentFields>(group.stream())?;
175        }
176    }
177
178    Ok(())
179}
180
181fn split_top_level_args(tokens: TokenStream2) -> Vec<TokenStream2> {
182    let mut args = Vec::new();
183    let mut current = TokenStream2::new();
184
185    for token in tokens {
186        match &token {
187            TokenTree::Punct(punct) if punct.as_char() == ',' => {
188                args.push(current);
189                current = TokenStream2::new();
190            },
191            _ => current.extend([token]),
192        }
193    }
194
195    if !current.is_empty() {
196        args.push(current);
197    }
198
199    args
200}
201
202fn fields_group(arg: &TokenStream2) -> Option<Group> {
203    let mut tokens = arg.clone().into_iter();
204    let Some(TokenTree::Ident(ident)) = tokens.next() else {
205        return None;
206    };
207    if ident != "fields" {
208        return None;
209    }
210
211    let Some(TokenTree::Group(group)) = tokens.next() else {
212        return None;
213    };
214    if group.delimiter() != Delimiter::Parenthesis || tokens.next().is_some() {
215        return None;
216    }
217
218    Some(group)
219}
220
221fn ends_with_comma(tokens: &TokenStream2) -> bool {
222    matches!(
223        tokens.clone().into_iter().last(),
224        Some(TokenTree::Punct(punct)) if punct.as_char() == ','
225    )
226}
227
228#[proc_macro]
229pub fn miden_span_record(input: TokenStream) -> TokenStream {
230    let records = parse_macro_input!(input as RecordFields);
231    let records = records.fields.into_iter().map(|field| {
232        let name = field.path.name();
233        let value = field
234            .value
235            .expect("record fields are parsed with required values")
236            .value_tokens();
237
238        quote! {
239            ::tracing::Span::current().record(#name, #value);
240        }
241    });
242
243    quote! {
244        __miden_span_record_must_be_used_within_miden_instrument!();
245        #(#records)*
246    }
247    .into()
248}
249
250fn validate_field_name(path: &FieldPath) -> Result<()> {
251    let name = path.name();
252
253    if ALLOWED_FIELD_NAMES.contains(&name.as_str()) {
254        Ok(())
255    } else {
256        Err(syn::Error::new_spanned(
257            path,
258            format!(
259                "unsupported tracing field `{name}`; use one of: {}",
260                ALLOWED_FIELD_NAMES.join(", "),
261            ),
262        ))
263    }
264}
265
266fn collect_recorded_fields(function: &ItemFn) -> Vec<FieldPath> {
267    let mut visitor = MacroVisitor::default();
268    visitor.visit_block(&function.block);
269
270    let mut names = BTreeSet::new();
271    visitor.fields.into_iter().filter(|field| names.insert(field.name())).collect()
272}
273
274#[derive(Default)]
275struct MacroVisitor {
276    fields: Vec<FieldPath>,
277}
278
279impl<'ast> Visit<'ast> for MacroVisitor {
280    fn visit_macro(&mut self, mac: &'ast Macro) {
281        if mac
282            .path
283            .segments
284            .last()
285            .is_some_and(|segment| segment.ident == "miden_span_record")
286        {
287            if let Ok(records) = syn::parse2::<RecordFields>(mac.tokens.clone()) {
288                self.fields.extend(records.fields.into_iter().map(|field| field.path));
289            }
290        }
291
292        syn::visit::visit_macro(self, mac);
293    }
294}
295
296type InstrumentFields = Fields<false>;
297type RecordFields = Fields<true>;
298
299struct Fields<const VALUE_REQUIRED: bool> {
300    fields: Punctuated<RecordField, Token![,]>,
301}
302
303impl<const VALUE_REQUIRED: bool> Parse for Fields<VALUE_REQUIRED> {
304    fn parse(input: ParseStream<'_>) -> Result<Self> {
305        Ok(Self {
306            fields: Punctuated::parse_terminated_with(input, |input| {
307                RecordField::parse(input, VALUE_REQUIRED)
308            })?,
309        })
310    }
311}
312
313struct RecordField {
314    path: FieldPath,
315    value: Option<RecordValue>,
316}
317
318impl RecordField {
319    fn parse(input: ParseStream<'_>, value_required: bool) -> Result<Self> {
320        let shorthand_formatter = if value_required {
321            None
322        } else {
323            Formatter::parse_optional(input)?
324        };
325        let path = input.parse()?;
326        validate_field_name(&path)?;
327        let value = if value_required || shorthand_formatter.is_none() && input.peek(Token![=]) {
328            input.parse::<Token![=]>()?;
329            Some(input.parse()?)
330        } else {
331            None
332        };
333
334        Ok(Self { path, value })
335    }
336}
337
338struct FieldPath {
339    first: Ident,
340    rest: Vec<(Dot, Ident)>,
341}
342
343impl FieldPath {
344    fn name(&self) -> String {
345        std::iter::once(&self.first)
346            .chain(self.rest.iter().map(|(_, ident)| ident))
347            .map(ToString::to_string)
348            .collect::<Vec<_>>()
349            .join(".")
350    }
351}
352
353impl Parse for FieldPath {
354    fn parse(input: ParseStream<'_>) -> Result<Self> {
355        let first = input.parse()?;
356        let mut rest = Vec::new();
357
358        while input.peek(Token![.]) {
359            rest.push((input.parse()?, input.parse()?));
360        }
361
362        Ok(Self { first, rest })
363    }
364}
365
366impl ToTokens for FieldPath {
367    fn to_tokens(&self, tokens: &mut TokenStream2) {
368        self.first.to_tokens(tokens);
369        for (dot, ident) in &self.rest {
370            dot.to_tokens(tokens);
371            ident.to_tokens(tokens);
372        }
373    }
374}
375
376struct RecordValue {
377    formatter: Formatter,
378    expr: Expr,
379}
380
381impl RecordValue {
382    fn value_tokens(&self) -> TokenStream2 {
383        let expr = &self.expr;
384
385        match self.formatter {
386            Formatter::Display => quote! { &::tracing::field::display(#expr) },
387            Formatter::Debug => quote! { &::tracing::field::debug(#expr) },
388            Formatter::Plain => quote! { &#expr },
389        }
390    }
391}
392
393impl Parse for RecordValue {
394    fn parse(input: ParseStream<'_>) -> Result<Self> {
395        let formatter = Formatter::parse_optional(input)?.unwrap_or(Formatter::Plain);
396        let expr = input.parse()?;
397
398        Ok(Self { formatter, expr })
399    }
400}
401
402enum Formatter {
403    Display,
404    Debug,
405    Plain,
406}
407
408impl Formatter {
409    fn parse_optional(input: ParseStream<'_>) -> Result<Option<Self>> {
410        if input.peek(Token![%]) {
411            input.parse::<Token![%]>()?;
412            Ok(Some(Self::Display))
413        } else if input.peek(Token![?]) {
414            input.parse::<Token![?]>()?;
415            Ok(Some(Self::Debug))
416        } else {
417            Ok(None)
418        }
419    }
420}