use super::*;
#[derive(Debug, Clone)]
pub(crate) struct MemoryRegion {
pub name: String,
pub lower: String,
pub upper: String,
pub buffer: String,
}
pub(crate) type MemoryError = CheckerError;
pub(crate) struct MemoryChecker {
buffers: HashMap<String, String>,
regions: Vec<MemoryRegion>,
}
impl MemoryChecker {
pub fn new() -> Self {
Self {
buffers: HashMap::new(),
regions: Vec::new(),
}
}
pub fn register_buffer(&mut self, name: String, capacity: String) {
self.buffers.insert(name, capacity);
}
pub fn register_region(&mut self, region: MemoryRegion) {
self.regions.push(region);
}
pub fn buffer_names(&self) -> Vec<String> {
self.buffers.keys().cloned().collect()
}
pub fn is_buffer(&self, name: &str) -> bool {
self.buffers.contains_key(name)
}
#[cfg(test)]
pub fn buffer_capacity(&self, name: &str) -> Option<&str> {
self.buffers.get(name).map(|s| s.as_str())
}
pub fn regions(&self) -> &[MemoryRegion] {
&self.regions
}
pub fn check_bounds_in_requires(
&self,
buffer_name: &str,
requires_exprs: &[&SpExpr],
span: &Range<usize>,
) -> Option<MemoryError> {
if !self.is_buffer(buffer_name) {
return None;
}
let has_bounds_check = requires_exprs
.iter()
.any(|expr| self.expr_has_bounds_check(expr, buffer_name));
if has_bounds_check {
None
} else {
Some(MemoryError {
code: "A08101".into(),
message: format!(
"buffer `{buffer_name}` accessed without bounds check: \
add a `requires` clause constraining index/offset \
to be within `{buffer_name}.len`"
),
span: span.clone(),
})
}
}
pub fn check_region_buffers(&self, span: &Range<usize>) -> Vec<MemoryError> {
let mut errors = Vec::new();
for region in &self.regions {
if !self.is_buffer(®ion.buffer) {
errors.push(MemoryError {
code: "A08103".into(),
message: format!(
"ghost region `{}` references non-existent buffer `{}`",
region.name, region.buffer,
),
span: span.clone(),
});
}
}
errors
}
pub fn check_region_containment(
&self,
sub_region: &str,
parent_region: &str,
span: &Range<usize>,
) -> Option<MemoryError> {
let sub = self.regions.iter().find(|r| r.name == sub_region);
let parent = self.regions.iter().find(|r| r.name == parent_region);
match (sub, parent) {
(Some(sub_r), Some(parent_r)) => {
if sub_r.lower.is_empty() || sub_r.upper.is_empty() {
return Some(MemoryError {
code: "A08102".into(),
message: format!(
"sub-region `{sub_region}` has incomplete bounds (lower=`{}`, upper=`{}`)",
sub_r.lower, sub_r.upper,
),
span: span.clone(),
});
}
if parent_r.lower.is_empty() || parent_r.upper.is_empty() {
return Some(MemoryError {
code: "A08102".into(),
message: format!(
"parent region `{parent_region}` has incomplete bounds (lower=`{}`, upper=`{}`)",
parent_r.lower, parent_r.upper,
),
span: span.clone(),
});
}
if sub_r.buffer != parent_r.buffer {
Some(MemoryError {
code: "A08102".into(),
message: format!(
"region `{sub_region}` (on buffer `{}`) cannot be contained in \
region `{parent_region}` (on buffer `{}`): different buffers",
sub_r.buffer, parent_r.buffer,
),
span: span.clone(),
})
} else {
None
}
}
(None, _) => Some(MemoryError {
code: "A08102".into(),
message: format!("sub-region `{sub_region}` is not defined"),
span: span.clone(),
}),
(_, None) => Some(MemoryError {
code: "A08102".into(),
message: format!("parent region `{parent_region}` is not defined"),
span: span.clone(),
}),
}
}
fn expr_has_bounds_check(&self, expr: &SpExpr, buffer_name: &str) -> bool {
match &expr.node {
Expr::BinOp { lhs, op, rhs } => {
match op {
BinOp::Lte | BinOp::Lt => {
self.references_buffer_capacity(rhs, buffer_name)
|| self.references_buffer_capacity(lhs, buffer_name)
}
BinOp::Gte | BinOp::Gt => {
self.references_buffer_capacity(lhs, buffer_name)
|| self.references_buffer_capacity(rhs, buffer_name)
}
BinOp::And => {
self.expr_has_bounds_check(lhs, buffer_name)
|| self.expr_has_bounds_check(rhs, buffer_name)
}
_ => false,
}
}
_ => false,
}
}
fn references_buffer_capacity(&self, expr: &SpExpr, buffer_name: &str) -> bool {
match &expr.node {
Expr::Field(receiver, field) => {
let is_len_field =
field == "len" || field == "capacity" || field == "length" || field == "size";
if is_len_field && let Expr::Ident(name) = &receiver.as_ref().node {
return name == buffer_name;
}
false
}
Expr::Ident(name) => {
if let Some(cap) = self.buffers.get(buffer_name) {
name == cap
} else {
false
}
}
Expr::BinOp { lhs, rhs, .. } => {
self.references_buffer_capacity(lhs, buffer_name)
|| self.references_buffer_capacity(rhs, buffer_name)
}
_ => false,
}
}
}
impl Default for MemoryChecker {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for MemoryChecker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MemoryChecker")
.field("buffers", &self.buffers)
.field("regions", &self.regions)
.finish()
}
}
pub fn expr_references_var(expr: &SpExpr, var_name: &str) -> bool {
struct VarRefChecker<'a> {
target: &'a str,
found: bool,
}
impl ExprVisitor for VarRefChecker<'_> {
fn visit_ident(&mut self, name: &str) {
if name == self.target {
self.found = true;
}
}
fn visit_raw(&mut self, tokens: &[String]) {
if tokens.iter().any(|t| t.trim() == self.target) {
self.found = true;
}
}
}
let mut c = VarRefChecker {
target: var_name,
found: false,
};
c.visit_expr(expr);
c.found
}
#[cfg(test)]
mod tests {
use super::*;
use assura_parser::ast::Spanned;
fn span() -> Range<usize> {
0..10
}
fn ident(s: &str) -> SpExpr {
Spanned::no_span(Expr::Ident(s.to_string()))
}
#[test]
fn register_buffer_and_query() {
let mut mc = MemoryChecker::new();
mc.register_buffer("buf".into(), "buf.len".into());
assert!(mc.is_buffer("buf"));
assert!(!mc.is_buffer("other"));
assert_eq!(mc.buffer_capacity("buf"), Some("buf.len"));
assert_eq!(mc.buffer_names().len(), 1);
}
#[test]
fn bounds_check_present_no_error() {
let mut mc = MemoryChecker::new();
mc.register_buffer("buf".into(), "buf.len".into());
let requires_expr = Spanned::no_span(Expr::BinOp {
lhs: Box::new(ident("offset")),
op: BinOp::Lte,
rhs: Box::new(Spanned::no_span(Expr::Field(
Box::new(ident("buf")),
"len".into(),
))),
});
let result = mc.check_bounds_in_requires("buf", &[&requires_expr], &span());
assert!(result.is_none());
}
#[test]
fn bounds_check_missing_a08101() {
let mut mc = MemoryChecker::new();
mc.register_buffer("buf".into(), "buf.len".into());
let requires_expr = Spanned::no_span(Expr::BinOp {
lhs: Box::new(ident("x")),
op: BinOp::Gt,
rhs: Box::new(ident("y")),
});
let result = mc.check_bounds_in_requires("buf", &[&requires_expr], &span());
assert_eq!(result.unwrap().code.as_ref(), "A08101");
}
#[test]
fn region_references_nonexistent_buffer_a08103() {
let mut mc = MemoryChecker::new();
mc.register_region(MemoryRegion {
name: "r1".into(),
lower: "0".into(),
upper: "10".into(),
buffer: "nonexistent".into(),
});
let errs = mc.check_region_buffers(&span());
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A08103");
}
#[test]
fn region_containment_same_buffer_ok() {
let mut mc = MemoryChecker::new();
mc.register_buffer("buf".into(), "buf.len".into());
mc.register_region(MemoryRegion {
name: "sub".into(),
lower: "0".into(),
upper: "5".into(),
buffer: "buf".into(),
});
mc.register_region(MemoryRegion {
name: "parent".into(),
lower: "0".into(),
upper: "10".into(),
buffer: "buf".into(),
});
let result = mc.check_region_containment("sub", "parent", &span());
assert!(result.is_none());
}
#[test]
fn region_containment_different_buffers_a08102() {
let mut mc = MemoryChecker::new();
mc.register_buffer("a".into(), "a.len".into());
mc.register_buffer("b".into(), "b.len".into());
mc.register_region(MemoryRegion {
name: "sub".into(),
lower: "0".into(),
upper: "5".into(),
buffer: "a".into(),
});
mc.register_region(MemoryRegion {
name: "parent".into(),
lower: "0".into(),
upper: "10".into(),
buffer: "b".into(),
});
let result = mc.check_region_containment("sub", "parent", &span());
assert_eq!(result.unwrap().code.as_ref(), "A08102");
}
#[test]
fn expr_references_var_found() {
assert!(expr_references_var(&ident("target"), "target"));
}
#[test]
fn expr_references_var_not_found() {
assert!(!expr_references_var(&ident("other"), "target"));
}
}