#![deny(missing_docs)]
use std::{collections::HashMap, fs::File, path::PathBuf};
use openapiv3::OpenAPI;
use proc_macro::TokenStream;
use progenitor_impl::{
CrateVers, GenerationSettings, Generator, InterfaceStyle, TagStyle, TypePatch, UnknownPolicy,
};
use quote::{ToTokens, quote};
use schemars::schema::SchemaObject;
use serde::Deserialize;
use serde_tokenstream::{OrderedMap, ParseWrapper};
use syn::LitStr;
use token_utils::TypeAndImpls;
mod token_utils;
#[derive(Debug, Clone, Copy, Deserialize)]
enum RelativeTo {
ManifestDir,
OutDir,
}
#[derive(Debug)]
struct SpecSource {
path: LitStr,
relative_to: RelativeTo,
}
impl syn::parse::Parse for SpecSource {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
#[derive(Deserialize)]
struct SpecSourceStruct {
path: ParseWrapper<LitStr>,
relative_to: RelativeTo,
}
let lookahead = input.lookahead1();
if lookahead.peek(LitStr) {
let path: LitStr = input.parse()?;
Ok(SpecSource {
path,
relative_to: RelativeTo::ManifestDir,
})
} else if lookahead.peek(syn::token::Brace) {
let content;
let brace_token = syn::braced!(content in input);
let stream: proc_macro2::TokenStream = content.parse()?;
let helper: SpecSourceStruct =
serde_tokenstream::from_tokenstream_spanned(&brace_token.span, &stream)?;
Ok(SpecSource {
path: helper.path.into_inner(),
relative_to: helper.relative_to,
})
} else {
Err(lookahead.error())
}
}
}
#[proc_macro]
pub fn generate_api(item: TokenStream) -> TokenStream {
match do_generate_api(item) {
Err(err) => err.to_compile_error().into(),
Ok(out) => out,
}
}
#[derive(Deserialize)]
struct MacroSettings {
spec: ParseWrapper<SpecSource>,
#[serde(default)]
interface: InterfaceStyle,
#[serde(default)]
tags: TagStyle,
inner_type: Option<ParseWrapper<syn::Type>>,
pre_hook: Option<ParseWrapper<ClosureOrPath>>,
pre_hook_async: Option<ParseWrapper<ClosureOrPath>>,
post_hook: Option<ParseWrapper<ClosureOrPath>>,
post_hook_async: Option<ParseWrapper<ClosureOrPath>>,
map_type: Option<ParseWrapper<syn::Type>>,
#[serde(default)]
derives: Vec<ParseWrapper<syn::Path>>,
#[serde(default)]
unknown_crates: UnknownPolicy,
#[serde(default)]
crates: HashMap<CrateName, MacroCrateSpec>,
#[serde(default)]
patch: HashMap<ParseWrapper<syn::Ident>, MacroPatch>,
#[serde(default)]
replace: HashMap<ParseWrapper<syn::Ident>, ParseWrapper<TypeAndImpls>>,
#[serde(default)]
convert: OrderedMap<SchemaObject, ParseWrapper<TypeAndImpls>>,
timeout: Option<u64>,
}
#[derive(Deserialize)]
struct MacroPatch {
#[serde(default)]
rename: Option<String>,
#[serde(default)]
derives: Vec<ParseWrapper<syn::Path>>,
}
impl From<MacroPatch> for TypePatch {
fn from(a: MacroPatch) -> Self {
let mut s = Self::default();
a.rename.iter().for_each(|rename| {
s.with_rename(rename);
});
a.derives.iter().for_each(|derive| {
s.with_derive(derive.to_token_stream().to_string());
});
s
}
}
#[derive(Debug)]
struct ClosureOrPath(proc_macro2::TokenStream);
impl syn::parse::Parse for ClosureOrPath {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let lookahead = input.lookahead1();
if lookahead.peek(syn::token::Paren) {
let group: proc_macro2::Group = input.parse()?;
return syn::parse2::<Self>(group.stream());
}
if let Ok(closure) = input.parse::<syn::ExprClosure>() {
return Ok(Self(closure.to_token_stream()));
}
input
.parse::<syn::Path>()
.map(|path| Self(path.to_token_stream()))
}
}
struct MacroCrateSpec {
original: Option<String>,
version: CrateVers,
}
impl<'de> Deserialize<'de> for MacroCrateSpec {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let ss = String::deserialize(deserializer)?;
let (original, vers_str) = if let Some(ii) = ss.find('@') {
let original_str = &ss[..ii];
let rest = &ss[ii + 1..];
if !is_crate(original_str) {
return Err(<D::Error as serde::de::Error>::invalid_value(
serde::de::Unexpected::Str(&ss),
&"valid crate name",
));
}
(Some(original_str.to_string()), rest)
} else {
(None, ss.as_ref())
};
let Some(version) = CrateVers::parse(vers_str) else {
return Err(<D::Error as serde::de::Error>::invalid_value(
serde::de::Unexpected::Str(&ss),
&"valid version",
));
};
Ok(Self { original, version })
}
}
#[derive(Hash, PartialEq, Eq)]
struct CrateName(String);
impl<'de> Deserialize<'de> for CrateName {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let ss = String::deserialize(deserializer)?;
if is_crate(&ss) {
Ok(Self(ss))
} else {
Err(<D::Error as serde::de::Error>::invalid_value(
serde::de::Unexpected::Str(&ss),
&"valid crate name",
))
}
}
}
fn is_crate(s: &str) -> bool {
!s.contains(|cc: char| !cc.is_alphanumeric() && cc != '_' && cc != '-')
}
fn open_file(path: PathBuf, span: proc_macro2::Span) -> Result<File, syn::Error> {
File::open(path.clone()).map_err(|e| {
let path_str = path.to_string_lossy();
syn::Error::new(span, format!("couldn't read file {}: {}", path_str, e))
})
}
fn do_generate_api(item: TokenStream) -> Result<TokenStream, syn::Error> {
let (spec_source, settings) = if let Ok(spec) = syn::parse::<LitStr>(item.clone()) {
let spec_source = SpecSource {
path: spec,
relative_to: RelativeTo::ManifestDir,
};
(spec_source, GenerationSettings::default())
} else {
let MacroSettings {
spec,
interface,
tags,
inner_type,
pre_hook,
pre_hook_async,
post_hook,
post_hook_async,
map_type,
unknown_crates,
crates,
derives,
patch,
replace,
convert,
timeout,
} = serde_tokenstream::from_tokenstream(&item.into())?;
let spec = spec.into_inner();
let mut settings = GenerationSettings::default();
settings.with_interface(interface);
settings.with_tag(tags);
inner_type.map(|inner_type| settings.with_inner_type(inner_type.to_token_stream()));
pre_hook.map(|pre_hook| settings.with_pre_hook(pre_hook.into_inner().0));
pre_hook_async
.map(|pre_hook_async| settings.with_pre_hook_async(pre_hook_async.into_inner().0));
post_hook.map(|post_hook| settings.with_post_hook(post_hook.into_inner().0));
post_hook_async
.map(|post_hook_async| settings.with_post_hook_async(post_hook_async.into_inner().0));
map_type.map(|map_type| settings.with_map_type(map_type.to_token_stream()));
settings.with_unknown_crates(unknown_crates);
crates.into_iter().for_each(
|(CrateName(crate_name), MacroCrateSpec { original, version })| {
if let Some(original_crate) = original {
settings.with_crate(original_crate, version, Some(&crate_name));
} else {
settings.with_crate(crate_name, version, None);
}
},
);
derives.into_iter().for_each(|derive| {
settings.with_derive(derive.to_token_stream());
});
patch.into_iter().for_each(|(type_name, patch)| {
settings.with_patch(type_name.to_token_stream().to_string(), &patch.into());
});
replace.into_iter().for_each(|(type_name, type_and_impls)| {
let type_name = type_name.to_token_stream();
let (replace_name, impls) = type_and_impls.into_inner().into_name_and_impls();
settings.with_replacement(type_name, replace_name, impls);
});
convert.into_iter().for_each(|(schema, type_and_impls)| {
let (type_name, impls) = type_and_impls.into_inner().into_name_and_impls();
settings.with_conversion(schema, type_name, impls);
});
if let Some(timeout) = timeout {
settings.with_timeout(timeout);
}
(spec, settings)
};
let spec_path = spec_source.path;
let base_dir = match spec_source.relative_to {
RelativeTo::ManifestDir => std::env::var("CARGO_MANIFEST_DIR")
.map_or_else(|_| std::env::current_dir().unwrap(), PathBuf::from),
RelativeTo::OutDir => {
let out_dir = std::env::var("OUT_DIR").map_err(|_| {
syn::Error::new(
spec_path.span(),
"relative_to = OutDir requires OUT_DIR to be set \
(are you using this from a build script?)",
)
})?;
PathBuf::from(out_dir)
}
};
let path = base_dir.join(spec_path.value());
let path_str = path.to_string_lossy();
let mut f = open_file(path.clone(), spec_path.span())?;
let oapi: OpenAPI = match serde_json::from_reader(f) {
Ok(json_value) => json_value,
_ => {
f = open_file(path.clone(), spec_path.span())?;
serde_yaml::from_reader(f).map_err(|e| {
syn::Error::new(
spec_path.span(),
format!("failed to parse {}: {}", path_str, e),
)
})?
}
};
let mut builder = Generator::new(&settings);
let code = builder.generate_tokens(&oapi).map_err(|e| {
syn::Error::new(
spec_path.span(),
format!("generation error for {}: {}", spec_path.value(), e),
)
})?;
let output = quote! {
use progenitor::progenitor_client;
#code
const _: &str = include_str!(#path_str);
};
Ok(output.into())
}