use std::collections::HashMap;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use crate::token::{MetaInfo, SignedNumType, UnsignedNumType};
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Program<T> {
pub const_deps: HashMap<String, HashMap<String, (T, MetaInfo)>>,
pub const_defs: HashMap<String, ConstDef>,
pub struct_defs: HashMap<String, StructDef>,
pub enum_defs: HashMap<String, EnumDef>,
pub fn_defs: HashMap<String, FnDef<T>>,
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct ConstDef {
pub ty: Type,
pub value: ConstExpr,
pub meta: MetaInfo,
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct ConstExpr(pub ConstExprEnum, pub MetaInfo);
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum ConstExprEnum {
True,
False,
NumUnsigned(u64, UnsignedNumType),
NumSigned(i64, SignedNumType),
ExternalValue {
party: String,
identifier: String,
},
Max(Vec<ConstExpr>),
Min(Vec<ConstExpr>),
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct StructDef {
pub fields: Vec<(String, Type)>,
pub meta: MetaInfo,
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct EnumDef {
pub variants: Vec<Variant>,
pub meta: MetaInfo,
}
impl EnumDef {
pub(crate) fn get_variant(&self, variant_name: &str) -> Option<&Variant> {
self.variants
.iter()
.find(|&variant| variant.variant_name() == variant_name)
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum Variant {
Unit(String),
Tuple(String, Vec<Type>),
}
impl Variant {
pub(crate) fn variant_name(&self) -> &str {
match self {
Variant::Unit(name) => name.as_str(),
Variant::Tuple(name, _) => name.as_str(),
}
}
pub(crate) fn types(&self) -> Option<Vec<Type>> {
match self {
Variant::Unit(_) => None,
Variant::Tuple(_, types) => Some(types.clone()),
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct FnDef<T> {
pub is_pub: bool,
pub identifier: String,
pub ty: Type,
pub params: Vec<ParamDef>,
pub body: Vec<Stmt<T>>,
pub meta: MetaInfo,
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct ParamDef {
pub mutability: Mutability,
pub name: String,
pub ty: Type,
}
#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum Mutability {
Mutable,
Immutable,
}
impl From<bool> for Mutability {
fn from(b: bool) -> Self {
if b {
Mutability::Mutable
} else {
Mutability::Immutable
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum Type {
Bool,
Unsigned(UnsignedNumType),
Signed(SignedNumType),
Fn(Vec<Type>, Box<Type>),
Array(Box<Type>, usize),
ArrayConst(Box<Type>, String),
Tuple(Vec<Type>),
UntypedTopLevelDefinition(String, MetaInfo),
Struct(String),
Enum(String),
}
impl std::fmt::Display for Type {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Type::Bool => f.write_str("bool"),
Type::Unsigned(n) => n.fmt(f),
Type::Signed(n) => n.fmt(f),
Type::Fn(params, ret_ty) => {
f.write_str("(")?;
let mut params = params.iter();
if let Some(param) = params.next() {
param.fmt(f)?;
}
for param in params {
f.write_str(", ")?;
param.fmt(f)?;
}
f.write_str(") -> ")?;
ret_ty.fmt(f)
}
Type::Array(ty, size) => {
f.write_str("[")?;
ty.fmt(f)?;
f.write_str("; ")?;
size.fmt(f)?;
f.write_str("]")
}
Type::ArrayConst(ty, size) => {
f.write_str("[")?;
ty.fmt(f)?;
f.write_str("; ")?;
size.fmt(f)?;
f.write_str("]")
}
Type::Tuple(fields) => {
f.write_str("(")?;
let mut fields = fields.iter();
if let Some(field) = fields.next() {
field.fmt(f)?;
}
for field in fields {
f.write_str(", ")?;
field.fmt(f)?;
}
f.write_str(")")
}
Type::UntypedTopLevelDefinition(name, _) => f.write_str(name),
Type::Struct(name) => f.write_str(name),
Type::Enum(name) => f.write_str(name),
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Stmt<T> {
pub inner: StmtEnum<T>,
pub meta: MetaInfo,
}
impl<T> Stmt<T> {
pub fn new(inner: StmtEnum<T>, meta: MetaInfo) -> Self {
Self { inner, meta }
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum StmtEnum<T> {
Let(Pattern<T>, Option<Type>, Expr<T>),
LetMut(String, Option<Type>, Expr<T>),
VarAssign(String, Vec<(Accessor<T>, MetaInfo)>, Expr<T>),
ForEachLoop(Pattern<T>, Expr<T>, Vec<Stmt<T>>),
JoinLoop(Pattern<T>, T, (Expr<T>, Expr<T>), Vec<Stmt<T>>),
Expr(Expr<T>),
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum Accessor<T> {
ArrayAccess {
array_ty: T,
index: Expr<T>,
},
TupleAccess {
tuple_ty: T,
index: usize,
},
StructAccess {
struct_ty: T,
field: String,
},
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Expr<T> {
pub inner: ExprEnum<T>,
pub meta: MetaInfo,
pub ty: T,
}
impl Expr<()> {
pub fn untyped(expr: ExprEnum<()>, meta: MetaInfo) -> Self {
Self {
inner: expr,
meta,
ty: (),
}
}
}
impl Expr<Type> {
pub fn typed(expr: ExprEnum<Type>, ty: Type, meta: MetaInfo) -> Self {
Self {
inner: expr,
meta,
ty,
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum ExprEnum<T> {
True,
False,
NumUnsigned(u64, UnsignedNumType),
NumSigned(i64, SignedNumType),
Identifier(String),
ArrayLiteral(Vec<Expr<T>>),
ArrayRepeatLiteral(Box<Expr<T>>, usize),
ArrayRepeatLiteralConst(Box<Expr<T>>, String),
ArrayAccess(Box<Expr<T>>, Box<Expr<T>>),
TupleLiteral(Vec<Expr<T>>),
TupleAccess(Box<Expr<T>>, usize),
StructAccess(Box<Expr<T>>, String),
StructLiteral(String, Vec<(String, Expr<T>)>),
EnumLiteral(String, String, VariantExprEnum<T>),
Match(Box<Expr<T>>, Vec<(Pattern<T>, Expr<T>)>),
UnaryOp(UnaryOp, Box<Expr<T>>),
Op(Op, Box<Expr<T>>, Box<Expr<T>>),
Block(Vec<Stmt<T>>),
FnCall(String, Vec<Expr<T>>),
If(Box<Expr<T>>, Box<Expr<T>>, Box<Expr<T>>),
Cast(Type, Box<Expr<T>>),
Range(u64, u64, UnsignedNumType),
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum VariantExprEnum<T> {
Unit,
Tuple(Vec<Expr<T>>),
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Pattern<T>(pub PatternEnum<T>, pub MetaInfo, pub T);
impl Pattern<()> {
pub fn untyped(pattern: PatternEnum<()>, meta: MetaInfo) -> Self {
Self(pattern, meta, ())
}
}
impl Pattern<Type> {
pub fn typed(pattern: PatternEnum<Type>, ty: Type, meta: MetaInfo) -> Self {
Self(pattern, meta, ty)
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum PatternEnum<T> {
Identifier(String),
True,
False,
NumUnsigned(u64, UnsignedNumType),
NumSigned(i64, SignedNumType),
Tuple(Vec<Pattern<T>>),
Struct(String, Vec<(String, Pattern<T>)>),
StructIgnoreRemaining(String, Vec<(String, Pattern<T>)>),
EnumUnit(String, String),
EnumTuple(String, String, Vec<Pattern<T>>),
UnsignedInclusiveRange(u64, u64, UnsignedNumType),
SignedInclusiveRange(i64, i64, SignedNumType),
}
impl<T> std::fmt::Display for Pattern<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &self.0 {
PatternEnum::Identifier(name) => f.write_str(name),
PatternEnum::True => f.write_str("true"),
PatternEnum::False => f.write_str("false"),
PatternEnum::NumUnsigned(n, suffix) => f.write_fmt(format_args!("{n}{suffix}")),
PatternEnum::NumSigned(n, suffix) => f.write_fmt(format_args!("{n}{suffix}")),
PatternEnum::Struct(struct_name, fields) => {
f.write_fmt(format_args!("{struct_name} {{ "))?;
let mut fields = fields.iter();
if let Some((field_name, field)) = fields.next() {
f.write_fmt(format_args!("{field_name}: {field}"))?;
}
for (field_name, field) in fields {
f.write_str(", ")?;
f.write_fmt(format_args!("{field_name}: {field}"))?;
}
f.write_str("}")
}
PatternEnum::StructIgnoreRemaining(struct_name, fields) => {
f.write_fmt(format_args!("{struct_name} {{ "))?;
for (field_name, field) in fields.iter() {
f.write_fmt(format_args!("{field_name}: {field}"))?;
f.write_str(", ")?;
}
f.write_str(".. }")
}
PatternEnum::Tuple(fields) => {
f.write_str("(")?;
let mut fields = fields.iter();
if let Some(field) = fields.next() {
field.fmt(f)?;
}
for field in fields {
f.write_str(", ")?;
field.fmt(f)?;
}
f.write_str(")")
}
PatternEnum::EnumUnit(enum_name, variant_name) => {
f.write_fmt(format_args!("{enum_name}::{variant_name}"))
}
PatternEnum::EnumTuple(enum_name, variant_name, fields) => {
f.write_fmt(format_args!("{enum_name}::{variant_name}("))?;
let mut fields = fields.iter();
if let Some(field) = fields.next() {
field.fmt(f)?;
}
for field in fields {
f.write_str(", ")?;
field.fmt(f)?;
}
f.write_str(")")
}
PatternEnum::UnsignedInclusiveRange(min, max, suffix) => {
if min == max {
f.write_fmt(format_args!("{min}{suffix}"))
} else if *min == 0 && Some(*max) == suffix.max() {
f.write_str("_")
} else {
f.write_fmt(format_args!("{min}{suffix}..={max}{suffix}"))
}
}
PatternEnum::SignedInclusiveRange(min, max, suffix) => {
if min == max {
f.write_fmt(format_args!("{min}{suffix}"))
} else if Some(*min) == suffix.min() && Some(*max) == suffix.max() {
f.write_str("_")
} else {
f.write_fmt(format_args!("{min}{suffix}..={max}{suffix}"))
}
}
}
}
}
#[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum UnaryOp {
Not,
Neg,
}
#[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum Op {
Add,
Sub,
Mul,
Div,
Mod,
BitAnd,
BitXor,
BitOr,
GreaterThan,
LessThan,
Eq,
NotEq,
ShiftLeft,
ShiftRight,
ShortCircuitAnd,
ShortCircuitOr,
}
impl std::fmt::Display for Op {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Op::Add => f.write_str("+"),
Op::Sub => f.write_str("-"),
Op::Mul => f.write_str("*"),
Op::Div => f.write_str("/"),
Op::Mod => f.write_str("%"),
Op::BitAnd => f.write_str("&"),
Op::BitXor => f.write_str("^"),
Op::BitOr => f.write_str("|"),
Op::GreaterThan => f.write_str(">"),
Op::LessThan => f.write_str("<"),
Op::Eq => f.write_str("=="),
Op::NotEq => f.write_str("!="),
Op::ShiftLeft => f.write_str("<<"),
Op::ShiftRight => f.write_str(">>"),
Op::ShortCircuitAnd => f.write_str("&&"),
Op::ShortCircuitOr => f.write_str("||"),
}
}
}