extern crate quote;
use quote::ToTokens;
extern crate proc_macro;
use proc_macro::TokenStream;
use proc_macro2::Span;
extern crate syn;
use syn::fold::Fold;
use syn::*;
use std::result::Result;
macro_rules! fold_expr_default {
($f:expr, $node:expr) => {
match $node {
Expr::Array(_binding_0) => Expr::Array($f.fold_expr_array(_binding_0)),
Expr::Assign(_binding_0) => Expr::Assign($f.fold_expr_assign(_binding_0)),
Expr::AssignOp(_binding_0) => Expr::AssignOp($f.fold_expr_assign_op(_binding_0)),
Expr::Async(_binding_0) => Expr::Async($f.fold_expr_async(_binding_0)),
Expr::Await(_binding_0) => Expr::Await($f.fold_expr_await(_binding_0)),
Expr::Binary(_binding_0) => Expr::Binary($f.fold_expr_binary(_binding_0)),
Expr::Block(_binding_0) => Expr::Block($f.fold_expr_block(_binding_0)),
Expr::Box(_binding_0) => Expr::Box($f.fold_expr_box(_binding_0)),
Expr::Break(_binding_0) => Expr::Break($f.fold_expr_break(_binding_0)),
Expr::Call(_binding_0) => Expr::Call($f.fold_expr_call(_binding_0)),
Expr::Cast(_binding_0) => Expr::Cast($f.fold_expr_cast(_binding_0)),
Expr::Closure(_binding_0) => Expr::Closure($f.fold_expr_closure(_binding_0)),
Expr::Continue(_binding_0) => Expr::Continue($f.fold_expr_continue(_binding_0)),
Expr::Field(_binding_0) => Expr::Field($f.fold_expr_field(_binding_0)),
Expr::ForLoop(_binding_0) => Expr::ForLoop($f.fold_expr_for_loop(_binding_0)),
Expr::Group(_binding_0) => Expr::Group($f.fold_expr_group(_binding_0)),
Expr::If(_binding_0) => Expr::If($f.fold_expr_if(_binding_0)),
Expr::Index(_binding_0) => Expr::Index($f.fold_expr_index(_binding_0)),
Expr::Let(_binding_0) => Expr::Let($f.fold_expr_let(_binding_0)),
Expr::Lit(_binding_0) => Expr::Lit($f.fold_expr_lit(_binding_0)),
Expr::Loop(_binding_0) => Expr::Loop($f.fold_expr_loop(_binding_0)),
Expr::Macro(_binding_0) => Expr::Macro($f.fold_expr_macro(_binding_0)),
Expr::Match(_binding_0) => Expr::Match($f.fold_expr_match(_binding_0)),
Expr::MethodCall(_binding_0) => Expr::MethodCall($f.fold_expr_method_call(_binding_0)),
Expr::Paren(_binding_0) => Expr::Paren($f.fold_expr_paren(_binding_0)),
Expr::Path(_binding_0) => Expr::Path($f.fold_expr_path(_binding_0)),
Expr::Range(_binding_0) => Expr::Range($f.fold_expr_range(_binding_0)),
Expr::Reference(_binding_0) => Expr::Reference($f.fold_expr_reference(_binding_0)),
Expr::Repeat(_binding_0) => Expr::Repeat($f.fold_expr_repeat(_binding_0)),
Expr::Return(_binding_0) => Expr::Return($f.fold_expr_return(_binding_0)),
Expr::Struct(_binding_0) => Expr::Struct($f.fold_expr_struct(_binding_0)),
Expr::Try(_binding_0) => Expr::Try($f.fold_expr_try(_binding_0)),
Expr::TryBlock(_binding_0) => Expr::TryBlock($f.fold_expr_try_block(_binding_0)),
Expr::Tuple(_binding_0) => Expr::Tuple($f.fold_expr_tuple(_binding_0)),
Expr::Type(_binding_0) => Expr::Type($f.fold_expr_type(_binding_0)),
Expr::Unary(_binding_0) => Expr::Unary($f.fold_expr_unary(_binding_0)),
Expr::Unsafe(_binding_0) => Expr::Unsafe($f.fold_expr_unsafe(_binding_0)),
Expr::Verbatim(_binding_0) => Expr::Verbatim(_binding_0),
Expr::While(_binding_0) => Expr::While($f.fold_expr_while(_binding_0)),
Expr::Yield(_binding_0) => Expr::Yield($f.fold_expr_yield(_binding_0)),
_ => unreachable!(),
}
};
}
pub fn modify_signature_for_helper_function(
input: TokenStream,
has_return: bool,
) -> Result<TokenStream, Vec<Error>> {
let maybe_ast = syn::parse::<ItemFn>(input.clone());
if let Ok(mut ast) = maybe_ast {
if has_return {
let input: proc_macro::TokenStream = quote! {
mut gpu: Gpu
}
.into();
ast.sig
.inputs
.insert(0, syn::parse::<FnArg>(input).unwrap());
if let ReturnType::Type(existing_output_arrow, existing_output_type) = ast.sig.output {
let output = quote! {
#existing_output_arrow (#existing_output_type, Gpu)
}
.into_token_stream();
ast.sig.output = syn::parse::<ReturnType>(output.into_token_stream().into())
.expect("could not change return type");
} else {
}
} else {
let input = quote! {
mut gpu: Gpu
}
.into();
ast.sig
.inputs
.insert(0, syn::parse::<FnArg>(input).unwrap());
let output = quote! {
-> ((), Gpu)
}
.into();
ast.sig.output = syn::parse::<ReturnType>(output).unwrap();
}
Ok(ast.to_token_stream().into())
} else {
Err(vec![Error::new(
Span::call_site().unwrap().into(),
"only functions that are items can be tagged with `#[gpu_use]`",
)])
}
}
pub fn modify_return_for_helper_function(
input: TokenStream,
has_return: bool,
) -> Result<TokenStream, Vec<Error>> {
let maybe_ast = syn::parse::<ItemFn>(input.clone());
if let Ok(mut ast) = maybe_ast {
if has_return {
let existing_body = ast.block;
let body = quote! {
{
(#existing_body, gpu)
}
};
ast.block = Box::new(
syn::parse::<Block>(body.into_token_stream().into())
.expect("could not change returns"),
);
} else {
let existing_body = ast.block;
let body = quote! {
{
#existing_body
((), gpu)
}
};
ast.block = Box::new(
syn::parse::<Block>(body.into_token_stream().into())
.expect("could not change returns"),
);
}
Ok(ast.to_token_stream().into())
} else {
Err(vec![Error::new(
Span::call_site().unwrap().into(),
"only functions that are items can be tagged with `#[gpu_use]`",
)])
}
}
pub struct HelperFunctionReturnModifier;
impl Fold for HelperFunctionReturnModifier {
fn fold_expr_return(&mut self, i: ExprReturn) -> ExprReturn {
let attrs = i.attrs;
let return_token = i.return_token;
let expr = i.expr;
let new_code = if expr.is_none() {
quote! {
#(#attrs)*
#return_token ((), gpu)
}
} else {
quote! {
#(#attrs)*
#return_token (#expr, gpu)
}
};
let new_ast = syn::parse_str::<ExprReturn>(&new_code.to_string())
.expect("could not modify return statements");
new_ast
}
fn fold_expr_closure(&mut self, i: ExprClosure) -> ExprClosure {
i
}
fn fold_item(&mut self, i: Item) -> Item {
i
}
}
pub fn modify_returns_for_helper_function(
input: TokenStream,
) -> Result<TokenStream, Vec<Error>> {
let maybe_ast = syn::parse::<ItemFn>(input.clone());
if let Ok(ast) = maybe_ast {
let mut helper_function_return_modifier = HelperFunctionReturnModifier {};
let new_ast = helper_function_return_modifier.fold_item_fn(ast);
Ok(new_ast.to_token_stream().into())
} else {
Err(vec![Error::new(
Span::call_site().unwrap().into(),
"only functions that are items can be tagged with `#[gpu_use]`",
)])
}
}
pub fn modify_for_not_a_helper_function(input: TokenStream) -> Result<TokenStream, Vec<Error>> {
let maybe_ast = syn::parse::<ItemFn>(input.clone());
if let Ok(mut ast) = maybe_ast {
let existing_body = ast.block;
let body = quote! {
{
use ocl::*;
let mut gpu = {
let new_platform = ocl::Platform::default();
let new_device = ocl::Device::first(new_platform).expect("no GPU found");
let new_context = ocl::Context::builder()
.platform(new_platform)
.devices(new_device.clone())
.build()
.expect("failed to build context for executing on GPU with OpenCL");
let new_queue = ocl::Queue::new(&new_context, new_device, None)
.expect("failed to create queue of commands to be sent to GPU");
Gpu {
device: new_device,
context: new_context,
queue: new_queue,
buffers: std::collections::HashMap::new(),
programs: std::collections::HashMap::new()
}
};
#existing_body
}
};
ast.block = Box::new(
syn::parse::<Block>(body.into_token_stream().into())
.expect("could not add boilerplate code for initialization of GPU"),
);
Ok(ast.to_token_stream().into())
} else {
Err(vec![Error::new(
Span::call_site().unwrap().into(),
"only functions that are items can be tagged with `#[gpu_use]`",
)])
}
}
pub struct HelperFunctionInvocationModifier {
pub helper_functions: Vec<Ident>,
}
impl Fold for HelperFunctionInvocationModifier {
fn fold_expr(&mut self, ii: Expr) -> Expr {
if let Expr::Call(mut i) = ii {
if let Expr::Path(path) = *i.func.clone() {
let mut is_helper_function_invocation = false;
for helper_function in &self.helper_functions {
if path.path.is_ident(helper_function) {
is_helper_function_invocation = true;
}
}
if is_helper_function_invocation {
let gpu_ident = quote! {gpu}.into();
i.args.insert(0, gpu_ident);
let new_code = quote! {
{
let result = #i;
gpu = result.1;
result.0
}
};
let new_ast = syn::parse_str::<Expr>(&new_code.to_string())
.expect("could not modify invocations of helper functions");
new_ast
} else {
fold_expr_default!(self, i.into())
}
} else {
fold_expr_default!(self, i.into())
}
} else {
fold_expr_default!(self, ii)
}
}
fn fold_item(&mut self, i: Item) -> Item {
i
}
}
pub fn modify_invocations(
input: TokenStream,
helper_functions: Vec<Ident>,
) -> Result<TokenStream, Vec<Error>> {
let maybe_ast = syn::parse::<ItemFn>(input.clone());
if let Ok(ast) = maybe_ast {
let mut helper_function_invocation_modifier = HelperFunctionInvocationModifier {
helper_functions: helper_functions,
};
let new_ast = helper_function_invocation_modifier.fold_item_fn(ast);
Ok(new_ast.to_token_stream().into())
} else {
Err(vec![Error::new(
Span::call_site().unwrap().into(),
"only functions that are items can be tagged with `#[gpu_use]`",
)])
}
}