use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct LocalVarId(
pub u32,
);
pub struct LocalVarInterner<'str> {
map: HashMap<&'str str, LocalVarId>,
count: u32,
}
impl<'str> Default for LocalVarInterner<'str> {
fn default() -> Self {
Self::new()
}
}
impl<'str> LocalVarInterner<'str> {
pub fn new() -> Self {
Self {
map: HashMap::new(),
count: 0,
}
}
pub fn intern(&mut self, name: &'str str) -> LocalVarId {
if let Some(&id) = self.map.get(name) {
return id;
}
let id = LocalVarId(self.count);
self.count += 1;
self.map.insert(name, id);
id
}
pub fn count(&self) -> u32 {
self.count
}
pub fn get(&self, name: &str) -> Option<LocalVarId> {
self.map.get(name).copied()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum Builtin {
Carry,
Scarry,
Sborrow,
Nan,
Abs,
Sqrt,
Floor,
Ceil,
Round,
Int2Float,
Float2Float,
Trunc,
Zext,
Sext,
Popcount,
Lzcount,
Cpool,
NewObject,
}
impl Builtin {
pub const ALL: &'static [Builtin] = &[
Builtin::Carry,
Builtin::Scarry,
Builtin::Sborrow,
Builtin::Nan,
Builtin::Abs,
Builtin::Sqrt,
Builtin::Floor,
Builtin::Ceil,
Builtin::Round,
Builtin::Int2Float,
Builtin::Float2Float,
Builtin::Trunc,
Builtin::Zext,
Builtin::Sext,
Builtin::Popcount,
Builtin::Lzcount,
Builtin::Cpool,
Builtin::NewObject,
];
pub fn as_str(self) -> &'static str {
match self {
Builtin::Carry => "carry",
Builtin::Scarry => "scarry",
Builtin::Sborrow => "sborrow",
Builtin::Nan => "nan",
Builtin::Abs => "abs",
Builtin::Sqrt => "sqrt",
Builtin::Floor => "floor",
Builtin::Ceil => "ceil",
Builtin::Round => "round",
Builtin::Int2Float => "int2float",
Builtin::Float2Float => "float2float",
Builtin::Trunc => "trunc",
Builtin::Zext => "zext",
Builtin::Sext => "sext",
Builtin::Popcount => "popcount",
Builtin::Lzcount => "lzcount",
Builtin::Cpool => "cpool",
Builtin::NewObject => "newobject",
}
}
pub fn from_name(s: &str) -> Option<Self> {
Self::ALL.iter().copied().find(|b| b.as_str() == s)
}
}
#[cfg(test)]
mod builtin_tests {
use super::Builtin;
#[test]
fn every_builtin_round_trips_through_all() {
for &builtin in Builtin::ALL {
assert_eq!(Builtin::from_name(builtin.as_str()), Some(builtin));
}
assert_eq!(Builtin::from_name("not_a_builtin"), None);
assert_eq!(Builtin::from_name("epsilon"), None);
assert_eq!(Builtin::from_name("float2int"), None);
}
}