use std::collections::BTreeMap;
use proc_macro2::{Span, TokenStream};
use quote::{ToTokens, format_ident, quote};
use syn::parse_quote;
use crate::util::{
PrivateField, RumaCommon, RumaCommonReexport, SerdeMetaItem, StructFieldExt, TypeExt,
expand_fields_as_list,
};
#[derive(Default)]
pub(super) struct Headers(BTreeMap<syn::Ident, syn::Field>);
impl Headers {
pub(super) fn insert(&mut self, header: syn::Ident, field: syn::Field) -> syn::Result<()> {
if self.0.contains_key(&header) {
return Err(syn::Error::new(
Span::call_site(),
format!("cannot have multiple values for `{header}` header"),
));
}
self.0.insert(header, field);
Ok(())
}
pub(super) fn expand_fields(&self) -> TokenStream {
expand_fields_as_list(self.0.values())
}
pub(super) fn expand_parse(
&self,
kind: MacroKind,
ruma_common: &RumaCommon,
) -> Option<TokenStream> {
if self.0.is_empty() {
return None;
}
let src = kind.as_variable_ident();
let decls = self
.0
.iter()
.map(|(header_name, field)| Self::expand_parse_header(header_name, field, ruma_common));
Some(quote! {
let headers = #src.headers();
#( #decls )*
})
}
pub(super) fn expand_parse_header(
header_name: &syn::Ident,
field: &syn::Field,
ruma_common: &RumaCommon,
) -> TokenStream {
let ident = field.ident();
let cfg_attrs = field.cfg_attrs();
let header_name_string = header_name.to_string();
let field_type = &field.ty;
let option_inner_type = field_type.option_inner_type();
let some_case = if let Some(field_type) = option_inner_type {
quote! {
str_value.parse::<#field_type>().ok()
}
} else {
quote! {
str_value
.parse::<#field_type>()
.map_err(|e| #ruma_common::api::error::HeaderDeserializationError::InvalidHeader(e.into()))?
}
};
let none_case = if option_inner_type.is_some() {
quote! { None }
} else {
quote! {
return Err(
#ruma_common::api::error::HeaderDeserializationError::MissingHeader(
#header_name_string.into()
).into(),
)
}
};
quote! {
#( #cfg_attrs )*
let #ident = match headers.get(#header_name) {
Some(header_value) => {
let str_value = header_value.to_str()?;
#some_case
}
None => #none_case,
};
}
}
pub(super) fn expand_serialize(
&self,
kind: MacroKind,
http: &TokenStream,
) -> Option<TokenStream> {
if self.0.is_empty() {
return None;
}
let mut serialize = TokenStream::new();
for (header_name, field) in &self.0 {
let ident = field.ident();
let cfg_attrs = field.cfg_attrs();
let header = if field.ty.option_inner_type().is_some() {
quote! {
#( #cfg_attrs )*
if let Some(header_val) = #ident.as_ref() {
headers.insert(
#header_name,
#http::header::HeaderValue::from_str(&header_val.to_string())?,
);
}
}
} else {
quote! {
#( #cfg_attrs )*
headers.insert(
#header_name,
#http::header::HeaderValue::from_str(&#ident.to_string())?,
);
}
};
serialize.extend(header);
}
let src = kind.as_variable_ident();
Some(quote! {{
let headers = #src.headers_mut();
#serialize
}})
}
}
#[derive(Default)]
pub(super) struct Body {
fields: BodyFields,
manual_serde: bool,
}
impl Body {
pub(super) fn push_json_field(&mut self, field: syn::Field) -> syn::Result<()> {
self.fields.push_json_field(field)
}
pub(super) fn set_json_all(&mut self, field: syn::Field) -> syn::Result<()> {
self.fields.set_json_all(field)
}
pub(super) fn set_raw(&mut self, field: syn::Field) -> syn::Result<()> {
self.fields.set_raw(field)
}
pub(super) fn set_manual_serde(&mut self, manual_serde: bool) {
self.manual_serde = manual_serde;
}
pub(super) fn is_empty(&self) -> bool {
matches!(self.fields, BodyFields::Empty)
}
pub(super) fn validate(&self) -> syn::Result<()> {
if let BodyFields::JsonFields(fields) = &self.fields
&& fields.len() == 1
&& let Some(single_field) = fields.first()
&& single_field.has_serde_meta_item(SerdeMetaItem::Flatten)
{
return Err(syn::Error::new_spanned(
single_field,
"Use `#[ruma_api(body)]` to represent the JSON body as a single field",
));
}
if matches!(self.fields, BodyFields::Raw(_)) && self.manual_serde {
return Err(syn::Error::new(
Span::call_site(),
"Cannot have a `manual_body_serde` container attribute with a `raw_body` field attribute",
));
}
Ok(())
}
pub(super) fn type_name(
&self,
kind: MacroKind,
ruma_common: &RumaCommon,
meta_ident: &syn::Ident,
) -> TokenStream {
match &self.fields {
BodyFields::Empty => {
match kind {
MacroKind::Request => {
let http = ruma_common.reexported(RumaCommonReexport::Http);
quote! {
#ruma_common::api::EmptyBody<{
static M: #http::Method = <#meta_ident as #ruma_common::api::Metadata>::METHOD;
match M {
#http::Method::GET => true,
_ => false
}
}>
}
}
MacroKind::Response => {
quote! { #ruma_common::api::EmptyBody<false> }
}
}
}
BodyFields::JsonFields(_) | BodyFields::JsonAll(_) => {
kind.as_struct_ident(StructSuffix::Body).into_token_stream()
}
BodyFields::Raw(_) => quote! { #ruma_common::api::BytesBody },
}
}
pub(super) fn expand_fields(&self) -> Option<TokenStream> {
self.fields.expand_fields()
}
pub(super) fn expand_serde_struct_definition(
&self,
kind: MacroKind,
ruma_common: &RumaCommon,
) -> Option<TokenStream> {
let fields = self.fields.json_fields()?.iter().map(PrivateField);
let ident = kind.as_struct_ident(StructSuffix::Body);
let ruma_macros = ruma_common.reexported(RumaCommonReexport::RumaMacros);
let mut extra_attrs = TokenStream::new();
if !self.manual_serde {
let serde = ruma_common.reexported(RumaCommonReexport::Serde);
let serialize_feature = match kind {
MacroKind::Request => "client",
MacroKind::Response => "server",
};
let deserialize_feature = match kind {
MacroKind::Request => "server",
MacroKind::Response => "client",
};
extra_attrs.extend(quote! {
#[cfg_attr(feature = #serialize_feature, derive(#serde::Serialize))]
#[cfg_attr(feature = #deserialize_feature, derive(#serde::Deserialize))]
});
}
if matches!(self.fields, BodyFields::JsonAll(_)) {
extra_attrs.extend(quote! { #[serde(transparent)] });
}
let outgoing_body_feature = match kind {
MacroKind::Request => "client",
MacroKind::Response => "server",
};
Some(quote! {
#[doc(hidden)]
#[derive(Debug, #ruma_macros::_FakeDeriveRumaApi, #ruma_macros::_FakeDeriveSerde)]
#[cfg_attr(feature = #outgoing_body_feature, derive(#ruma_macros::OutgoingBodyJson))]
#extra_attrs
pub struct #ident { #( #fields ),* }
})
}
pub(super) fn expand_parse(
&self,
kind: MacroKind,
ruma_common: &RumaCommon,
) -> Option<TokenStream> {
match &self.fields {
BodyFields::Empty => None,
BodyFields::JsonFields(fields) => {
Some(Self::expand_parse_json_body(fields, kind, ruma_common))
}
BodyFields::JsonAll(field) => {
Some(Self::expand_parse_json_body(std::slice::from_ref(field), kind, ruma_common))
}
BodyFields::Raw(field) => {
let src = kind.as_variable_ident();
let ident = field.ident();
let cfg_attrs = field.cfg_attrs();
Some(quote! {
#( #cfg_attrs )*
let #ident =
::std::convert::AsRef::<[u8]>::as_ref(#src.body()).to_vec();
})
}
}
}
fn expand_parse_json_body(
fields: &[syn::Field],
kind: MacroKind,
ruma_common: &RumaCommon,
) -> TokenStream {
let src = kind.as_variable_ident();
let body_ident = kind.as_struct_ident(StructSuffix::Body);
let serde_json = ruma_common.reexported(RumaCommonReexport::SerdeJson);
let body_fields = expand_fields_as_list(fields);
quote! {
let body: #body_ident = #serde_json::from_slice(match *#src.body() {
[] => b"{}",
b => b,
})?;
let #body_ident {
#body_fields
} = body;
}
}
pub(super) fn body_expr(&self, kind: MacroKind, ruma_common: &RumaCommon) -> TokenStream {
match &self.fields {
BodyFields::Empty => quote! { #ruma_common::api::EmptyBody },
BodyFields::JsonFields(_) | BodyFields::JsonAll(_) => {
let serde_struct = kind.as_struct_ident(StructSuffix::Body);
let fields = self.expand_fields();
quote! {
#serde_struct { #fields }
}
}
BodyFields::Raw(field) => {
let field_ident = field.ident.to_token_stream();
quote! {
#ruma_common::api::BytesBody(#field_ident)
}
}
}
}
}
#[derive(Default)]
enum BodyFields {
#[default]
Empty,
JsonFields(Vec<syn::Field>),
JsonAll(syn::Field),
Raw(syn::Field),
}
impl BodyFields {
fn push_json_field(&mut self, field: syn::Field) -> syn::Result<()> {
let error_msg = match self {
Self::Empty => {
*self = Self::JsonFields(vec![field]);
return Ok(());
}
Self::JsonFields(fields) => {
fields.push(field);
return Ok(());
}
Self::JsonAll(_) => "cannot have both a `body` field and regular body fields",
Self::Raw(_) => "cannot have both a `raw_body` field and regular body fields",
};
Err(syn::Error::new(Span::call_site(), error_msg))
}
fn set_json_all(&mut self, field: syn::Field) -> syn::Result<()> {
let error_msg = match self {
Self::Empty => {
*self = Self::JsonAll(field);
return Ok(());
}
Self::JsonFields(_) => "cannot have both a `body` field and regular body fields",
Self::JsonAll(_) => "cannot have multiple `body` fields",
Self::Raw(_) => "cannot have both a `raw_body` field and a `body` field",
};
Err(syn::Error::new(Span::call_site(), error_msg))
}
fn set_raw(&mut self, field: syn::Field) -> syn::Result<()> {
let error_msg = match self {
Self::Empty => {
*self = Self::Raw(field);
return Ok(());
}
Self::JsonFields(_) => "cannot have both a `raw_body` field and regular body fields",
Self::JsonAll(_) => "cannot have both a `raw_body` field and a `body` field",
Self::Raw(_) => "cannot have multiple `raw_body` fields",
};
Err(syn::Error::new(Span::call_site(), error_msg))
}
fn json_fields(&self) -> Option<&[syn::Field]> {
let fields = match self {
Self::Empty | Self::Raw(_) => return None,
Self::JsonFields(fields) => fields.as_slice(),
Self::JsonAll(field) => std::slice::from_ref(field),
};
Some(fields)
}
fn expand_fields(&self) -> Option<TokenStream> {
let fields = match self {
Self::Empty => return None,
Self::JsonFields(fields) => fields.as_slice(),
Self::JsonAll(field) => std::slice::from_ref(field),
Self::Raw(field) => std::slice::from_ref(field),
};
Some(expand_fields_as_list(fields))
}
}
#[derive(Clone, Copy)]
pub(super) enum MacroKind {
Request,
Response,
}
impl MacroKind {
pub(super) fn as_variable_ident(&self) -> syn::Ident {
match self {
Self::Request => parse_quote! { request },
Self::Response => parse_quote! { response },
}
}
pub(super) fn as_struct_ident(&self, suffix: StructSuffix) -> syn::Ident {
let prefix = match self {
Self::Request => "Request",
Self::Response => "Response",
};
format_ident!("{prefix}{}", suffix.as_str())
}
}
pub(super) enum StructSuffix {
Body,
Query,
}
impl StructSuffix {
fn as_str(&self) -> &'static str {
match self {
Self::Body => "Body",
Self::Query => "Query",
}
}
}