Skip to content

Commit 0ee9e2e

Browse files
Resolve oneapi_rs crate name at macro-expansion time
1 parent 757b0b4 commit 0ee9e2e

3 files changed

Lines changed: 66 additions & 5 deletions

File tree

Cargo.lock

Lines changed: 49 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

oneapi-rs-derive/Cargo.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ edition = "2024"
77
proc-macro = true
88

99
[dependencies]
10+
proc-macro-crate = "3.5.0"
1011
proc-macro2 = "1.0.107"
1112
quote = "1.0.47"
1213
syn = "3.0.3"

oneapi-rs-derive/src/lib.rs

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,29 +1,40 @@
11
use proc_macro::TokenStream;
2+
use proc_macro_crate::{FoundCrate, crate_name};
23
use quote::{format_ident, quote};
34
use syn::{
4-
Data, DataStruct, DeriveInput, Error, Field, LitInt, WhereClause, parse_macro_input,
5+
Data, DataStruct, DeriveInput, Error, Field, Ident, LitInt, WhereClause, parse_macro_input,
56
parse_quote,
67
};
78

9+
fn find_oneapi() -> Ident {
10+
let crate_name =
11+
crate_name("oneapi-rs").expect("Expected oneapi-rs to be present in Cargo.toml");
12+
match crate_name {
13+
FoundCrate::Itself => format_ident!("oneapi_rs"),
14+
FoundCrate::Name(name) => format_ident!("{name}"),
15+
}
16+
}
17+
818
/// Derive macro generating an impl of the `KernelArgumentList` trait for a given struct.
919
#[proc_macro_derive(KernelArgumentList)]
1020
pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream {
1121
let mut input = parse_macro_input!(input as DeriveInput);
22+
let oneapi = find_oneapi();
1223

1324
let Data::Struct(data) = &input.data else {
1425
return Error::new_spanned(input, "This derive macro only works on structs.")
1526
.into_compile_error()
1627
.into();
1728
};
18-
expand_where_clause(input.generics.make_where_clause(), data);
29+
expand_where_clause(input.generics.make_where_clause(), data, &oneapi);
1930
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
2031

2132
let ident = input.ident;
2233
let argc = data.fields.len();
2334
let members = data.fields.members();
2435

2536
let expanded = quote! {
26-
unsafe impl #impl_generics oneapi_rs::kernel::KernelArgumentList<#argc>
37+
unsafe impl #impl_generics #oneapi::kernel::KernelArgumentList<#argc>
2738
for #ident #ty_generics #where_clause {
2839
unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] {
2940
[ #(unsafe { self.#members.as_raw_arg() }),* ]
@@ -34,11 +45,11 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream {
3445
TokenStream::from(expanded)
3546
}
3647

37-
fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct) {
48+
fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct, oneapi: &Ident) {
3849
for Field { ty, .. } in &data.fields {
3950
where_clause
4051
.predicates
41-
.push(parse_quote!(#ty: oneapi_rs::kernel::KernelArgument));
52+
.push(parse_quote!(#ty: #oneapi::kernel::KernelArgument));
4253
}
4354
}
4455

0 commit comments

Comments
 (0)