Skip to content

Commit 06db8bb

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

3 files changed

Lines changed: 65 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: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,29 +1,39 @@
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 = crate_name("oneapi_rs").expect("oneapi_rs is present in Cargo.toml");
11+
match crate_name {
12+
FoundCrate::Itself => format_ident!("crate"),
13+
FoundCrate::Name(name) => format_ident!("{name}"),
14+
}
15+
}
16+
817
/// Derive macro generating an impl of the `KernelArgumentList` trait for a given struct.
918
#[proc_macro_derive(KernelArgumentList)]
1019
pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream {
1120
let mut input = parse_macro_input!(input as DeriveInput);
21+
let oneapi = find_oneapi();
1222

1323
let Data::Struct(data) = &input.data else {
1424
return Error::new_spanned(input, "This derive macro only works on structs.")
1525
.into_compile_error()
1626
.into();
1727
};
18-
expand_where_clause(input.generics.make_where_clause(), data);
28+
expand_where_clause(input.generics.make_where_clause(), data, &oneapi);
1929
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
2030

2131
let ident = input.ident;
2232
let argc = data.fields.len();
2333
let members = data.fields.members();
2434

2535
let expanded = quote! {
26-
unsafe impl #impl_generics oneapi_rs::kernel::KernelArgumentList<#argc>
36+
unsafe impl #impl_generics #oneapi::kernel::KernelArgumentList<#argc>
2737
for #ident #ty_generics #where_clause {
2838
unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] {
2939
[ #(unsafe { self.#members.as_raw_arg() }),* ]
@@ -34,11 +44,11 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream {
3444
TokenStream::from(expanded)
3545
}
3646

37-
fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct) {
47+
fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct, oneapi: &Ident) {
3848
for Field { ty, .. } in &data.fields {
3949
where_clause
4050
.predicates
41-
.push(parse_quote!(#ty: oneapi_rs::kernel::KernelArgument));
51+
.push(parse_quote!(#ty: #oneapi::kernel::KernelArgument));
4252
}
4353
}
4454

0 commit comments

Comments
 (0)