Skip to main content

tenferro_extension_macros/
lib.rs

1//! Procedural macros for tenferro extension crates.
2//!
3//! # Examples
4//!
5//! ```
6//! use tenferro_extension_macros::ExtensionFamilyId;
7//!
8//! #[derive(ExtensionFamilyId)]
9//! #[tenferro_extension(namespace = "my-crate", name = "fft", version = 1)]
10//! struct FftOp;
11//!
12//! assert_eq!(FftOp::FAMILY_ID, "my-crate.fft.v1");
13//! ```
14
15use proc_macro::TokenStream;
16use quote::{format_ident, quote};
17use syn::parse::{Parse, ParseStream};
18use syn::{parse_macro_input, DeriveInput, Expr, ExprLit, Ident, Lit, Path, Token};
19
20#[derive(Debug, Default)]
21struct ExtensionArgs {
22    namespace: Option<String>,
23    name: Option<String>,
24    version: Option<u64>,
25}
26
27struct RuntimeArgs {
28    runtime: Ident,
29    family_id: Path,
30    op_type: Path,
31    execute: Option<Path>,
32    execute_reads: Option<Path>,
33    execute_in_session: Option<Path>,
34    session_supported: Option<Path>,
35    backend_bound: Path,
36}
37
38impl Parse for ExtensionArgs {
39    fn parse(input: ParseStream<'_>) -> syn::Result<Self> {
40        let mut args = Self::default();
41        while !input.is_empty() {
42            let key: syn::Ident = input.parse()?;
43            input.parse::<Token![=]>()?;
44            let value: Expr = input.parse()?;
45            match key.to_string().as_str() {
46                "namespace" => args.namespace = Some(expect_string(value, "namespace")?),
47                "name" => args.name = Some(expect_string(value, "name")?),
48                "version" => args.version = Some(expect_u64(value, "version")?),
49                other => {
50                    return Err(syn::Error::new(
51                        key.span(),
52                        format!("unsupported tenferro_extension argument {other:?}"),
53                    ));
54                }
55            }
56            if input.is_empty() {
57                break;
58            }
59            input.parse::<Token![,]>()?;
60        }
61        Ok(args)
62    }
63}
64
65impl Parse for RuntimeArgs {
66    fn parse(input: ParseStream<'_>) -> syn::Result<Self> {
67        let mut runtime = None;
68        let mut family_id = None;
69        let mut op_type = None;
70        let mut execute = None;
71        let mut execute_reads = None;
72        let mut execute_in_session = None;
73        let mut session_supported = None;
74        let mut backend_bound = None;
75
76        while !input.is_empty() {
77            let key: Ident = input.parse()?;
78            input.parse::<Token![=]>()?;
79            match key.to_string().as_str() {
80                "runtime" => runtime = Some(input.parse()?),
81                "family_id" => family_id = Some(input.parse()?),
82                "op_type" => op_type = Some(input.parse()?),
83                "execute" => execute = Some(input.parse()?),
84                "execute_reads" => execute_reads = Some(input.parse()?),
85                "execute_in_session" => execute_in_session = Some(input.parse()?),
86                "session_supported" => session_supported = Some(input.parse()?),
87                "backend_bound" => backend_bound = Some(input.parse()?),
88                other => {
89                    return Err(syn::Error::new(
90                        key.span(),
91                        format!("unsupported define_extension_runtime argument {other:?}"),
92                    ));
93                }
94            }
95            if input.is_empty() {
96                break;
97            }
98            input.parse::<Token![,]>()?;
99        }
100
101        Ok(Self {
102            runtime: required(runtime, "runtime")?,
103            family_id: required(family_id, "family_id")?,
104            op_type: required(op_type, "op_type")?,
105            execute,
106            execute_reads,
107            execute_in_session,
108            session_supported,
109            backend_bound: backend_bound
110                .unwrap_or_else(|| syn::parse_quote!(tenferro_tensor::TensorBackend)),
111        })
112    }
113}
114
115/// Derive an inherent `FAMILY_ID` constant for an extension payload type.
116///
117/// The required attribute is:
118/// `#[tenferro_extension(namespace = "...", version = N)]`.
119/// `name = "..."` is optional; when omitted, the Rust type name is converted
120/// to snake_case.
121#[proc_macro_derive(ExtensionFamilyId, attributes(tenferro_extension))]
122pub fn derive_extension_family_id(input: TokenStream) -> TokenStream {
123    let input = parse_macro_input!(input as DeriveInput);
124    match expand_extension_family_id(input) {
125        Ok(tokens) => tokens.into(),
126        Err(err) => err.to_compile_error().into(),
127    }
128}
129
130/// Generate a standard extension module, preparation engine, and prepared
131/// operation.
132///
133/// Exactly one execution route must be supplied.
134///
135/// * `execute_reads` runs the operation inside the runtime-formed context and
136///   has this signature:
137///   `fn<B: BackendSession + ?Sized>(&OpType, &[TensorRead<'_>], &mut ExtensionExecutionContext<'_, B>)`.
138///   The generated owner entry opens the backend session and the context
139///   borrows it (`B` is `dyn BackendSession`), so the callback never runs on
140///   the owner. Prefer the session route below for new extensions.
141/// * the session route (`execute_in_session` + `session_supported`) is the
142///   preferred one: `session_supported` has signature
143///   `fn<B: BackendBound + 'static>(&OpType) -> bool` and `execute_in_session`
144///   has signature
145///   `fn(&OpType, &mut dyn BackendSession, &mut ExtensionCacheStore, &[TensorRead<'_>])`.
146///   With it, the generated owner entry forms the session itself and hands the
147///   extension only a session, so no operation ever runs on the owner.
148///
149/// The legacy `execute` argument is accepted but unused and may be omitted.
150#[proc_macro]
151pub fn define_extension_runtime(input: TokenStream) -> TokenStream {
152    let args = parse_macro_input!(input as RuntimeArgs);
153    match expand_extension_runtime(args) {
154        Ok(tokens) => tokens.into(),
155        Err(err) => err.to_compile_error().into(),
156    }
157}
158
159fn expand_extension_family_id(input: DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
160    let mut parsed = None;
161    for attr in &input.attrs {
162        if attr.path().is_ident("tenferro_extension") {
163            let args = attr.parse_args::<ExtensionArgs>()?;
164            parsed = Some(args);
165        }
166    }
167    let args = parsed.ok_or_else(|| {
168        syn::Error::new_spanned(
169            &input.ident,
170            "missing #[tenferro_extension(namespace = \"...\", version = N)]",
171        )
172    })?;
173    let namespace = args.namespace.ok_or_else(|| {
174        syn::Error::new_spanned(&input.ident, "missing tenferro_extension namespace")
175    })?;
176    let version = args.version.ok_or_else(|| {
177        syn::Error::new_spanned(&input.ident, "missing tenferro_extension version")
178    })?;
179    let name = args
180        .name
181        .unwrap_or_else(|| to_snake_case(&input.ident.to_string()));
182    let family_id = format!("{namespace}.{name}.v{version}");
183    let ident = input.ident;
184
185    Ok(quote! {
186        impl #ident {
187            /// Stable extension family identifier generated by `ExtensionFamilyId`.
188            pub const FAMILY_ID: &'static str = #family_id;
189        }
190    })
191}
192
193fn expand_extension_runtime(args: RuntimeArgs) -> syn::Result<proc_macro2::TokenStream> {
194    let RuntimeArgs {
195        runtime,
196        family_id,
197        op_type,
198        execute: _execute,
199        execute_reads,
200        execute_in_session,
201        session_supported,
202        backend_bound,
203    } = args;
204    let module = format_ident!("{}Module", runtime);
205    let planning_config = format_ident!("{}PlanningConfig", runtime);
206    let prepared_operation = format_ident!("{}PreparedOperation", runtime);
207    let owner_execute = match (&execute_reads, &execute_in_session) {
208        (None, Some(execute_in_session)) => quote! {
209            fn execute(
210                &self,
211                context: &mut tenferro_runtime::ErasedExecutionContext<'_>,
212                extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
213                inputs: &[tenferro_tensor::TensorRead<'_>],
214            ) -> tenferro_runtime::Result<Vec<tenferro_tensor::Tensor>> {
215                // The runtime owns the region: the owner entry forms the session
216                // and the extension only ever runs through it.
217                let backend = context
218                    .downcast_mut::<B>(self.binding.context_identity())
219                    .map_err(|source| tenferro_runtime::Error::runtime_state_source(
220                        "extension",
221                        tenferro_runtime::ErrorPhase::Execution,
222                        source,
223                    ))?;
224                backend.with_backend_session(
225                    |session| -> tenferro_runtime::Result<Vec<tenferro_tensor::Tensor>> {
226                        Ok(#execute_in_session(&self.op, session, extension_caches, inputs)?)
227                    },
228                )?
229            }
230        },
231        (Some(execute_reads), None) => quote! {
232            fn execute(
233                &self,
234                context: &mut tenferro_runtime::ErasedExecutionContext<'_>,
235                extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
236                inputs: &[tenferro_tensor::TensorRead<'_>],
237            ) -> tenferro_runtime::Result<Vec<tenferro_tensor::Tensor>> {
238                let backend = context
239                    .downcast_mut::<B>(self.binding.context_identity())
240                    .map_err(|source| tenferro_runtime::Error::runtime_state_source(
241                        "extension",
242                        tenferro_runtime::ErrorPhase::Execution,
243                        source,
244                    ))?;
245                // As in the session route, the owner entry forms the session and
246                // the callback's context borrows that session, never the owner.
247                backend.with_backend_session(
248                    |session| -> tenferro_runtime::Result<Vec<tenferro_tensor::Tensor>> {
249                        let mut ctx = tenferro_runtime::ExtensionExecutionContext::new(
250                            session,
251                            extension_caches,
252                        );
253                        Ok(#execute_reads(&self.op, inputs, &mut ctx)?)
254                    },
255                )?
256            }
257        },
258        (Some(execute_reads), Some(_)) => {
259            return Err(syn::Error::new_spanned(
260                execute_reads,
261                "supply either execute_reads or execute_in_session, not both",
262            ))
263        }
264        (None, None) => {
265            return Err(syn::Error::new_spanned(
266                &op_type,
267                "either execute_reads or execute_in_session must be supplied",
268            ))
269        }
270    };
271    let session_methods = match (execute_in_session, session_supported) {
272        (Some(execute_in_session), Some(session_supported)) => quote! {
273            fn supports_session(&self) -> bool {
274                #session_supported::<B>(&self.op)
275            }
276
277            fn execute_in_session(
278                &self,
279                session: &mut dyn tenferro_tensor::BackendSession,
280                extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
281                inputs: &[tenferro_tensor::TensorRead<'_>],
282            ) -> tenferro_runtime::Result<Vec<tenferro_tensor::Tensor>> {
283                Ok(#execute_in_session(&self.op, session, extension_caches, inputs)?)
284            }
285        },
286        (None, None) => quote! {},
287        (Some(path), None) | (None, Some(path)) => {
288            return Err(syn::Error::new_spanned(
289                path,
290                "execute_in_session and session_supported must be supplied together",
291            ))
292        }
293    };
294    Ok(quote! {
295        pub(crate) struct #runtime<B: #backend_bound + 'static> {
296            engine_id: tenferro_runtime::EngineId,
297            _backend: std::marker::PhantomData<fn() -> B>,
298        }
299
300        pub(crate) struct #module<B: #backend_bound + 'static> {
301            module_id: tenferro_runtime::ExtensionModuleId,
302            engine_id: tenferro_runtime::EngineId,
303            _backend: std::marker::PhantomData<fn() -> B>,
304        }
305
306        #[derive(Debug, Default)]
307        pub(crate) struct #planning_config;
308
309        pub(crate) struct #prepared_operation<B: #backend_bound + 'static> {
310            binding: tenferro_runtime::PreparedOperationBinding,
311            specialization: tenferro_runtime::SpecializationProjection,
312            op: #op_type,
313            _backend: std::marker::PhantomData<fn() -> B>,
314        }
315
316        impl<B: #backend_bound + 'static> std::fmt::Debug for #runtime<B> {
317            fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
318                formatter
319                    .debug_struct(stringify!(#runtime))
320                    .field("family_id", &#family_id)
321                    .field("engine_id", &self.engine_id)
322                    .field("backend_type", &std::any::type_name::<B>())
323                    .finish()
324            }
325        }
326
327        impl<B: #backend_bound + 'static> std::fmt::Debug for #module<B> {
328            fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
329                formatter
330                    .debug_struct(stringify!(#module))
331                    .field("module_id", &self.module_id)
332                    .field("engine_id", &self.engine_id)
333                    .field("backend_type", &std::any::type_name::<B>())
334                    .finish()
335            }
336        }
337
338        impl<B: #backend_bound + 'static> std::fmt::Debug for #prepared_operation<B> {
339            fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
340                formatter
341                    .debug_struct(stringify!(#prepared_operation))
342                    .field("family_id", &#family_id)
343                    .field("binding", &self.binding)
344                    .field("specialization", &self.specialization)
345                    .field("backend_type", &std::any::type_name::<B>())
346                    .finish_non_exhaustive()
347            }
348        }
349
350        impl<B: #backend_bound + 'static> tenferro_runtime::ExtensionEngine for #runtime<B>
351        where
352            #op_type: Clone + Send + Sync + 'static,
353        {
354            fn family_id(&self) -> &'static str {
355                #family_id
356            }
357
358            fn engine_id(&self) -> &tenferro_runtime::EngineId {
359                &self.engine_id
360            }
361
362            fn context_identity(&self) -> tenferro_runtime::ExecutionContextIdentity {
363                tenferro_runtime::ExecutionContextIdentity::of::<B>()
364            }
365
366            fn prepare(
367                &self,
368                request: tenferro_runtime::ExtensionPrepareRequest<'_>,
369            ) -> std::result::Result<tenferro_runtime::PrepareCapability, tenferro_runtime::PrepareError> {
370                let op = request
371                    .operation()
372                    .as_any()
373                    .downcast_ref::<#op_type>()
374                    .cloned()
375                    .ok_or_else(|| tenferro_runtime::PrepareError::ProviderContract {
376                        source: tenferro_runtime::ProviderContractError::WrongOperationFamily {
377                            expected: tenferro_runtime::CoreCapabilityKind::Elementwise,
378                            operation: #family_id,
379                        },
380                    })?;
381                let prepared = std::sync::Arc::new(#prepared_operation::<B> {
382                    binding: request.binding().clone(),
383                    specialization: request.specialization().clone(),
384                    op,
385                    _backend: std::marker::PhantomData,
386                });
387                Ok(tenferro_runtime::PrepareCapability::Prepared(
388                    tenferro_runtime::PreparedOperationPlan::executable(prepared.clone(), prepared)
389                ))
390            }
391        }
392
393        impl tenferro_runtime::ExtensionPlanningConfig for #planning_config {
394            fn family_id(&self) -> &'static str {
395                #family_id
396            }
397
398            fn as_any(&self) -> &dyn std::any::Any {
399                self
400            }
401
402            fn payload_hash(&self, state: &mut dyn std::hash::Hasher) {
403                state.write_u8(0);
404            }
405
406            fn payload_eq(&self, other: &dyn tenferro_runtime::ExtensionPlanningConfig) -> bool {
407                other.as_any().downcast_ref::<Self>().is_some()
408            }
409
410            fn retained_bytes(&self) -> usize {
411                0
412            }
413        }
414
415        impl<B: #backend_bound + 'static> tenferro_runtime::PreparedOperation for #prepared_operation<B>
416        where
417            #op_type: Clone + Send + Sync + 'static,
418        {
419            fn binding(&self) -> &tenferro_runtime::PreparedOperationBinding {
420                &self.binding
421            }
422
423            fn specialization(&self) -> &tenferro_runtime::SpecializationProjection {
424                &self.specialization
425            }
426
427            fn retained_bytes(&self) -> usize {
428                0
429            }
430
431        }
432
433        impl<B: #backend_bound + 'static> tenferro_runtime::PreparedOperationExecutor for #prepared_operation<B>
434        where
435            #op_type: Clone + Send + Sync + 'static,
436        {
437            #owner_execute
438
439            #session_methods
440        }
441
442        impl<B: #backend_bound + 'static> tenferro_runtime::ExtensionModule for #module<B>
443        where
444            #op_type: Clone + Send + Sync + 'static,
445        {
446            fn module_id(&self) -> &tenferro_runtime::ExtensionModuleId {
447                &self.module_id
448            }
449
450            fn configure(
451                &self,
452                registrar: &mut tenferro_runtime::ExtensionModuleRegistrar<'_>,
453            ) -> std::result::Result<(), tenferro_runtime::ExtensionModuleError> {
454                registrar.register_engine(std::sync::Arc::new(#runtime::<B> {
455                    engine_id: self.engine_id.clone(),
456                    _backend: std::marker::PhantomData,
457                }))?;
458                registrar.register_planning_config(
459                    self.engine_id.clone(),
460                    std::sync::Arc::new(#planning_config),
461                )?;
462                Ok(())
463            }
464        }
465
466        #[doc = "Build this extension module for one runtime engine."]
467        #[doc = "\n# Errors\n\nReturns `RuntimeConfigError::MalformedIdentity` when the generated module identifier is invalid."]
468        pub fn extension_module<B: #backend_bound + 'static>(
469            engine_id: tenferro_runtime::EngineId,
470        ) -> std::result::Result<
471            std::sync::Arc<dyn tenferro_runtime::ExtensionModule>,
472            tenferro_runtime::RuntimeConfigError,
473        >
474        where
475            #op_type: Clone + Send + Sync + 'static,
476        {
477            Ok(std::sync::Arc::new(#module::<B> {
478                module_id: tenferro_runtime::ExtensionModuleId::new(format!("{}.module", #family_id))?,
479                engine_id,
480                _backend: std::marker::PhantomData,
481            }))
482        }
483
484    })
485}
486
487fn expect_string(value: Expr, field: &str) -> syn::Result<String> {
488    match value {
489        Expr::Lit(ExprLit {
490            lit: Lit::Str(value),
491            ..
492        }) => Ok(value.value()),
493        other => Err(syn::Error::new_spanned(
494            other,
495            format!("{field} must be a string literal"),
496        )),
497    }
498}
499
500fn expect_u64(value: Expr, field: &str) -> syn::Result<u64> {
501    match value {
502        Expr::Lit(ExprLit {
503            lit: Lit::Int(value),
504            ..
505        }) => value.base10_parse(),
506        other => Err(syn::Error::new_spanned(
507            other,
508            format!("{field} must be an integer literal"),
509        )),
510    }
511}
512
513fn required<T>(value: Option<T>, field: &str) -> syn::Result<T> {
514    value.ok_or_else(|| syn::Error::new(proc_macro2::Span::call_site(), format!("missing {field}")))
515}
516
517fn to_snake_case(input: &str) -> String {
518    let mut out = String::new();
519    let mut prev_lower_or_digit = false;
520    for ch in input.chars() {
521        if ch.is_ascii_uppercase() {
522            if prev_lower_or_digit {
523                out.push('_');
524            }
525            out.push(ch.to_ascii_lowercase());
526            prev_lower_or_digit = false;
527        } else {
528            prev_lower_or_digit = ch.is_ascii_lowercase() || ch.is_ascii_digit();
529            out.push(ch);
530        }
531    }
532    out
533}
534
535#[cfg(test)]
536mod tests;