use crate::{
ArrayType,
CompositeType,
FutureType,
Identifier,
IntegerType,
Location,
MappingType,
OptionalType,
Path,
ProgramId,
TupleType,
Type,
TypeInterner,
VectorType,
};
use itertools::Itertools;
use leo_span::Span;
use serde::Serialize;
use snarkvm::prelude::{
LiteralType,
Network,
PlaintextType,
PlaintextType::{Array, ExternalStruct, Literal, Struct},
};
use std::fmt;
#[derive(Clone, Debug, Eq, Serialize)]
pub struct TypeNode {
kind: TypeKind,
pub span: Span,
#[serde(skip)]
type_: Type,
}
impl TypeNode {
pub fn new(interner: &TypeInterner, kind: TypeKind, span: Span) -> Self {
let type_ = interner.intern(&kind);
Self { kind, span, type_ }
}
pub fn unchecked(kind: TypeKind, span: Span) -> Self {
Self { kind, span, type_: Type::default() }
}
pub fn kind(&self) -> &TypeKind {
&self.kind
}
pub fn ty(&self) -> Type {
self.type_
}
pub fn into_parts(self) -> (TypeKind, Span, Type) {
(self.kind, self.span, self.type_)
}
pub fn from_parts(kind: TypeKind, span: Span, type_: Type) -> Self {
Self { kind, span, type_ }
}
pub fn types_equivalent(&self, other: &TypeNode) -> bool {
if self.type_ != Type::ERR && self.type_ == other.type_ {
return true;
}
self.kind.types_equivalent(&other.kind)
}
}
impl Default for TypeNode {
fn default() -> Self {
Self::unchecked(TypeKind::Err, Span::default())
}
}
impl PartialEq for TypeNode {
fn eq(&self, other: &Self) -> bool {
assert!(
self.kind == TypeKind::Err || self.type_ != Type::ERR,
"TypeNode with non-Err kind must have been interned before equality use",
);
assert!(
other.kind == TypeKind::Err || other.type_ != Type::ERR,
"TypeNode with non-Err kind must have been interned before equality use",
);
self.type_ == other.type_
}
}
impl std::hash::Hash for TypeNode {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
assert!(
self.kind == TypeKind::Err || self.type_ != Type::ERR,
"TypeNode with non-Err kind must have been interned before hashing",
);
self.type_.hash(state);
}
}
impl fmt::Display for TypeNode {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
self.kind.fmt(f)
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Hash, Serialize)]
pub enum TypeKind {
Address,
Array(ArrayType),
Boolean,
Composite(CompositeType),
Field,
Future(FutureType),
Group,
Identifier,
DynRecord,
Ident(Identifier),
Integer(IntegerType),
Mapping(MappingType),
Optional(OptionalType),
Scalar,
Signature,
String,
Tuple(TupleType),
Vector(VectorType),
Numeric,
Unit,
#[default]
Err,
}
impl TypeKind {
pub fn types_equivalent(&self, other: &TypeKind) -> bool {
match (self, other) {
(TypeKind::Err, _)
| (_, TypeKind::Err)
| (TypeKind::Address, TypeKind::Address)
| (TypeKind::Boolean, TypeKind::Boolean)
| (TypeKind::Field, TypeKind::Field)
| (TypeKind::Group, TypeKind::Group)
| (TypeKind::Scalar, TypeKind::Scalar)
| (TypeKind::Signature, TypeKind::Signature)
| (TypeKind::String, TypeKind::String)
| (TypeKind::Identifier, TypeKind::Identifier)
| (TypeKind::DynRecord, TypeKind::DynRecord)
| (TypeKind::Unit, TypeKind::Unit) => true,
(TypeKind::Array(left), TypeKind::Array(right)) => {
(match (left.length.as_u32(), right.length.as_u32()) {
(Some(l1), Some(l2)) => l1 == l2,
_ => {
true
}
}) && left.element_type().types_equivalent(right.element_type())
}
(TypeKind::Ident(left), TypeKind::Ident(right)) => left.name == right.name,
(TypeKind::Integer(left), TypeKind::Integer(right)) => left == right,
(TypeKind::Mapping(left), TypeKind::Mapping(right)) => {
left.key.types_equivalent(&right.key) && left.value.types_equivalent(&right.value)
}
(TypeKind::Optional(left), TypeKind::Optional(right)) => left.inner.types_equivalent(&right.inner),
(TypeKind::Tuple(left), TypeKind::Tuple(right)) if left.length() == right.length() => left
.elements()
.iter()
.zip_eq(right.elements().iter())
.all(|(left_type, right_type)| left_type.types_equivalent(right_type)),
(TypeKind::Vector(left), TypeKind::Vector(right)) => {
left.element_type.types_equivalent(&right.element_type)
}
(TypeKind::Composite(left), TypeKind::Composite(right)) => {
if !left.const_arguments.is_empty() || !right.const_arguments.is_empty() {
return true;
}
match (&left.path.try_global_location(), &right.path.try_global_location()) {
(Some(l), Some(r)) => l == r,
_ => false,
}
}
(TypeKind::Future(left), TypeKind::Future(right)) if !left.is_explicit || !right.is_explicit => true,
(TypeKind::Future(left), TypeKind::Future(right)) if left.inputs.len() == right.inputs.len() => left
.inputs()
.iter()
.zip_eq(right.inputs().iter())
.all(|(left_type, right_type)| left_type.types_equivalent(right_type)),
_ => false,
}
}
pub fn from_snarkvm<N: Network>(t: &PlaintextType<N>, program_id: ProgramId) -> Self {
match t {
Literal(lit) => (*lit).into(),
Struct(s) => TypeKind::Composite(CompositeType {
path: {
let ident = Identifier::from(s);
Path::from(ident).to_global(Location::new(program_id.as_symbol(), vec![ident.name]))
},
const_arguments: Vec::new(),
}),
ExternalStruct(l) => TypeKind::Composite(CompositeType {
path: {
let external_program = ProgramId::from(l.program_id());
let name = Identifier::from(l.resource());
Path::from(name)
.with_user_program(external_program)
.to_global(Location::new(external_program.as_symbol(), vec![name.name]))
},
const_arguments: Vec::new(),
}),
Array(array) => TypeKind::Array(ArrayType::from_snarkvm(array, program_id)),
}
}
pub fn to_snarkvm<N: Network>(&self) -> anyhow::Result<PlaintextType<N>> {
match self {
TypeKind::Address => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::Address)),
TypeKind::Boolean => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::Boolean)),
TypeKind::Field => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::Field)),
TypeKind::Group => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::Group)),
TypeKind::Integer(int_type) => match int_type {
IntegerType::U8 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::U8)),
IntegerType::U16 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::U16)),
IntegerType::U32 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::U32)),
IntegerType::U64 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::U64)),
IntegerType::U128 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::U128)),
IntegerType::I8 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::I8)),
IntegerType::I16 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::I16)),
IntegerType::I32 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::I32)),
IntegerType::I64 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::I64)),
IntegerType::I128 => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::I128)),
},
TypeKind::Scalar => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::Scalar)),
TypeKind::Signature => Ok(PlaintextType::Literal(snarkvm::prelude::LiteralType::Signature)),
TypeKind::Array(array_type) => Ok(PlaintextType::<N>::Array(array_type.to_snarkvm()?)),
_ => anyhow::bail!("Converting from type {self} to snarkVM type is not supported"),
}
}
pub fn size_in_bits<N: Network, F0, F1>(
&self,
is_raw: bool,
get_structs: F0,
get_external_structs: F1,
) -> anyhow::Result<usize>
where
F0: Fn(&snarkvm::prelude::Identifier<N>) -> anyhow::Result<snarkvm::prelude::StructType<N>>,
F1: Fn(&snarkvm::prelude::Locator<N>) -> anyhow::Result<snarkvm::prelude::StructType<N>>,
{
match is_raw {
false => self.to_snarkvm::<N>()?.size_in_bits(&get_structs, &get_external_structs),
true => self.to_snarkvm::<N>()?.size_in_bits_raw(&get_structs, &get_external_structs),
}
}
pub fn can_coerce_to(&self, expected: &TypeKind) -> bool {
use TypeKind::*;
match (self, expected) {
(Optional(actual_opt), Optional(expected_opt)) => actual_opt.inner.can_coerce_to(&expected_opt.inner),
(a, Optional(opt)) => a.can_coerce_to(&opt.inner),
(Array(a_arr), Array(e_arr)) => {
let lengths_equal = match (a_arr.length.as_u32(), e_arr.length.as_u32()) {
(Some(l1), Some(l2)) => l1 == l2,
_ => true,
};
lengths_equal && a_arr.element_type().can_coerce_to(e_arr.element_type())
}
_ => self.types_equivalent(expected),
}
}
pub fn is_optional(&self) -> bool {
matches!(self, Self::Optional(_))
}
pub fn is_vector(&self) -> bool {
matches!(self, Self::Vector(_))
}
pub fn is_mapping(&self) -> bool {
matches!(self, Self::Mapping(_))
}
pub fn to_optional(&self) -> TypeKind {
TypeKind::Optional(OptionalType { inner: Box::new(self.clone()) })
}
pub fn is_empty(&self) -> bool {
match self {
TypeKind::Unit => true,
TypeKind::Array(array_type) => {
if let Some(length) = array_type.length.as_u32() {
length == 0
} else {
false
}
}
_ => false,
}
}
pub fn is_valid_const_generic_type(&self) -> bool {
matches!(
self,
TypeKind::Boolean
| TypeKind::Integer(_)
| TypeKind::Address
| TypeKind::Scalar
| TypeKind::Group
| TypeKind::Field
| TypeKind::Identifier
)
}
}
impl From<LiteralType> for TypeKind {
fn from(value: LiteralType) -> Self {
match value {
LiteralType::Identifier => TypeKind::Identifier,
LiteralType::Address => TypeKind::Address,
LiteralType::Boolean => TypeKind::Boolean,
LiteralType::Field => TypeKind::Field,
LiteralType::Group => TypeKind::Group,
LiteralType::U8 => TypeKind::Integer(IntegerType::U8),
LiteralType::U16 => TypeKind::Integer(IntegerType::U16),
LiteralType::U32 => TypeKind::Integer(IntegerType::U32),
LiteralType::U64 => TypeKind::Integer(IntegerType::U64),
LiteralType::U128 => TypeKind::Integer(IntegerType::U128),
LiteralType::I8 => TypeKind::Integer(IntegerType::I8),
LiteralType::I16 => TypeKind::Integer(IntegerType::I16),
LiteralType::I32 => TypeKind::Integer(IntegerType::I32),
LiteralType::I64 => TypeKind::Integer(IntegerType::I64),
LiteralType::I128 => TypeKind::Integer(IntegerType::I128),
LiteralType::Scalar => TypeKind::Scalar,
LiteralType::Signature => TypeKind::Signature,
LiteralType::String => TypeKind::String,
}
}
}
impl fmt::Display for TypeKind {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match *self {
TypeKind::Address => write!(f, "address"),
TypeKind::Identifier => write!(f, "identifier"),
TypeKind::DynRecord => write!(f, "dyn record"),
TypeKind::Array(ref array_type) => write!(f, "{array_type}"),
TypeKind::Boolean => write!(f, "bool"),
TypeKind::Field => write!(f, "field"),
TypeKind::Future(ref future_type) => write!(f, "{future_type}"),
TypeKind::Group => write!(f, "group"),
TypeKind::Ident(ref variable) => write!(f, "{variable}"),
TypeKind::Integer(ref integer_type) => write!(f, "{integer_type}"),
TypeKind::Mapping(ref mapping_type) => write!(f, "{mapping_type}"),
TypeKind::Optional(ref optional_type) => write!(f, "{optional_type}"),
TypeKind::Scalar => write!(f, "scalar"),
TypeKind::Signature => write!(f, "signature"),
TypeKind::String => write!(f, "string"),
TypeKind::Composite(ref composite_type) => write!(f, "{composite_type}"),
TypeKind::Tuple(ref tuple) => write!(f, "{tuple}"),
TypeKind::Vector(ref vector_type) => write!(f, "{vector_type}"),
TypeKind::Numeric => write!(f, "numeric"),
TypeKind::Unit => write!(f, "()"),
TypeKind::Err => write!(f, "error"),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[should_panic(expected = "TypeNode with non-Err kind must have been interned")]
fn eq_panics_on_unchecked_non_err_kind() {
let interner = TypeInterner::new();
let unchecked = TypeNode::unchecked(TypeKind::Address, Span::default());
let interned = TypeNode::new(&interner, TypeKind::Boolean, Span::default());
let _ = unchecked == interned;
}
#[test]
#[should_panic(expected = "TypeNode with non-Err kind must have been interned")]
fn hash_panics_on_unchecked_non_err_kind() {
use std::hash::Hash;
let unchecked = TypeNode::unchecked(TypeKind::Address, Span::default());
let mut hasher = std::collections::hash_map::DefaultHasher::new();
unchecked.hash(&mut hasher);
}
#[test]
fn eq_accepts_defaulted_err_nodes() {
let a = TypeNode::default();
let b = TypeNode::default();
assert_eq!(a, b);
}
}