use proc_macro2::{Delimiter, Ident, TokenStream, TokenTree};
use quote::quote;
use crate::apply::Type;
use crate::types::*;
pub(crate) struct Cursor<'a> {
tokens: &'a [TokenTree],
pos: usize,
}
impl<'a> Cursor<'a> {
pub(crate) fn new(tokens: &'a [TokenTree]) -> Self {
Self { tokens, pos: 0 }
}
pub(crate) fn at_end(&self) -> bool {
self.pos >= self.tokens.len()
}
pub(crate) fn peek(&self) -> Option<&'a TokenTree> {
self.tokens.get(self.pos)
}
pub(crate) fn peek_at(&self, offset: usize) -> Option<&'a TokenTree> {
self.tokens.get(self.pos + offset)
}
pub(crate) fn is_punct(&self, ch: char) -> bool {
matches!(self.tokens.get(self.pos), Some(t) if is_punct(t, ch))
}
pub(crate) fn bump(&mut self) {
self.pos += 1;
}
pub(crate) fn pos(&self) -> usize {
self.pos
}
pub(crate) fn slice_since(&self, start: usize) -> &'a [TokenTree] {
&self.tokens[start..self.pos]
}
fn take_segment(&mut self, stop: &[char]) -> &'a [TokenTree] {
let tokens = self.tokens;
let rest = &tokens[self.pos..];
let end = scan_stop(rest, stop).unwrap_or(rest.len());
self.pos += end;
&rest[..end]
}
fn take_rest(&mut self) -> &'a [TokenTree] {
let tokens = self.tokens;
let rest = &tokens[self.pos..];
self.pos = tokens.len();
rest
}
}
fn scan_stop(tokens: &[TokenTree], stop: &[char]) -> Option<usize> {
let mut depth = 0usize;
for (index, token) in tokens.iter().enumerate() {
if is_punct(token, '<') {
depth += 1;
} else if is_punct(token, '>') && !(index > 0 && is_punct(&tokens[index - 1], '-')) {
depth = depth.saturating_sub(1);
} else if depth == 0 && matches!(token, TokenTree::Punct(p) if stop.contains(&p.as_char())) {
let is_arrow_dash = is_punct(token, '-')
&& matches!(tokens.get(index + 1), Some(next) if is_punct(next, '>'));
if !is_arrow_dash {
return Some(index);
}
}
}
None
}
fn is_punct(token: &TokenTree, punctuation: char) -> bool {
matches!(token, TokenTree::Punct(p) if p.as_char() == punctuation)
}
fn contains_punct(tokens: &[TokenTree], punctuation: char) -> bool {
tokens.iter().any(|token| is_punct(token, punctuation))
}
pub(crate) fn parse_item(cursor: &mut Cursor, level: Op, trait_name: Option<&Ident>) -> Option<Ty> {
match level {
Op::Semi | Op::Comma => loop {
if let Some(item) = parse_operand(cursor, level, trait_name) {
return Some(item);
}
if cursor.is_punct(',') {
cursor.bump();
} else {
return None;
}
},
Op::Dash => {
let mut result = parse_operand(cursor, Op::Dash, trait_name)?;
while cursor.is_punct('-') {
cursor.bump();
result = result.apply(parse_operand(cursor, Op::Dash, trait_name)?);
}
Some(result)
}
Op::Caret => {
let mut items = vec![parse_operand(cursor, Op::Caret, trait_name)?];
while cursor.is_punct('^') {
cursor.bump();
items.push(parse_operand(cursor, Op::Caret, trait_name)?);
}
let mut result = items.pop()?;
while let Some(left) = items.pop() {
result = left.apply(result);
}
Some(result)
}
Op::Prim => Some(parse_primitive(cursor.take_rest(), trait_name)),
}
}
fn parse_operand(cursor: &mut Cursor, level: Op, trait_name: Option<&Ident>) -> Option<Ty> {
if cursor.at_end() {
return None;
}
let segment = cursor.take_segment(level.stop_chars());
parse_item(&mut Cursor::new(segment), level.next()?, trait_name)
}
pub(crate) fn parse_primitive(tokens: &[TokenTree], trait_name: Option<&Ident>) -> Ty {
let (tokens, body) = split_trailing_body(tokens);
attach_body(parse_primary(tokens, trait_name), body)
}
fn split_trailing_body(tokens: &[TokenTree]) -> (&[TokenTree], Option<TokenStream>) {
match tokens.last() {
Some(TokenTree::Group(group)) if group.delimiter() == Delimiter::Brace => {
if tokens.len() >= 2
&& let TokenTree::Punct(p) = &tokens[tokens.len() - 2]
&& p.as_char() == '!'
{
return (tokens, None);
}
(&tokens[..tokens.len() - 1], Some(group.stream()))
}
_ => (tokens, None),
}
}
fn attach_body(ty: Ty, body: Option<TokenStream>) -> Ty {
match body {
Some(body) => TyCodeBlock(body).apply(ty),
None => ty,
}
}
fn parse_primary(tokens: &[TokenTree], trait_name: Option<&Ident>) -> Ty {
if let Some((attr, rest)) = parse_attribute(tokens) {
let inner = if rest.is_empty() {
Ty::Attr(TyAttr(attr))
} else {
TyAttr(attr).apply(parse_primitive(rest, trait_name))
};
return inner;
}
if let Some(function) = parse_function(tokens, trait_name) {
return function;
}
if let Some((prefix, rest)) = parse_prefix(tokens) {
let inner = if rest.is_empty() {
Ty::Prefix(prefix)
} else {
prefix.apply(parse_primitive(rest, trait_name))
};
return inner;
}
if let Some(range) = parse_range(tokens) {
return range;
}
if let [TokenTree::Literal(literal)] = tokens
&& let Ok(number) = literal.to_string().parse::<u8>()
{
return Ty::Num(TyNum(number));
}
if let [TokenTree::Group(group)] = tokens {
return parse_group(group, trait_name);
}
if let Some((base, args, rest)) = parse_generic(tokens) {
let params = parse_angle_bracket_contents(args, trait_name);
let generic = if is_trait_base(base, trait_name) {
Ty::Trait(TyTrait(base.iter().cloned().collect(), params))
} else {
if !rest.is_empty() {
return primitive(tokens);
}
Ty::Generic(TyGeneric(Box::new(primitive(base)), params))
};
return if rest.is_empty() {
generic
} else {
generic.apply(parse_primitive(rest, trait_name))
};
}
if let Some((args, rest)) = parse_type_params(tokens) {
let params = Ty::TypeParam(parse_angle_bracket_contents(args, trait_name));
return if rest.is_empty() {
params
} else {
params.apply(parse_primitive(rest, trait_name))
};
}
primitive(tokens)
}
fn parse_attribute(tokens: &[TokenTree]) -> Option<(TokenStream, &[TokenTree])> {
match tokens {
[TokenTree::Punct(hash), TokenTree::Group(group), rest @ ..]
if hash.as_char() == '#' && group.delimiter() == Delimiter::Bracket =>
{
Some((group.stream(), rest))
}
_ => None,
}
}
fn parse_function(tokens: &[TokenTree], trait_name: Option<&Ident>) -> Option<Ty> {
let [TokenTree::Ident(name), TokenTree::Group(args), rest @ ..] = tokens else {
return None;
};
if name != "fn" || args.delimiter() != Delimiter::Parenthesis {
return None;
}
let args_tokens: Vec<_> = args.stream().into_iter().collect();
let mut cursor = Cursor::new(&args_tokens);
let mut parameters = Vec::new();
while let Some(parameter) = parse_item(&mut cursor, Op::Comma, trait_name) {
parameters.push(parameter);
}
let return_type = match rest {
[TokenTree::Punct(dash), TokenTree::Punct(arrow), return_tokens @ ..]
if dash.as_char() == '-' && arrow.as_char() == '>' && !return_tokens.is_empty() =>
{
Some(Box::new(parse_primitive(return_tokens, trait_name)))
}
_ => None,
};
Some(Ty::Fn(TyFn(parameters, return_type)))
}
fn parse_prefix(tokens: &[TokenTree]) -> Option<(TyPrefix, &[TokenTree])> {
match tokens {
[TokenTree::Punct(p), TokenTree::Ident(name), rest @ ..]
if p.as_char() == '&' && name == "mut" =>
{
Some((TyPrefix::RefMut, rest))
}
[TokenTree::Punct(p), rest @ ..] if p.as_char() == '&' => Some((TyPrefix::Ref, rest)),
[TokenTree::Punct(p), TokenTree::Ident(name), rest @ ..]
if p.as_char() == '*' && name == "const" =>
{
Some((TyPrefix::PtrConst, rest))
}
[TokenTree::Punct(p), TokenTree::Ident(name), rest @ ..]
if p.as_char() == '*' && name == "mut" =>
{
Some((TyPrefix::PtrMut, rest))
}
[TokenTree::Ident(name), rest @ ..] if name == "self" => Some((TyPrefix::SelfType, rest)),
[TokenTree::Ident(name), rest @ ..] if name == "fn" => Some((TyPrefix::Fn, rest)),
[TokenTree::Ident(name), rest @ ..] if name == "unsafe" => Some((TyPrefix::Unsafe, rest)),
_ => None,
}
}
fn parse_range(tokens: &[TokenTree]) -> Option<Ty> {
let [TokenTree::Literal(start), TokenTree::Punct(first_dot), TokenTree::Punct(second_dot), rest @ ..] =
tokens
else {
return None;
};
if first_dot.as_char() != '.' || second_dot.as_char() != '.' {
return None;
}
let start = start.to_string().parse::<u8>().ok()?;
let (inclusive, end) = match rest {
[TokenTree::Literal(end)] => (false, end),
[TokenTree::Punct(eq), TokenTree::Literal(end)] if eq.as_char() == '=' => (true, end),
_ => return None,
};
Some(Ty::Range(TyRange {
start,
end: end.to_string().parse().ok()?,
inclusive,
}))
}
fn parse_group(group: &proc_macro2::Group, trait_name: Option<&Ident>) -> Ty {
let contents: Vec<_> = group.stream().into_iter().collect();
match group.delimiter() {
Delimiter::Parenthesis => {
if contents.is_empty() || contains_punct(&contents, ',') {
Ty::Tuple(TyTuple(parse_list(&contents, Op::Comma, trait_name)))
} else {
Ty::Group(TyGroup(Box::new(
parse_item(&mut Cursor::new(&contents), Op::Dash, trait_name)
.unwrap_or_else(empty),
)))
}
}
Delimiter::Bracket => {
if contains_punct(&contents, ',') {
Ty::Array(TyArray(parse_list(&contents, Op::Comma, trait_name)))
} else {
let mut cursor = Cursor::new(&contents);
let element = parse_item(&mut cursor, Op::Semi, trait_name).unwrap_or_else(empty);
if cursor.is_punct(';') {
cursor.bump();
let length = cursor.take_rest().iter().cloned().collect();
Ty::FixedArray(TyFixedArray(Box::new(element), length))
} else {
Ty::Slice(TySlice(Box::new(element)))
}
}
}
Delimiter::Brace => Ty::CodeBlock(TyCodeBlock(group.stream())),
_ => empty(),
}
}
fn parse_list(tokens: &[TokenTree], level: Op, trait_name: Option<&Ident>) -> Vec<Ty> {
let mut cursor = Cursor::new(tokens);
let mut items = Vec::new();
while let Some(item) = parse_item(&mut cursor, level, trait_name) {
items.push(item);
}
items
}
fn parse_generic(tokens: &[TokenTree]) -> Option<(&[TokenTree], &[TokenTree], &[TokenTree])> {
let open = tokens.iter().position(|token| is_punct(token, '<'))?;
if open == 0 {
return None;
}
let close = matching_angle(tokens, open)?;
Some((
&tokens[..open],
&tokens[open + 1..close],
&tokens[close + 1..],
))
}
fn parse_type_params(tokens: &[TokenTree]) -> Option<(&[TokenTree], &[TokenTree])> {
if !matches!(tokens.first(), Some(token) if is_punct(token, '<')) {
return None;
}
let close = matching_angle(tokens, 0)?;
Some((&tokens[1..close], &tokens[close + 1..]))
}
fn matching_angle(tokens: &[TokenTree], open: usize) -> Option<usize> {
let mut depth = 0usize;
for (index, token) in tokens.iter().enumerate().skip(open) {
if is_punct(token, '<') {
depth += 1;
} else if is_punct(token, '>') {
if index > open && is_punct(&tokens[index - 1], '-') {
continue;
}
depth = depth.checked_sub(1)?;
if depth == 0 {
return Some(index);
}
}
}
None
}
fn is_trait_base(base: &[TokenTree], trait_name: Option<&Ident>) -> bool {
trait_name
.is_some_and(|name| matches!(base.last(), Some(TokenTree::Ident(last)) if last == name))
}
fn split_at_depth0(tokens: &[TokenTree], separator: char) -> Vec<&[TokenTree]> {
let mut chunks = Vec::new();
let mut rest = tokens;
while let Some(index) = scan_stop(rest, &[separator]) {
chunks.push(&rest[..index]);
rest = &rest[index + 1..];
}
chunks.push(rest);
chunks
}
fn find_colon_at_depth0(tokens: &[TokenTree]) -> Option<usize> {
scan_stop(tokens, &[':']).filter(|index| {
!(*index > 0 && is_punct(&tokens[*index - 1], ':'))
&& !(*index + 1 < tokens.len() && is_punct(&tokens[*index + 1], ':'))
})
}
fn parse_angle_bracket_contents(tokens: &[TokenTree], trait_name: Option<&Ident>) -> TyTypeParam {
let mut params = Vec::new();
let mut bindings = Vec::new();
for chunk in split_at_depth0(tokens, ',') {
if chunk.is_empty() {
continue;
}
if let Some(eq) = scan_stop(chunk, &['=']) {
bindings.push((
chunk[..eq].iter().cloned().collect(),
chunk[eq + 1..].iter().cloned().collect(),
));
} else if let Some(colon) = find_colon_at_depth0(chunk) {
params.push((
chunk[..colon].iter().cloned().collect(),
Some(
parse_item(&mut Cursor::new(&chunk[colon + 1..]), Op::Dash, trait_name)
.unwrap_or_else(empty),
),
));
} else {
params.push((chunk.iter().cloned().collect(), None));
}
}
TyTypeParam { params, bindings }
}
fn primitive(tokens: &[TokenTree]) -> Ty {
Ty::Primitive(TyPrimitive(tokens.iter().cloned().collect()))
}
fn empty() -> Ty {
Ty::Primitive(TyPrimitive(quote![]))
}