use shape_ast::ast::{Span, TypeAnnotation};
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct SlotId(pub u16);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct FieldIdx(pub u16);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct BasicBlockId(pub u32);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct Point(pub u32);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct LoanId(pub u32);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ProjectionStep {
Field(FieldIdx),
Index,
}
impl fmt::Display for SlotId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "_{}", self.0)
}
}
impl fmt::Display for BasicBlockId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "bb{}", self.0)
}
}
impl fmt::Display for Point {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "p{}", self.0)
}
}
impl fmt::Display for LoanId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "L{}", self.0)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum Place {
Local(SlotId),
Field(Box<Place>, FieldIdx),
Index(Box<Place>, Box<Operand>),
Deref(Box<Place>),
}
impl Place {
pub fn root_local(&self) -> SlotId {
match self {
Place::Local(slot) => *slot,
Place::Field(base, _) | Place::Index(base, _) | Place::Deref(base) => base.root_local(),
}
}
pub fn is_prefix_of(&self, other: &Place) -> bool {
if self == other {
return true;
}
match other {
Place::Local(_) => false,
Place::Field(base, _) | Place::Index(base, _) | Place::Deref(base) => {
self.is_prefix_of(base)
}
}
}
pub fn conflicts_with(&self, other: &Place) -> bool {
if self.root_local() != other.root_local() {
return false;
}
self.is_prefix_of(other) || other.is_prefix_of(self) || self.overlaps(other)
}
fn overlaps(&self, other: &Place) -> bool {
match (self, other) {
(Place::Local(a), Place::Local(b)) => a == b,
(Place::Field(base_a, field_a), Place::Field(base_b, field_b)) => {
if base_a == base_b {
field_a == field_b
} else {
base_a.overlaps(base_b)
}
}
(Place::Index(base_a, _), Place::Index(base_b, _)) => base_a.overlaps(base_b),
_ => self.is_prefix_of(other) || other.is_prefix_of(self),
}
}
pub fn projection_steps(&self) -> Vec<ProjectionStep> {
let mut steps = Vec::new();
self.collect_projection_steps(&mut steps);
steps
}
fn collect_projection_steps(&self, steps: &mut Vec<ProjectionStep>) {
match self {
Place::Local(_) => {}
Place::Field(base, field) => {
base.collect_projection_steps(steps);
steps.push(ProjectionStep::Field(*field));
}
Place::Index(base, _) => {
base.collect_projection_steps(steps);
steps.push(ProjectionStep::Index);
}
Place::Deref(base) => base.collect_projection_steps(steps),
}
}
}
impl fmt::Display for Place {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Place::Local(slot) => write!(f, "{}", slot),
Place::Field(base, field) => write!(f, "{}.{}", base, field.0),
Place::Index(base, idx) => write!(f, "{}[{}]", base, idx),
Place::Deref(base) => write!(f, "*{}", base),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum Operand {
Copy(Place),
Move(Place),
MoveExplicit(Place),
Constant(MirConstant),
}
impl fmt::Display for Operand {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Operand::Copy(p) => write!(f, "copy {}", p),
Operand::Move(p) => write!(f, "move {}", p),
Operand::MoveExplicit(p) => write!(f, "move! {}", p),
Operand::Constant(c) => write!(f, "{}", c),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum MirConstant {
Int(i64),
Bool(bool),
None,
StringId(u32),
Str(String),
Float(u64),
Decimal(String),
Char(char),
Function(String),
Method(String),
ClosurePlaceholder,
}
impl fmt::Display for MirConstant {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
MirConstant::Int(v) => write!(f, "{}", v),
MirConstant::Bool(v) => write!(f, "{}", v),
MirConstant::None => write!(f, "none"),
MirConstant::StringId(id) => write!(f, "str#{}", id),
MirConstant::Str(s) => write!(f, "\"{}\"", s),
MirConstant::Float(bits) => write!(f, "{}", f64::from_bits(*bits)),
MirConstant::Decimal(s) => write!(f, "{}D", s),
MirConstant::Char(c) => write!(f, "'{}'", c.escape_default()),
MirConstant::Function(name) => write!(f, "fn:{}", name),
MirConstant::Method(name) => write!(f, "method:{}", name),
MirConstant::ClosurePlaceholder => write!(f, "closure_placeholder"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum BorrowKind {
Shared,
Exclusive,
}
impl fmt::Display for BorrowKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
BorrowKind::Shared => write!(f, "&"),
BorrowKind::Exclusive => write!(f, "&mut"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum VariantTag {
Ok,
Err,
Some_,
None_,
}
impl fmt::Display for VariantTag {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
VariantTag::Ok => write!(f, "Ok"),
VariantTag::Err => write!(f, "Err"),
VariantTag::Some_ => write!(f, "Some"),
VariantTag::None_ => write!(f, "None"),
}
}
}
impl VariantTag {
#[inline]
pub fn from_name(name: &str) -> Option<Self> {
match name {
"Ok" => Some(VariantTag::Ok),
"Err" => Some(VariantTag::Err),
"Some" => Some(VariantTag::Some_),
"None" => Some(VariantTag::None_),
_ => None,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Rvalue {
Use(Operand),
Borrow(BorrowKind, Place),
BinaryOp(BinOp, Operand, Operand),
UnaryOp(UnOp, Operand),
Aggregate(Vec<Operand>),
Clone(Operand),
EnumTest {
operand: Operand,
variant: VariantTag,
},
EnumPayload {
operand: Operand,
variant: VariantTag,
},
TypePatternTest {
operand: Operand,
type_annotation: TypeAnnotation,
},
EnumDiscriminantTest {
operand: Operand,
enum_name: Option<String>,
variant_name: String,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BinOp {
Add,
Sub,
Mul,
Div,
Mod,
Pow,
BitAnd,
BitOr,
BitXor,
BitShl,
BitShr,
Eq,
Ne,
Lt,
Le,
Gt,
Ge,
And,
Or,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UnOp {
Neg,
Not,
BitNot,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TaskBoundaryKind {
Detached,
Structured,
}
#[derive(Debug, Clone, PartialEq)]
pub struct MirStatement {
pub kind: StatementKind,
pub span: Span,
pub point: Point,
}
#[derive(Debug, Clone, PartialEq)]
pub enum StatementKind {
Assign(Place, Rvalue),
Drop(Place),
TaskBoundary(Vec<Operand>, TaskBoundaryKind),
ClosureCapture {
closure_slot: SlotId,
operands: Vec<Operand>,
function_id: Option<u16>,
},
ArrayStore {
container_slot: SlotId,
operands: Vec<Operand>,
},
ObjectStore {
container_slot: SlotId,
operands: Vec<Operand>,
field_names: Vec<String>,
schema_id: Option<u32>,
},
EnumStore {
container_slot: SlotId,
operands: Vec<Operand>,
variant_name: Option<String>,
},
Nop,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Terminator {
pub kind: TerminatorKind,
pub span: Span,
}
#[derive(Debug, Clone, PartialEq)]
pub enum TerminatorKind {
Goto(BasicBlockId),
SwitchBool {
operand: Operand,
true_bb: BasicBlockId,
false_bb: BasicBlockId,
},
Call {
func: Operand,
args: Vec<Operand>,
destination: Place,
next: BasicBlockId,
},
Return,
Unreachable,
}
#[derive(Debug, Clone)]
pub struct BasicBlock {
pub id: BasicBlockId,
pub statements: Vec<MirStatement>,
pub terminator: Terminator,
}
#[derive(Debug, Clone)]
pub struct MirFunction {
pub name: String,
pub blocks: Vec<BasicBlock>,
pub num_locals: u16,
pub param_slots: Vec<SlotId>,
pub param_reference_kinds: Vec<Option<BorrowKind>>,
pub local_types: Vec<LocalTypeInfo>,
pub span: Span,
pub field_name_table: std::collections::HashMap<FieldIdx, String>,
pub local_struct_type_names:
std::collections::HashMap<SlotId, String>,
pub local_typed_array_element_types:
std::collections::HashMap<SlotId, shape_value::v2::ConcreteType>,
pub local_declared_scalar_types:
std::collections::HashMap<SlotId, shape_value::v2::ConcreteType>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum LocalTypeInfo {
Copy,
NonCopy,
Unknown,
}
impl MirFunction {
pub fn entry_block(&self) -> BasicBlockId {
BasicBlockId(0)
}
pub fn iter_blocks(&self) -> impl Iterator<Item = &BasicBlock> {
self.blocks.iter()
}
pub fn block(&self, id: BasicBlockId) -> &BasicBlock {
&self.blocks[id.0 as usize]
}
pub fn all_points(&self) -> Vec<(Point, BasicBlockId, usize)> {
let mut points = Vec::new();
for block in &self.blocks {
for (i, stmt) in block.statements.iter().enumerate() {
points.push((stmt.point, block.id, i));
}
}
points
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_place_root_local() {
let p = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(1));
assert_eq!(p.root_local(), SlotId(0));
}
#[test]
fn test_place_prefix() {
let x = Place::Local(SlotId(0));
let xa = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(0));
assert!(x.is_prefix_of(&xa));
assert!(!xa.is_prefix_of(&x));
}
#[test]
fn test_disjoint_fields_no_conflict() {
let xa = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(0));
let xb = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(1));
assert!(!xa.overlaps(&xb));
}
#[test]
fn test_same_field_conflicts() {
let xa1 = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(0));
let xa2 = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(0));
assert!(xa1.conflicts_with(&xa2));
}
#[test]
fn test_different_locals_no_conflict() {
let x = Place::Local(SlotId(0));
let y = Place::Local(SlotId(1));
assert!(!x.conflicts_with(&y));
}
#[test]
fn test_parent_child_conflict() {
let x = Place::Local(SlotId(0));
let xa = Place::Field(Box::new(Place::Local(SlotId(0))), FieldIdx(0));
assert!(x.conflicts_with(&xa));
assert!(xa.conflicts_with(&x));
}
}