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    "snapshot.block_num",
85    "snapshot.lifetime_ms",
86    "snapshots.live",
87    "transaction.id",
88    "transaction.expires_at",
89    "transaction.input_notes.count",
90    "transaction.output_notes.count",
91    "transaction.reference_block.commitment",
92    "transaction.reference_block.number",
93    "tip.number",
94    "transactions.count",
95    "transactions.ids",
96    "transactions.input_notes.count",
97    "transactions.output_notes.count",
98    "transactions.unauthenticated_notes.count",
99    "workers.active",
100    "workers.capacity",
101    "workers.count",
102];
103
104#[proc_macro_attribute]
105pub fn miden_instrument(attr: TokenStream, item: TokenStream) -> TokenStream {
106    let attr = TokenStream2::from(attr);
107    let mut function = parse_macro_input!(item as ItemFn);
108    let fields = collect_recorded_fields(&function);
109    let args = match merge_inferred_fields(attr, &fields) {
110        Ok(args) => args,
111        Err(error) => return error.into_compile_error().into(),
112    };
113    let statements = &function.block.stmts;
114    let block: Block = parse_quote! {{
115        #[allow(unused_macros)]
116        macro_rules! __miden_span_record_must_be_used_within_miden_instrument {
117            () => {};
118        }
119
120        #(#statements)*
121    }};
122    *function.block = block;
123
124    let expanded = quote! {
125        #[::tracing::instrument(#args)]
126        #function
127    };
128
129    expanded.into()
130}
131
132fn merge_inferred_fields(attr: TokenStream2, fields: &[FieldPath]) -> Result<TokenStream2> {
133    validate_explicit_fields(&attr)?;
134
135    let mut args = split_top_level_args(attr);
136    reject_skip_directives(&args)?;
137
138    // Function arguments often contain large or sensitive values. Always skip them so spans only
139    // contain fields explicitly declared by the caller or inferred from `miden_span_record!`.
140    args.push(quote! { skip_all });
141
142    if fields.is_empty() {
143        return Ok(quote! { #(#args),* });
144    }
145
146    let inferred_fields = quote! { #(#fields = ::tracing::field::Empty),* };
147    let mut merged_existing_fields = false;
148    let args = args
149        .into_iter()
150        .map(|arg| {
151            if let Some(group) = fields_group(&arg) {
152                merged_existing_fields = true;
153                let existing_fields = group.stream();
154                let merged_fields = if existing_fields.is_empty() {
155                    inferred_fields.clone()
156                } else if ends_with_comma(&existing_fields) {
157                    quote! { #existing_fields #inferred_fields }
158                } else {
159                    quote! { #existing_fields, #inferred_fields }
160                };
161                let mut merged_group = Group::new(Delimiter::Parenthesis, merged_fields);
162                merged_group.set_span(group.span());
163                quote! { fields #merged_group }
164            } else {
165                arg
166            }
167        })
168        .collect::<Vec<_>>();
169
170    if merged_existing_fields {
171        Ok(quote! { #(#args),* })
172    } else {
173        Ok(quote! { #(#args,)* fields(#inferred_fields) })
174    }
175}
176
177fn reject_skip_directives(args: &[TokenStream2]) -> Result<()> {
178    for arg in args {
179        let Some(TokenTree::Ident(ident)) = arg.clone().into_iter().next() else {
180            continue;
181        };
182        if ident == "skip" || ident == "skip_all" {
183            return Err(syn::Error::new_spanned(
184                arg,
185                format!(
186                    "`{ident}` is not supported by `miden_instrument`; function arguments are \
187                     always skipped, record fields explicitly with `fields(...)`"
188                ),
189            ));
190        }
191    }
192
193    Ok(())
194}
195
196fn validate_explicit_fields(attr: &TokenStream2) -> Result<()> {
197    for arg in split_top_level_args(attr.clone()) {
198        if let Some(group) = fields_group(&arg) {
199            syn::parse2::<InstrumentFields>(group.stream())?;
200        }
201    }
202
203    Ok(())
204}
205
206fn split_top_level_args(tokens: TokenStream2) -> Vec<TokenStream2> {
207    let mut args = Vec::new();
208    let mut current = TokenStream2::new();
209
210    for token in tokens {
211        match &token {
212            TokenTree::Punct(punct) if punct.as_char() == ',' => {
213                args.push(current);
214                current = TokenStream2::new();
215            },
216            _ => current.extend([token]),
217        }
218    }
219
220    if !current.is_empty() {
221        args.push(current);
222    }
223
224    args
225}
226
227fn fields_group(arg: &TokenStream2) -> Option<Group> {
228    let mut tokens = arg.clone().into_iter();
229    let Some(TokenTree::Ident(ident)) = tokens.next() else {
230        return None;
231    };
232    if ident != "fields" {
233        return None;
234    }
235
236    let Some(TokenTree::Group(group)) = tokens.next() else {
237        return None;
238    };
239    if group.delimiter() != Delimiter::Parenthesis || tokens.next().is_some() {
240        return None;
241    }
242
243    Some(group)
244}
245
246fn ends_with_comma(tokens: &TokenStream2) -> bool {
247    matches!(
248        tokens.clone().into_iter().last(),
249        Some(TokenTree::Punct(punct)) if punct.as_char() == ','
250    )
251}
252
253#[proc_macro]
254pub fn miden_span_record(input: TokenStream) -> TokenStream {
255    let records = parse_macro_input!(input as RecordFields);
256    let records = records.fields.into_iter().map(|field| {
257        let name = field.path.name();
258        let value = field
259            .value
260            .expect("record fields are parsed with required values")
261            .value_tokens();
262
263        quote! {
264            ::tracing::Span::current().record(#name, #value);
265        }
266    });
267
268    quote! {
269        __miden_span_record_must_be_used_within_miden_instrument!();
270        #(#records)*
271    }
272    .into()
273}
274
275fn validate_field_name(path: &FieldPath) -> Result<()> {
276    let name = path.name();
277
278    if ALLOWED_FIELD_NAMES.contains(&name.as_str()) {
279        Ok(())
280    } else {
281        Err(syn::Error::new_spanned(
282            path,
283            format!(
284                "unsupported tracing field `{name}`; use one of: {}",
285                ALLOWED_FIELD_NAMES.join(", "),
286            ),
287        ))
288    }
289}
290
291fn collect_recorded_fields(function: &ItemFn) -> Vec<FieldPath> {
292    let mut visitor = MacroVisitor::default();
293    visitor.visit_block(&function.block);
294
295    let mut names = BTreeSet::new();
296    visitor.fields.into_iter().filter(|field| names.insert(field.name())).collect()
297}
298
299#[derive(Default)]
300struct MacroVisitor {
301    fields: Vec<FieldPath>,
302}
303
304impl<'ast> Visit<'ast> for MacroVisitor {
305    fn visit_macro(&mut self, mac: &'ast Macro) {
306        if mac
307            .path
308            .segments
309            .last()
310            .is_some_and(|segment| segment.ident == "miden_span_record")
311        {
312            if let Ok(records) = syn::parse2::<RecordFields>(mac.tokens.clone()) {
313                self.fields.extend(records.fields.into_iter().map(|field| field.path));
314            }
315        }
316
317        syn::visit::visit_macro(self, mac);
318    }
319}
320
321type InstrumentFields = Fields<false>;
322type RecordFields = Fields<true>;
323
324struct Fields<const VALUE_REQUIRED: bool> {
325    fields: Punctuated<RecordField, Token![,]>,
326}
327
328impl<const VALUE_REQUIRED: bool> Parse for Fields<VALUE_REQUIRED> {
329    fn parse(input: ParseStream<'_>) -> Result<Self> {
330        Ok(Self {
331            fields: Punctuated::parse_terminated_with(input, |input| {
332                RecordField::parse(input, VALUE_REQUIRED)
333            })?,
334        })
335    }
336}
337
338struct RecordField {
339    path: FieldPath,
340    value: Option<RecordValue>,
341}
342
343impl RecordField {
344    fn parse(input: ParseStream<'_>, value_required: bool) -> Result<Self> {
345        let shorthand_formatter = if value_required {
346            None
347        } else {
348            Formatter::parse_optional(input)?
349        };
350        let path = input.parse()?;
351        validate_field_name(&path)?;
352        let value = if value_required || shorthand_formatter.is_none() && input.peek(Token![=]) {
353            input.parse::<Token![=]>()?;
354            Some(input.parse()?)
355        } else {
356            None
357        };
358
359        Ok(Self { path, value })
360    }
361}
362
363struct FieldPath {
364    first: Ident,
365    rest: Vec<(Dot, Ident)>,
366}
367
368impl FieldPath {
369    fn name(&self) -> String {
370        std::iter::once(&self.first)
371            .chain(self.rest.iter().map(|(_, ident)| ident))
372            .map(ToString::to_string)
373            .collect::<Vec<_>>()
374            .join(".")
375    }
376}
377
378impl Parse for FieldPath {
379    fn parse(input: ParseStream<'_>) -> Result<Self> {
380        let first = input.parse()?;
381        let mut rest = Vec::new();
382
383        while input.peek(Token![.]) {
384            rest.push((input.parse()?, input.parse()?));
385        }
386
387        Ok(Self { first, rest })
388    }
389}
390
391impl ToTokens for FieldPath {
392    fn to_tokens(&self, tokens: &mut TokenStream2) {
393        self.first.to_tokens(tokens);
394        for (dot, ident) in &self.rest {
395            dot.to_tokens(tokens);
396            ident.to_tokens(tokens);
397        }
398    }
399}
400
401struct RecordValue {
402    formatter: Formatter,
403    expr: Expr,
404}
405
406impl RecordValue {
407    fn value_tokens(&self) -> TokenStream2 {
408        let expr = &self.expr;
409
410        match self.formatter {
411            Formatter::Display => quote! { &::tracing::field::display(#expr) },
412            Formatter::Debug => quote! { &::tracing::field::debug(#expr) },
413            Formatter::Plain => quote! { &#expr },
414        }
415    }
416}
417
418impl Parse for RecordValue {
419    fn parse(input: ParseStream<'_>) -> Result<Self> {
420        let formatter = Formatter::parse_optional(input)?.unwrap_or(Formatter::Plain);
421        let expr = input.parse()?;
422
423        Ok(Self { formatter, expr })
424    }
425}
426
427enum Formatter {
428    Display,
429    Debug,
430    Plain,
431}
432
433impl Formatter {
434    fn parse_optional(input: ParseStream<'_>) -> Result<Option<Self>> {
435        if input.peek(Token![%]) {
436            input.parse::<Token![%]>()?;
437            Ok(Some(Self::Display))
438        } else if input.peek(Token![?]) {
439            input.parse::<Token![?]>()?;
440            Ok(Some(Self::Debug))
441        } else {
442            Ok(None)
443        }
444    }
445}