Skip to main content

ptx_syntax_proc_macros/
lib.rs

1//! Procedural macros used by the PTX parser implementation.
2
3use proc_macro::TokenStream;
4use proc_macro2::Span as ProcSpan;
5use quote::{format_ident, quote};
6use syn::{
7    Data, DeriveInput, Expr, Fields, FieldsNamed, Ident, Path, Token,
8    parse::{Parse, ParseStream},
9    parse_macro_input,
10    punctuated::Punctuated,
11};
12
13/// A field in the cmap macro - either `field = expr`, `field`, or `_`
14enum CmapField {
15    /// `field = expr` - assign expr to field
16    Assign { name: Ident, expr: Expr },
17    /// `field` - shorthand for `field: field`
18    Shorthand { name: Ident },
19    /// `_` - wildcard, ignored in pattern
20    Wildcard,
21}
22
23impl Parse for CmapField {
24    fn parse(input: ParseStream) -> syn::Result<Self> {
25        // Check for wildcard first
26        if input.peek(Token![_]) {
27            let _: Token![_] = input.parse()?;
28            return Ok(CmapField::Wildcard);
29        }
30
31        let name: Ident = input.parse()?;
32
33        // Check if there's an `=` token
34        if input.peek(Token![=]) {
35            let _eq: Token![=] = input.parse()?;
36            let expr: Expr = input.parse()?;
37            Ok(CmapField::Assign { name, expr })
38        } else {
39            Ok(CmapField::Shorthand { name })
40        }
41    }
42}
43
44/// Input structure for the cmap macro
45///
46/// Syntax: `Type { field1 = expr1, field2, ... }`
47struct CmapInput {
48    type_path: Path,
49    fields: Punctuated<CmapField, Token![,]>,
50}
51
52impl Parse for CmapInput {
53    fn parse(input: ParseStream) -> syn::Result<Self> {
54        // Parse Type path
55        let type_path: Path = input.parse()?;
56
57        // Parse { fields } - now required
58        let fields_content;
59        syn::braced!(fields_content in input);
60        let fields = fields_content.parse_terminated(CmapField::parse, Token![,])?;
61
62        Ok(CmapInput { type_path, fields })
63    }
64}
65
66/// Constructor mapping macro - eliminates field repetition
67///
68/// Syntax:
69/// ```text
70/// cclosure!(Type { field1, field2=expr })
71/// ```
72///
73/// Example:
74/// ```text
75/// cclosure!(Operand::SymbolOffset { symbol, _, offset=offset.unwrap() })
76/// // Generates: |(symbol, _, offset), span| c!(Operand::SymbolOffset { symbol, _, offset=offset.unwrap() })
77/// ```
78#[proc_macro]
79pub fn cclosure(input: TokenStream) -> TokenStream {
80    let CmapInput { type_path, fields } = parse_macro_input!(input as CmapInput);
81
82    // Extract pattern elements from field names (in order)
83    let mut pattern_elements = Vec::new();
84    for field in &fields {
85        match field {
86            CmapField::Assign { name, .. } => {
87                pattern_elements.push(quote! { #name });
88            }
89            CmapField::Shorthand { name } => {
90                pattern_elements.push(quote! { #name });
91            }
92            CmapField::Wildcard => {
93                pattern_elements.push(quote! { _ });
94            }
95        }
96    }
97
98    // Build the tuple pattern from the field names
99    let pattern = if pattern_elements.len() == 1 {
100        let elem = &pattern_elements[0];
101        quote! { #elem }
102    } else {
103        quote! { (#(#pattern_elements),*) }
104    };
105
106    // Generate field assignments (skip wildcards)
107    let field_assignments = fields.iter().filter_map(|field| match field {
108        CmapField::Assign { name, expr } => Some(quote! { #name: #expr }),
109        CmapField::Shorthand { name } => Some(quote! { #name: #name }),
110        CmapField::Wildcard => None,
111    });
112
113    // Generate the closure
114    let expanded = quote! {
115        move |#pattern, span| {
116            #type_path {
117                #(#field_assignments,)*
118                span
119            }
120        }
121    };
122
123    TokenStream::from(expanded)
124}
125
126/// Constructor macro - builds a struct with automatic span field
127///
128/// Syntax:
129/// ```text
130/// c!(Type { field1 = expr1, field2, ... })
131/// ```
132///
133/// Example:
134/// ```text
135/// c!(VariableDirective { name=something, foo })
136/// // Expands to: VariableDirective { name: something, foo: foo, span: span }
137/// ```
138#[proc_macro]
139pub fn c(input: TokenStream) -> TokenStream {
140    // Parse: [span_expr =>] Type { fields }
141    let input_parsed = parse_macro_input!(input with parse_c_input);
142
143    let (span_expr, type_path, fields) = input_parsed;
144
145    // Process fields similar to cmap
146    let field_assignments = fields.iter().filter_map(|field| match field {
147        CmapField::Assign { name, expr } => Some(quote! { #name: #expr }),
148        CmapField::Shorthand { name } => Some(quote! { #name: #name }),
149        CmapField::Wildcard => None,
150    });
151
152    // Generate the struct construction
153    let expanded = quote! {
154        #type_path {
155            #(#field_assignments,)*
156            span: #span_expr
157        }
158    };
159
160    TokenStream::from(expanded)
161}
162
163fn parse_c_input(
164    input: ParseStream,
165) -> syn::Result<(Expr, Path, Punctuated<CmapField, Token![,]>)> {
166    // Fork the input to try parsing with span expression first
167    let fork = input.fork();
168
169    // Try to parse: Expr => Path { fields }
170    let (span_expr, type_path) = if let Ok(_expr) = fork.parse::<Expr>() {
171        if fork.peek(Token![=>]) {
172            // Successfully parsed expr followed by =>, so advance the main input
173            let expr: Expr = input.parse()?;
174            let _arrow: Token![=>] = input.parse()?;
175            let type_path: Path = input.parse()?;
176            (expr, type_path)
177        } else {
178            // No =>, so the first thing must be the type path
179            let type_path: Path = input.parse()?;
180            (syn::parse_quote!(span), type_path)
181        }
182    } else {
183        // Couldn't parse as expr, try as path
184        let type_path: Path = input.parse()?;
185        (syn::parse_quote!(span), type_path)
186    };
187
188    // Parse { fields } if present
189    let fields = if input.peek(syn::token::Brace) {
190        let fields_content;
191        syn::braced!(fields_content in input);
192        fields_content.parse_terminated(CmapField::parse, Token![,])?
193    } else {
194        Punctuated::new()
195    };
196
197    Ok((span_expr, type_path, fields))
198}
199
200/// Ok wrapper macro - wraps result in Ok(...) with automatic span field
201///
202/// Syntax:
203/// ```text
204/// ok!(Type { field1 = expr1, field2, ... })
205/// ```
206///
207/// Example:
208/// ```text
209/// ok!(TexHandler2 { operands, name=foo.to_string() })
210/// // Expands to: Ok(TexHandler2 { operands: operands, name: foo.to_string(), span: span })
211/// ```
212#[proc_macro]
213pub fn ok(input: TokenStream) -> TokenStream {
214    // Reuse c! macro logic but wrap in Ok(value)
215    // Used inside try_map closures where try_with_span will add the outer tuple
216    let c_result = c(input);
217    let c_tokens: proc_macro2::TokenStream = c_result.into();
218
219    let expanded = quote! {
220        Ok(#c_tokens)
221    };
222
223    TokenStream::from(expanded)
224}
225
226/// Error constructor macro - builds a PtxParseError with automatic span field
227///
228/// Syntax:
229/// ```text
230/// err!(ErrorKind)
231/// ```
232///
233/// Example:
234/// ```text
235/// err!(ParseErrorKind::InvalidLiteral("message".into()))
236/// // Expands to: Err(PtxParseError { kind: ParseErrorKind::InvalidLiteral("message".into()), span: span })
237/// ```
238#[proc_macro]
239pub fn err(input: TokenStream) -> TokenStream {
240    // Parse: [span_expr =>] error_kind
241    let input_parsed = parse_macro_input!(input with parse_err_input);
242
243    let (span_expr, error_kind) = input_parsed;
244
245    // Generate the error construction
246    // Note: We use crate::parser::PtxParseError which will resolve in the calling code's context
247    let expanded = quote! {
248        Err(crate::parser::PtxParseError {
249            kind: #error_kind,
250            span: #span_expr
251        })
252    };
253
254    TokenStream::from(expanded)
255}
256
257fn parse_err_input(input: ParseStream) -> syn::Result<(Expr, Expr)> {
258    // Try to parse span expression followed by =>
259    let span_expr = if input.peek2(Token![=>]) {
260        let expr: Expr = input.parse()?;
261        let _arrow: Token![=>] = input.parse()?;
262        expr
263    } else {
264        // If no span => provided, use `span` identifier
265        syn::parse_quote!(span)
266    };
267
268    // Parse error kind expression
269    let error_kind: Expr = input.parse()?;
270
271    Ok((span_expr, error_kind))
272}
273
274/// Constructor mapping macro with Ok wrapper - like cclosure! but wraps result with Ok
275///
276/// Syntax:
277/// ```text
278/// okmap!(Type { field1, field2=expr })
279/// ```
280///
281/// Example:
282/// ```text
283/// okmap!(Operand::SymbolOffset { symbol, _, offset=offset.unwrap() })
284/// // Generates: |(symbol, _, offset), span| Ok(Operand::SymbolOffset { symbol, offset: offset.unwrap(), span })
285/// ```
286#[proc_macro]
287pub fn okmap(input: TokenStream) -> TokenStream {
288    let CmapInput { type_path, fields } = parse_macro_input!(input as CmapInput);
289
290    // Extract pattern elements from field names (in order)
291    let mut pattern_elements = Vec::new();
292    for field in &fields {
293        match field {
294            CmapField::Assign { name, .. } => {
295                pattern_elements.push(quote! { #name });
296            }
297            CmapField::Shorthand { name } => {
298                pattern_elements.push(quote! { #name });
299            }
300            CmapField::Wildcard => {
301                pattern_elements.push(quote! { _ });
302            }
303        }
304    }
305
306    // Build the tuple pattern from the field names
307    let pattern = if pattern_elements.len() == 1 {
308        let elem = &pattern_elements[0];
309        quote! { #elem }
310    } else {
311        quote! { (#(#pattern_elements),*) }
312    };
313
314    // Generate field assignments (skip wildcards)
315    let field_assignments = fields.iter().filter_map(|field| match field {
316        CmapField::Assign { name, expr } => Some(quote! { #name: #expr }),
317        CmapField::Shorthand { name } => Some(quote! { #name: #name }),
318        CmapField::Wildcard => None,
319    });
320
321    // Generate the closure wrapped in Ok
322    let expanded = quote! {
323        move |#pattern, span| {
324            Ok(#type_path {
325                #(#field_assignments,)*
326                span
327            })
328        }
329    };
330
331    TokenStream::from(expanded)
332}
333
334/// Function macro - adds span parameter to closures
335///
336/// Syntax:
337/// ```text
338/// func!(|param1, param2| body)
339/// ```
340///
341/// Example:
342/// ```text
343/// func!(|x, y| x + y)
344/// // Expands to: |x, y, span| x + y
345/// ```
346#[proc_macro]
347pub fn func(input: TokenStream) -> TokenStream {
348    use syn::{ExprClosure, Pat};
349
350    let closure = parse_macro_input!(input as ExprClosure);
351
352    // Extract the existing parameters
353    let mut params = closure.inputs.clone();
354
355    // Add span parameter
356    let span_param: Pat = syn::parse_quote!(span);
357    params.push(span_param);
358
359    // Get the body
360    let body = closure.body;
361
362    // Generate the expanded closure
363    let expanded = quote! {
364        |#params| #body
365    };
366
367    TokenStream::from(expanded)
368}
369
370/// Implement the parser's `Spanned` trait and inherent span helpers.
371///
372/// The input must be a struct or enum whose variants use named fields and
373/// contain a field named `span`. Invalid inputs produce a compile error.
374#[proc_macro_derive(Spanned)]
375pub fn derive_spanned(input: TokenStream) -> TokenStream {
376    let input = parse_macro_input!(input as DeriveInput);
377    match impl_spanned(&input) {
378        Ok(tokens) => tokens.into(),
379        Err(error) => error.to_compile_error().into(),
380    }
381}
382
383fn impl_spanned(input: &DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
384    let name = &input.ident;
385    let generics = &input.generics;
386    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
387
388    let (span_arms, set_arms) = match &input.data {
389        Data::Struct(data) => {
390            let (span_arm, set_arm) = build_match_arm(quote! { Self }, &data.fields)?;
391            (vec![span_arm], vec![set_arm])
392        }
393        Data::Enum(data) => {
394            let mut span_arms = Vec::new();
395            let mut set_arms = Vec::new();
396            for variant in &data.variants {
397                let ident = &variant.ident;
398                let path = quote! { Self::#ident };
399                let (span_arm, set_arm) = build_match_arm(path, &variant.fields)?;
400                span_arms.push(span_arm);
401                set_arms.push(set_arm);
402            }
403            (span_arms, set_arms)
404        }
405        Data::Union(_) => {
406            return Err(syn::Error::new_spanned(
407                &input.ident,
408                "Spanned cannot be derived for unions",
409            ));
410        }
411    };
412
413    let span_ty = quote! { crate::parser::Span };
414    let trait_path = quote! { crate::span::Spanned };
415    let span_match = quote! {
416        match self {
417            #(#span_arms)*
418        }
419    };
420    let set_match = quote! {
421        match self {
422            #(#set_arms)*
423        }
424    };
425
426    Ok(quote! {
427        impl #impl_generics #trait_path for #name #ty_generics #where_clause {
428            fn span(&self) -> #span_ty {
429                #span_match
430            }
431
432            fn set_span(&mut self, span: #span_ty) {
433                #set_match
434            }
435        }
436
437        impl #impl_generics #name #ty_generics #where_clause {
438            pub fn span(&self) -> #span_ty {
439                <Self as #trait_path>::span(self)
440            }
441
442            pub fn with_span(mut self, span: #span_ty) -> Self {
443                <Self as #trait_path>::set_span(&mut self, span);
444                self
445            }
446        }
447    })
448}
449
450fn build_match_arm(
451    path: proc_macro2::TokenStream,
452    fields: &Fields,
453) -> syn::Result<(proc_macro2::TokenStream, proc_macro2::TokenStream)> {
454    match fields {
455        Fields::Named(named) => build_named_arm(path, named),
456        _ => Err(syn::Error::new(
457            ProcSpan::call_site(),
458            "Spanned derive only supports structs/enums with named `span` fields",
459        )),
460    }
461}
462
463fn build_named_arm(
464    path: proc_macro2::TokenStream,
465    fields: &FieldsNamed,
466) -> syn::Result<(proc_macro2::TokenStream, proc_macro2::TokenStream)> {
467    let has_span = fields
468        .named
469        .iter()
470        .any(|field| field.ident.as_ref().is_some_and(|ident| ident == "span"));
471    if !has_span {
472        return Err(syn::Error::new(
473            ProcSpan::call_site(),
474            "Spanned derive requires a field named `span`",
475        ));
476    }
477
478    let binding = format_ident!("__span_field");
479    let span_arm = quote! {
480        #path { span: #binding, .. } => #binding.clone(),
481    };
482    let set_arm = quote! {
483        #path { span: #binding, .. } => {
484            *#binding = span;
485        }
486    };
487
488    Ok((span_arm, set_arm))
489}