use std::collections::HashMap;
use syn::{Expr, Lit};
use crate::interpreter::bytecode::ScalarTy;
use crate::interpreter::numeric::IntWidth;
pub(super) struct TyEnv<'a> {
locals: &'a HashMap<String, ScalarTy>,
fn_returns: &'a HashMap<String, ScalarTy>,
param: Option<(&'a str, &'a ScalarTy)>,
outer: Option<&'a TyEnv<'a>>,
}
impl<'a> TyEnv<'a> {
pub(super) fn new(
locals: &'a HashMap<String, ScalarTy>,
fn_returns: &'a HashMap<String, ScalarTy>,
) -> Self {
Self {
locals,
fn_returns,
param: None,
outer: None,
}
}
fn with_param<'b>(&'b self, name: &'b str, ty: &'b ScalarTy) -> TyEnv<'b> {
TyEnv {
locals: self.locals,
fn_returns: self.fn_returns,
param: Some((name, ty)),
outer: Some(self),
}
}
fn lookup(&self, name: &str) -> Option<&ScalarTy> {
let mut env = Some(self);
while let Some(current) = env {
if let Some((param, ty)) = current.param
&& param == name
{
return Some(ty);
}
env = current.outer;
}
self.locals.get(name)
}
}
fn closure_param(closure: &syn::ExprClosure, item: Option<ScalarTy>) -> Option<(String, ScalarTy)> {
let mut pattern = closure.inputs.first()?;
let mut stated = None;
if let syn::Pat::Type(typed) = pattern {
stated = ScalarTy::lower(&typed.ty);
pattern = &typed.pat;
}
match pattern {
syn::Pat::Ident(id) => Some((id.ident.to_string(), stated.or(item)?)),
_ => None,
}
}
fn in_closure<T>(
closure: &syn::ExprClosure,
item: Option<ScalarTy>,
env: &TyEnv,
walk: impl Fn(&Expr, &TyEnv) -> Option<T>,
) -> Option<T> {
match closure_param(closure, item) {
Some((name, ty)) => walk(&closure.body, &env.with_param(&name, &ty)),
None => walk(&closure.body, env),
}
}
pub(super) fn option_payload(expr: &Expr, env: &TyEnv) -> Option<ScalarTy> {
match expr {
Expr::Paren(inner) => option_payload(&inner.expr, env),
Expr::Group(inner) => option_payload(&inner.expr, env),
Expr::Block(block) => block_tail(&block.block).and_then(|e| option_payload(e, env)),
Expr::If(sel) => block_tail(&sel.then_branch)
.and_then(|e| option_payload(e, env))
.or_else(|| {
sel.else_branch
.as_ref()
.and_then(|(_, e)| option_payload(e, env))
}),
Expr::Path(path) => {
let segment = path.path.segments.last()?;
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
return turbofish_scalar(Some(args));
}
match env.lookup(&segment.ident.to_string()) {
Some(ScalarTy::Opt(payload)) => Some((**payload).clone()),
_ => None,
}
}
Expr::Call(call) => {
let Expr::Path(path) = &*call.func else {
return None;
};
let last = path.path.segments.last()?;
(last.ident == "Some")
.then(|| call.args.first().and_then(|a| written_ty(a, env)))
.flatten()
}
Expr::MethodCall(call) => match call.method.to_string().as_str() {
"then_some" => call.args.first().and_then(|a| written_ty(a, env)),
"parse" => turbofish_scalar(call.turbofish.as_ref()),
"or" => call
.args
.first()
.and_then(|a| option_payload(a, env))
.or_else(|| option_payload(&call.receiver, env)),
"map" | "and_then" => {
let payload = option_payload(&call.receiver, env)?;
match call.args.first() {
Some(Expr::Closure(closure)) if call.method == "map" => {
in_closure(closure, Some(payload), env, written_ty)
}
Some(Expr::Closure(closure)) => {
in_closure(closure, Some(payload), env, option_payload)
}
_ => None,
}
}
"clone" | "cloned" | "copied" | "take" | "as_ref" | "as_mut" | "filter" | "ok" => {
option_payload(&call.receiver, env)
}
"unwrap_or_default" | "unwrap" | "expect" => {
option_payload(&call.receiver, env)?.payload().cloned()
}
"unwrap_or" => option_payload(&call.receiver, env)
.and_then(|payload| payload.payload().cloned())
.or_else(|| call.args.first().and_then(|a| option_payload(a, env))),
"get" => element_ty(&call.receiver, env).or_else(|| map_value_ty(&call.receiver, env)),
"remove" => map_value_ty(&call.receiver, env),
"to_digit" => Some(ScalarTy::Int(IntWidth::U32)),
"position" | "rposition" => Some(ScalarTy::Int(IntWidth::USize)),
"find" | "rfind" => match call.args.first() {
Some(Expr::Closure(_)) => element_ty(&call.receiver, env),
_ => Some(ScalarTy::Int(IntWidth::USize)),
},
"first" | "last" | "pop" | "next" | "reduce" | "min_by_key" | "max_by_key" => {
element_ty(&call.receiver, env)
}
"min" | "max" if call.args.is_empty() => element_ty(&call.receiver, env),
"checked_add" | "checked_sub" | "checked_mul" | "checked_div" | "checked_rem"
| "checked_neg" | "checked_abs" | "checked_pow" | "checked_shl" | "checked_shr"
| "checked_div_euclid" | "checked_rem_euclid" => {
match written_ty(&call.receiver, env) {
Some(ty @ ScalarTy::Int(_)) => Some(ty),
_ => None,
}
}
_ => None,
},
_ => None,
}
}
fn block_tail(block: &syn::Block) -> Option<&Expr> {
match block.stmts.last()? {
syn::Stmt::Expr(expr, None) => Some(expr),
_ => None,
}
}
fn element_ty(expr: &Expr, env: &TyEnv) -> Option<ScalarTy> {
match expr {
Expr::Paren(inner) => element_ty(&inner.expr, env),
Expr::Group(inner) => element_ty(&inner.expr, env),
Expr::Block(block) => block_tail(&block.block).and_then(|e| element_ty(e, env)),
Expr::Range(range) => range
.start
.as_ref()
.and_then(|e| written_ty(e, env))
.or_else(|| range.end.as_ref().and_then(|e| written_ty(e, env))),
Expr::If(sel) => block_tail(&sel.then_branch)
.and_then(|e| element_ty(e, env))
.or_else(|| {
sel.else_branch
.as_ref()
.and_then(|(_, e)| element_ty(e, env))
}),
Expr::Path(path) => {
let segment = path.path.segments.last()?;
match env.lookup(&segment.ident.to_string()) {
Some(ScalarTy::List(element) | ScalarTy::Set(element)) => Some((**element).clone()),
_ => None,
}
}
Expr::Macro(mac) if mac.mac.path.is_ident("vec") => vec_macro_element(&mac.mac, env),
Expr::Call(_) => match written_ty(expr, env) {
Some(ScalarTy::List(element) | ScalarTy::Set(element)) => Some(*element),
_ => None,
},
Expr::MethodCall(call) => match call.method.to_string().as_str() {
"iter" | "into_iter" | "iter_mut" | "cloned" | "copied" | "clone" | "to_vec"
| "rev" | "filter" | "take" | "skip" | "take_while" | "skip_while" | "peekable"
| "by_ref" => element_ty(&call.receiver, env),
"map" => match call.args.first() {
Some(Expr::Closure(closure)) => {
in_closure(closure, element_ty(&call.receiver, env), env, written_ty)
}
_ => None,
},
"filter_map" => match call.args.first() {
Some(Expr::Closure(closure)) => in_closure(
closure,
element_ty(&call.receiver, env),
env,
option_payload,
),
_ => None,
},
"values" | "into_values" | "values_mut" => map_value_ty(&call.receiver, env),
"chars" => Some(ScalarTy::Char),
"bytes" => Some(ScalarTy::Int(IntWidth::U8)),
"collect" => match turbofish_scalar(call.turbofish.as_ref()) {
Some(ScalarTy::List(element) | ScalarTy::Set(element)) => Some(*element),
_ => None,
},
"unwrap" | "unwrap_or" | "unwrap_or_default" => {
let from_receiver = match option_payload(&call.receiver, env) {
Some(ScalarTy::List(element)) => Some(*element),
_ => None,
};
from_receiver.or_else(|| call.args.first().and_then(|a| element_ty(a, env)))
}
_ => None,
},
_ => None,
}
}
fn vec_macro_element(mac: &syn::Macro, env: &TyEnv) -> Option<ScalarTy> {
use syn::Token;
use syn::punctuated::Punctuated;
if let Ok(elements) = mac.parse_body_with(Punctuated::<Expr, Token![,]>::parse_terminated) {
return elements.iter().find_map(|e| written_ty(e, env));
}
mac.parse_body_with(Punctuated::<Expr, Token![;]>::parse_terminated)
.ok()?
.first()
.and_then(|e| written_ty(e, env))
}
fn keeps_receiver_ty(method: &str) -> bool {
matches!(
method,
"clone"
| "to_ascii_lowercase"
| "to_ascii_uppercase"
| "saturating_add"
| "saturating_sub"
| "saturating_mul"
| "wrapping_add"
| "wrapping_sub"
| "wrapping_mul"
| "rotate_left"
| "rotate_right"
| "rem_euclid"
| "div_euclid"
| "pow"
| "powi"
| "powf"
| "abs"
| "signum"
| "isqrt"
)
}
pub(super) fn written_ty(expr: &Expr, env: &TyEnv) -> Option<ScalarTy> {
match expr {
Expr::Paren(inner) => written_ty(&inner.expr, env),
Expr::Group(inner) => written_ty(&inner.expr, env),
Expr::Block(block) => block_tail(&block.block).and_then(|e| written_ty(e, env)),
Expr::If(sel) => block_tail(&sel.then_branch)
.and_then(|e| written_ty(e, env))
.or_else(|| {
sel.else_branch
.as_ref()
.and_then(|(_, e)| written_ty(e, env))
}),
Expr::Cast(cast) => ScalarTy::lower(&cast.ty),
Expr::Binary(bin) => {
use syn::BinOp::{
Add, And, BitAnd, BitOr, BitXor, Div, Eq, Ge, Gt, Le, Lt, Mul, Ne, Or, Rem, Shl,
Shr, Sub,
};
match bin.op {
Add(_) | Sub(_) | Mul(_) | Div(_) | Rem(_) | BitAnd(_) | BitOr(_) | BitXor(_) => {
written_ty(&bin.left, env).or_else(|| written_ty(&bin.right, env))
}
Shl(_) | Shr(_) => written_ty(&bin.left, env),
Eq(_) | Ne(_) | Lt(_) | Le(_) | Gt(_) | Ge(_) | And(_) | Or(_) => {
Some(ScalarTy::Bool)
}
_ => None,
}
}
Expr::Unary(un) => match un.op {
syn::UnOp::Neg(_) | syn::UnOp::Not(_) => written_ty(&un.expr, env),
_ => None,
},
Expr::Lit(lit) => match &lit.lit {
Lit::Str(_) => Some(ScalarTy::Str),
Lit::Bool(_) => Some(ScalarTy::Bool),
Lit::Char(_) => Some(ScalarTy::Char),
Lit::Int(int) => IntWidth::parse(int.suffix()).map(ScalarTy::Int),
Lit::Float(float) => match float.suffix() {
"f32" => Some(ScalarTy::F32),
"f64" => Some(ScalarTy::F64),
_ => None,
},
_ => None,
},
Expr::MethodCall(call) if keeps_receiver_ty(&call.method.to_string()) => {
written_ty(&call.receiver, env)
}
Expr::MethodCall(call)
if call.method == "collect" && turbofish_scalar(call.turbofish.as_ref()).is_some() =>
{
turbofish_scalar(call.turbofish.as_ref())
}
Expr::MethodCall(call) if call.method == "fold" => {
call.args.first().and_then(|init| written_ty(init, env))
}
Expr::MethodCall(call)
if matches!(
call.method.to_string().as_str(),
"unwrap" | "expect" | "unwrap_or" | "unwrap_or_default"
) && option_payload(&call.receiver, env).is_some() =>
{
option_payload(&call.receiver, env)
}
Expr::Call(_) | Expr::Path(_) | Expr::MethodCall(_) => {
if let Some(payload) = option_payload(expr, env) {
Some(ScalarTy::Opt(Box::new(payload)))
} else if is_none_path(expr) {
Some(ScalarTy::Opt(Box::new(ScalarTy::Other)))
} else if let Some(element) = vec_new_element(expr) {
Some(ScalarTy::List(Box::new(element)))
} else if let Some(container) = container_new_ty(expr) {
Some(container)
} else if is_string_call(expr) {
Some(ScalarTy::Str)
} else if let Expr::Path(path) = expr
&& path.path.segments.len() == 1
&& let Some(declared) = env.lookup(&path.path.segments[0].ident.to_string())
{
Some(declared.clone())
} else {
fn_return_ty(expr, env)
}
}
Expr::Macro(mac) if mac.mac.path.is_ident("vec") => Some(ScalarTy::List(Box::new(
vec_macro_element(&mac.mac, env).unwrap_or(ScalarTy::Other),
))),
_ => None,
}
}
fn fn_return_ty(expr: &Expr, env: &TyEnv) -> Option<ScalarTy> {
let Expr::Call(call) = expr else {
return None;
};
let Expr::Path(path) = &*call.func else {
return None;
};
let segment = path.path.segments.last()?;
env.fn_returns.get(&segment.ident.to_string()).cloned()
}
fn is_string_call(expr: &Expr) -> bool {
let Expr::Call(call) = expr else {
return false;
};
let Expr::Path(path) = &*call.func else {
return false;
};
let mut segments = path.path.segments.iter().rev();
let is_ctor = segments
.next()
.is_some_and(|s| s.ident == "from" || s.ident == "new");
is_ctor && segments.next().is_some_and(|s| s.ident == "String")
}
fn vec_new_element(expr: &Expr) -> Option<ScalarTy> {
let Expr::Call(call) = expr else {
return None;
};
let Expr::Path(path) = &*call.func else {
return None;
};
let mut segments = path.path.segments.iter().rev();
let last = segments.next()?;
if last.ident != "new" {
return None;
}
let container = segments.next()?;
if container.ident != "Vec" && container.ident != "VecDeque" {
return None;
}
let syn::PathArguments::AngleBracketed(args) = &container.arguments else {
return None;
};
turbofish_scalar(Some(args))
}
fn container_new_ty(expr: &Expr) -> Option<ScalarTy> {
let Expr::Call(call) = expr else {
return None;
};
let Expr::Path(path) = &*call.func else {
return None;
};
let mut segments = path.path.segments.iter().rev();
let last = segments.next()?;
if last.ident != "new" {
return None;
}
let container = segments.next()?;
let name = container.ident.to_string();
if !matches!(
name.as_str(),
"HashMap" | "BTreeMap" | "HashSet" | "BTreeSet"
) {
return None;
}
ScalarTy::lower_segment(container)
}
fn map_value_ty(expr: &Expr, env: &TyEnv) -> Option<ScalarTy> {
match expr {
Expr::Paren(inner) => map_value_ty(&inner.expr, env),
Expr::Group(inner) => map_value_ty(&inner.expr, env),
Expr::Block(block) => block_tail(&block.block).and_then(|e| map_value_ty(e, env)),
Expr::Path(path) => {
let segment = path.path.segments.last()?;
match env.lookup(&segment.ident.to_string()) {
Some(ScalarTy::Map(value)) => Some((**value).clone()),
_ => None,
}
}
Expr::Call(_) => match container_new_ty(expr) {
Some(ScalarTy::Map(value)) => Some(*value),
_ => None,
},
Expr::MethodCall(call) if call.method == "clone" => map_value_ty(&call.receiver, env),
Expr::MethodCall(call) if call.method == "collect" => {
match turbofish_scalar(call.turbofish.as_ref()) {
Some(ScalarTy::Map(value)) => Some(*value),
_ => None,
}
}
_ => None,
}
}
fn is_none_path(expr: &Expr) -> bool {
matches!(expr, Expr::Path(path)
if path.path.segments.last().is_some_and(|s| s.ident == "None"))
}
pub(super) fn turbofish_scalar(
args: Option<&syn::AngleBracketedGenericArguments>,
) -> Option<ScalarTy> {
args?
.args
.iter()
.find_map(|arg| match arg {
syn::GenericArgument::Type(ty) => Some(ty),
_ => None,
})
.and_then(ScalarTy::lower)
}