#![deny(missing_docs)]
#![deny(rustdoc::broken_intra_doc_links)]
extern crate proc_macro;
mod codegen;
mod parse;
use parse::{ToolAttr, ToolDecl, ToolFn};
use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, spanned::Spanned, DeriveInput, FnArg, ItemFn, Pat, Type};
#[proc_macro_attribute]
pub fn tool(attr: TokenStream, item: TokenStream) -> TokenStream {
let attr = parse_macro_input!(attr as ToolAttr);
let item_fn = parse_macro_input!(item as ItemFn);
let func = match ToolFn::parse(item_fn) {
Ok(f) => f,
Err(e) => return e.to_compile_error().into(),
};
let decl = ToolDecl { attr, func };
codegen::expand(decl).into()
}
#[proc_macro_attribute]
pub fn test(_attr: TokenStream, item: TokenStream) -> TokenStream {
let item_fn = parse_macro_input!(item as ItemFn);
if item_fn.sig.asyncness.is_none() {
return syn::Error::new_spanned(
item_fn.sig.fn_token,
"#[klieo::test] requires an async fn",
)
.to_compile_error()
.into();
}
let fn_name = &item_fn.sig.ident;
let body = &item_fn.block;
let attrs = &item_fn.attrs;
let vis = &item_fn.vis;
let inputs = &item_fn.sig.inputs;
let injects_ctx = match inputs.len() {
0 => false,
1 => {
let arg = inputs.first().unwrap();
let pat_ty = match arg {
FnArg::Typed(pt) => pt,
FnArg::Receiver(_) => {
return syn::Error::new_spanned(arg, "#[klieo::test] cannot decorate methods")
.to_compile_error()
.into();
}
};
if !matches!(*pat_ty.pat, Pat::Ident(_)) {
return syn::Error::new_spanned(
&pat_ty.pat,
"#[klieo::test] requires the argument to be `name: TestContext`",
)
.to_compile_error()
.into();
}
if !is_test_context_type(&pat_ty.ty) {
return syn::Error::new_spanned(
&pat_ty.ty,
"#[klieo::test] argument must be `TestContext` (the type-name tail must be `TestContext`)",
)
.to_compile_error()
.into();
}
true
}
_ => {
return syn::Error::new(
inputs.span(),
"#[klieo::test] accepts at most one argument: `ctx: TestContext`",
)
.to_compile_error()
.into();
}
};
let runtime_build = quote! {
::tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("tokio current-thread runtime builds")
};
let expanded = if injects_ctx {
let arg = inputs.first().unwrap();
let pat_ty = match arg {
FnArg::Typed(pt) => pt,
_ => unreachable!("guarded above"),
};
let arg_pat = &pat_ty.pat;
let arg_ty = &pat_ty.ty;
quote! {
#( #attrs )*
#[::core::prelude::v1::test]
#vis fn #fn_name() {
#runtime_build .block_on(async move {
#[allow(unused_mut)]
let mut #arg_pat: #arg_ty =
::klieo::__private::klieo_core::test_utils::TestContext::default();
#body
});
}
}
} else {
quote! {
#( #attrs )*
#[::core::prelude::v1::test]
#vis fn #fn_name() {
#runtime_build .block_on(async move {
#body
});
}
}
};
expanded.into()
}
fn is_test_context_type(ty: &Type) -> bool {
let Type::Path(tp) = ty else {
return false;
};
tp.path
.segments
.last()
.map(|s| s.ident == "TestContext")
.unwrap_or(false)
}
#[proc_macro_derive(KlieoResponse)]
pub fn derive_klieo_response(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let expanded = quote! {
impl #impl_generics ::klieo::__private::klieo_core::response::KlieoResponse for #name #ty_generics
#where_clause
{
fn json_schema() -> ::klieo::__private::serde_json::Value {
let schema = ::klieo::__private::schemars::schema_for!(#name);
::klieo::__private::serde_json::to_value(&schema)
.expect("schemars produces a valid serde_json::Value")
}
}
};
expanded.into()
}