use crate::{Error, Result};
use syn::{spanned::Spanned, Type};
pub fn from_owned(ty: Type) -> Result<Type> {
let span = ty.span();
let err = || Error::custom("expected an `Option`").with_span(&span);
let Type::Path(path) = ty else {
return Err(err());
};
if path.qself.is_some() {
return Err(err());
}
let Some(last_segment) = path.path.segments.last() else {
return Err(err());
};
if last_segment.ident != "Option" {
return Err(err());
}
let syn::PathArguments::AngleBracketed(ty) = last_segment.clone().arguments else {
return Err(err());
};
let args = ty.args.into_iter().collect::<Vec<_>>();
if args.len() != 1 {
return Err(err());
}
let arg = args
.into_iter()
.next()
.expect("just checked that `.len() == 1`");
let syn::GenericArgument::Type(ty) = arg else {
return Err(err());
};
Ok(ty)
}
pub fn from_mut(ty: &mut Type) -> Result<&mut Type> {
let span = ty.span();
let err = || Error::custom("expected an `Option`").with_span(&span);
let Type::Path(path) = ty else {
return Err(err());
};
if path.qself.is_some() {
return Err(err());
}
let Some(last_segment) = path.path.segments.last_mut() else {
return Err(err());
};
if last_segment.ident != "Option" {
return Err(err());
}
let syn::PathArguments::AngleBracketed(ty) = &mut last_segment.arguments else {
return Err(err());
};
let args = ty.args.iter_mut().collect::<Vec<_>>();
if args.len() != 1 {
return Err(err());
}
let arg = args
.into_iter()
.next()
.expect("just checked that `.len() == 1`");
let syn::GenericArgument::Type(ty) = arg else {
return Err(err());
};
Ok(ty)
}
pub fn from_ref(ty: &Type) -> Result<&Type> {
let span = ty.span();
let err = || Error::custom("expected an `Option`").with_span(&span);
let Type::Path(path) = ty else {
return Err(err());
};
if path.qself.is_some() {
return Err(err());
}
let Some(last_segment) = &path.path.segments.last() else {
return Err(err());
};
if &last_segment.ident != "Option" {
return Err(err());
}
let syn::PathArguments::AngleBracketed(ty) = &last_segment.arguments else {
return Err(err());
};
let args = ty.args.iter().collect::<Vec<_>>();
if args.len() != 1 {
return Err(err());
}
let arg = args
.into_iter()
.next()
.expect("just checked that `.len() == 1`");
let syn::GenericArgument::Type(ty) = arg else {
return Err(err());
};
Ok(ty)
}
#[cfg(test)]
mod tests {
use super::*;
use syn::Type;
macro_rules! test_all {
($test:literal, $method:ident) => {
from_owned(syn::parse_str::<Type>($test).unwrap()).$method();
from_ref(&syn::parse_str::<Type>($test).unwrap()).$method();
from_mut(&mut syn::parse_str::<Type>($test).unwrap()).$method();
};
}
#[test]
fn simple() {
test_all!("Option<String>", unwrap);
}
#[test]
fn fully_qualified() {
test_all!("std::option::Option<String>", unwrap);
test_all!("core::option::Option<String>", unwrap);
}
#[test]
fn absolute_path() {
test_all!("::std::option::Option<String>", unwrap);
test_all!("::core::option::Option<String>", unwrap);
}
#[test]
fn submodule() {
test_all!("option::Option<String>", unwrap);
}
#[test]
fn wrong_arg_count() {
test_all!("Option<String, u8>", unwrap_err);
}
#[test]
fn rejects_qself() {
test_all!("<T as Option>::Option<u32>", unwrap_err);
}
#[test]
fn non_path() {
test_all!("&'a Option<u8>", unwrap_err);
}
}