#![allow(incomplete_features)]
#![feature(iter_intersperse)]
#![feature(adt_const_params)]
#![feature(proc_macro_diagnostic)]
extern crate proc_macro;
use proc_macro::Diagnostic;
use std::collections::HashMap;
use syn::spanned::Spanned;
pub mod derivatives;
pub use derivatives::*;
mod dict;
pub use dict::*;
pub mod traits;
use traits::*;
pub const FUNCTION_PREFIX: &str = "f";
pub const RECEIVER_PREFIX: &str = "r";
pub const RETURN_SUFFIX: &str = "rtn";
pub fn append_insert(key: &str, value: String, map: &mut HashMap<String, Vec<String>>) {
if let Some(entry) = map.get_mut(key) {
entry.push(value);
} else {
map.insert(String::from(key), vec![value]);
}
}
pub fn expr_type(
expr: &syn::Expr,
type_map: &HashMap<String, String>,
) -> Result<String, PassError> {
match expr {
syn::Expr::Path(path_expr) => {
let var = path_expr.path.segments[0].ident.to_string();
match type_map.get(&var) {
Some(ident) => Ok(ident.clone()),
None => {
Diagnostic::spanned(
path_expr.span().unwrap(),
proc_macro::Level::Error,
format!("variable not found in type map ({:?})", type_map),
)
.emit();
Err(String::from("expr_type"))
}
}
}
syn::Expr::Lit(lit_expr) => literal_type(lit_expr),
syn::Expr::Call(call_expr) => {
let function_sig = pass!(function_signature(call_expr, type_map), "expr_type");
match SUPPORTED_FUNCTIONS.get(&function_sig) {
Some(out_sig) => Ok(out_sig.output_type.clone()),
None => {
let error = format!("expr_type: unsupported function: {}", function_sig);
Diagnostic::spanned(call_expr.span().unwrap(), proc_macro::Level::Error, error)
.emit();
Err(String::from("expr_type"))
}
}
}
syn::Expr::MethodCall(method_expr) => {
let method_sig = pass!(method_signature(method_expr, type_map), "expr_type");
match SUPPORTED_METHODS.get(&method_sig) {
Some(out_sig) => Ok(out_sig.output_type.clone()),
None => {
let error = format!("unsupported method: {}", method_sig);
Diagnostic::spanned(
method_expr.span().unwrap(),
proc_macro::Level::Error,
error,
)
.emit();
Err(String::from("expr_type"))
}
}
}
syn::Expr::Binary(bin_expr) => {
let operation_sig = match operation_signature(bin_expr, type_map) {
Ok(types) => types,
Err(e) => return Err(e),
};
match SUPPORTED_OPERATIONS.get(&operation_sig) {
Some(out_sig) => Ok(out_sig.output_type.clone()),
None => {
Diagnostic::spanned(
bin_expr.span().unwrap(),
proc_macro::Level::Error,
format!("expr_type: unsupported binary operation: {}", operation_sig),
)
.emit();
Err(String::from("expr_type"))
}
}
}
_ => {
Diagnostic::spanned(
expr.span().unwrap(),
proc_macro::Level::Error,
"expr_type: unsupported expression type",
)
.emit();
Err(String::from("expr_type"))
}
}
}
pub fn literal_type(expr_lit: &syn::ExprLit) -> Result<String, PassError> {
match &expr_lit.lit {
syn::Lit::Float(float_lit) => {
let float_str = float_lit.to_string();
let n = float_str.len();
if n <= 3 {
Diagnostic::spanned(
expr_lit.span().unwrap(),
proc_macro::Level::Error,
"All literals need a type suffix e.g. `10.2f32` -- Bad float literal (len)",
)
.emit();
return Err(String::from("literal_type"));
}
let float_type_str = &float_str[n - 3..n];
if !(float_type_str == "f32" || float_type_str == "f64") {
Diagnostic::spanned(
expr_lit.span().unwrap(),
proc_macro::Level::Error,
"All literals need a type suffix e.g. `10.2f32` -- Bad float literal (type)",
)
.emit();
return Err(String::from("literal_type"));
}
Ok(String::from(float_type_str))
}
syn::Lit::Int(int_lit) => {
let int_str = int_lit.to_string();
let n = int_str.len();
let large_type = if n > 4 {
let large_int_str = &int_str[n - 4..n];
match large_int_str {
"i128" | "u128" => Some(String::from(large_int_str)),
_ => None,
}
} else {
None
};
let standard_type = if n > 3 {
let standard_int_str = &int_str[n - 3..n];
match standard_int_str {
"u16" | "u32" | "u64" | "i16" | "i32" | "i64" | "f32" | "f64" => {
Some(String::from(standard_int_str))
}
_ => None,
}
} else {
None
};
let short_type = if n > 2 {
let short_int_str = &int_str[n - 2..n];
match short_int_str {
"i8" | "u8" => Some(String::from(short_int_str)),
_ => None,
}
} else {
None
};
match large_type.or(standard_type).or(short_type) {
Some(int_lit_some) => Ok(int_lit_some),
None => {
Diagnostic::spanned(
expr_lit.span().unwrap(),
proc_macro::Level::Error,
"All literals need a type suffix e.g. `10.2f32` -- Bad integer literal",
)
.emit();
Err(String::from("literal_type"))
}
}
}
_ => {
Diagnostic::spanned(
expr_lit.span().unwrap(),
proc_macro::Level::Error,
"Unsupported literal (only integer and float literals are supported)",
)
.emit();
Err(String::from("literal_type"))
}
}
}
#[macro_export]
macro_rules! rtn {
($a:expr) => {{
format!("{}{}", rust_ad_consts::REVERSE_RETURN_DERIVATIVE, $a)
}};
}
#[macro_export]
macro_rules! der {
($a:expr) => {{
format!("{}{}", rust_ad_consts::DERIVATIVE_PREFIX, $a)
}};
}
#[macro_export]
macro_rules! wrtn {
($a:expr, $b:expr, $c: expr) => {{
format!("{}_wrt_{}_{}", $a, $b, $c)
}};
}
#[macro_export]
macro_rules! wrt {
($a:expr,$b:expr) => {{
format!("{}_wrt_{}", $a, $b)
}};
}
#[macro_export]
macro_rules! pass {
($result: expr,$prefix:expr) => {
match $result {
Ok(res) => res,
Err(err) => {
return Err(format!("{}->{}", $prefix, err));
}
}
};
}
pub type PassError = String;
pub fn method_signature(
method_expr: &syn::ExprMethodCall,
type_map: &HashMap<String, String>,
) -> Result<MethodSignature, PassError> {
let method_str = method_expr.method.to_string();
let receiver_type_str = pass!(
expr_type(&*method_expr.receiver, type_map),
"method_signature"
);
let arg_types_res = method_expr
.args
.iter()
.map(|p| expr_type(p, type_map))
.collect::<Result<Vec<_>, _>>();
let arg_types = pass!(arg_types_res, "method_signature");
Ok(MethodSignature::new(
method_str,
receiver_type_str,
arg_types,
))
}
pub fn function_signature(
function_expr: &syn::ExprCall,
type_map: &HashMap<String, String>,
) -> Result<FunctionSignature, PassError> {
let arg_types_res = function_expr
.args
.iter()
.map(|arg| expr_type(arg, type_map))
.collect::<Result<Vec<_>, _>>();
let arg_types = pass!(arg_types_res, "function_signature");
let func_ident_str = function_expr
.func
.path()
.expect("function_signature: func not path")
.path
.segments[0]
.ident
.to_string();
Ok(FunctionSignature::new(func_ident_str, arg_types))
}
pub fn operation_signature(
operation_expr: &syn::ExprBinary,
type_map: &HashMap<String, String>,
) -> Result<OperationSignature, PassError> {
let (left, right) = (
expr_type(&*operation_expr.left, type_map),
expr_type(&*operation_expr.right, type_map),
);
if left.is_err() {
Diagnostic::spanned(
operation_expr.left.span().unwrap(),
proc_macro::Level::Error,
"operation_signature: unsupported left type",
)
.emit();
}
if right.is_err() {
Diagnostic::spanned(
operation_expr.right.span().unwrap(),
proc_macro::Level::Error,
"operation_signature: unsupported right type",
)
.emit();
}
match (left, right) {
(Ok(l), Ok(r)) => Ok(OperationSignature::from((l, operation_expr.op, r))),
_ => Err(String::from("operation_signature")),
}
}