1use 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#[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#[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 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 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 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;