use super::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum TaintLabel {
Untrusted,
Validated,
Trusted,
}
impl std::fmt::Display for TaintLabel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TaintLabel::Untrusted => write!(f, "untrusted"),
TaintLabel::Validated => write!(f, "validated"),
TaintLabel::Trusted => write!(f, "trusted"),
}
}
}
pub(crate) fn extract_taint_label_from_tokens(type_tokens: &[String]) -> Option<TaintLabel> {
for window in type_tokens.windows(4) {
if window[0] == "@" && window[1] == "taint" && window[2] == ":" {
return match window[3].as_str() {
"untrusted" => Some(TaintLabel::Untrusted),
"validated" => Some(TaintLabel::Validated),
"trusted" => Some(TaintLabel::Trusted),
_ => None,
};
}
}
for window in type_tokens.windows(2) {
if window[0] == "@" {
return match window[1].as_str() {
"untrusted" => Some(TaintLabel::Untrusted),
"validated" => Some(TaintLabel::Validated),
"trusted" => Some(TaintLabel::Trusted),
_ => None,
};
}
}
None
}
pub(crate) fn extract_taint_label(
type_expr: &Option<assura_parser::ast::TypeExpr>,
) -> Option<TaintLabel> {
let tokens = type_expr
.as_ref()
.map(|t| t.to_tokens())
.unwrap_or_default();
extract_taint_label_from_tokens(&tokens)
}
#[derive(Debug, Clone)]
pub(crate) struct TaintChecker {
labels: HashMap<String, TaintLabel>,
validation_fns: std::collections::HashSet<String>,
trusted_sinks: HashMap<String, Vec<Option<TaintLabel>>>,
}
impl TaintChecker {
pub fn new() -> Self {
let mut validation_fns = std::collections::HashSet::new();
validation_fns.insert("validate".to_string());
validation_fns.insert("sanitize".to_string());
Self {
labels: HashMap::new(),
validation_fns,
trusted_sinks: HashMap::new(),
}
}
pub fn declare(&mut self, name: String, label: TaintLabel) {
self.labels.insert(name, label);
}
pub fn register_validator(&mut self, name: String) {
self.validation_fns.insert(name);
}
pub fn register_trusted_sink(&mut self, name: String, param_labels: Vec<Option<TaintLabel>>) {
self.trusted_sinks.insert(name, param_labels);
}
pub fn get_label(&self, name: &str) -> Option<TaintLabel> {
self.labels.get(name).copied()
}
pub fn has_taint_info(&self) -> bool {
!self.labels.is_empty()
}
pub fn infer_taint(&self, expr: &SpExpr) -> TaintLabel {
match &expr.node {
Expr::Ident(name) => self.get_label(name).unwrap_or(TaintLabel::Trusted),
Expr::Literal(_) => TaintLabel::Trusted,
Expr::Field(receiver, _) => self.infer_taint(receiver),
Expr::BinOp { lhs, rhs, .. } => {
std::cmp::min(self.infer_taint(lhs), self.infer_taint(rhs))
}
Expr::UnaryOp { expr: inner, .. } => self.infer_taint(inner),
Expr::Call { func, args } => {
if let Expr::Ident(name) = &func.as_ref().node
&& self.validation_fns.contains(name)
{
return TaintLabel::Validated;
}
args.iter().fold(TaintLabel::Trusted, |acc, arg| {
std::cmp::min(acc, self.infer_taint(arg))
})
}
Expr::MethodCall {
receiver,
method,
args,
} => {
if self.validation_fns.contains(method) {
return TaintLabel::Validated;
}
let r = self.infer_taint(receiver);
args.iter()
.fold(r, |acc, arg| std::cmp::min(acc, self.infer_taint(arg)))
}
Expr::Index { expr: base, index } => {
std::cmp::min(self.infer_taint(base), self.infer_taint(index))
}
Expr::Old(inner) | Expr::Cast { expr: inner, .. } => self.infer_taint(inner),
Expr::If {
cond,
then_branch,
else_branch,
} => {
let mut r = std::cmp::min(self.infer_taint(cond), self.infer_taint(then_branch));
if let Some(e) = else_branch {
r = std::cmp::min(r, self.infer_taint(e));
}
r
}
Expr::List(items) => items.iter().fold(TaintLabel::Trusted, |a, i| {
std::cmp::min(a, self.infer_taint(i))
}),
Expr::Block(exprs) => exprs.iter().fold(TaintLabel::Trusted, |a, e| {
std::cmp::min(a, self.infer_taint(e))
}),
Expr::Forall { body, .. } | Expr::Exists { body, .. } => self.infer_taint(body),
Expr::Apply { args, .. } => args.iter().fold(TaintLabel::Trusted, |a, arg| {
std::cmp::min(a, self.infer_taint(arg))
}),
Expr::Match { scrutinee, arms } => {
let mut r = self.infer_taint(scrutinee);
for arm in arms {
r = std::cmp::min(r, self.infer_taint(&arm.body));
}
r
}
Expr::Let { value, body, .. } => {
std::cmp::min(self.infer_taint(value), self.infer_taint(body))
}
Expr::Tuple(elems) => elems.iter().fold(TaintLabel::Trusted, |a, e| {
std::cmp::min(a, self.infer_taint(e))
}),
Expr::Ghost(_) | Expr::Raw(_) => TaintLabel::Trusted,
}
}
pub fn check_expr(&self, expr: &SpExpr, span: &Range<usize>) -> Vec<TypeError> {
let mut errors = Vec::new();
self.check_expr_inner(expr, span, &mut errors);
errors
}
fn check_expr_inner(&self, expr: &SpExpr, span: &Range<usize>, errors: &mut Vec<TypeError>) {
match &expr.node {
Expr::Index { expr: base, index } => {
let index_taint = self.infer_taint(index);
if index_taint == TaintLabel::Untrusted {
errors.push(TypeError {
code: "A09101".into(),
message: "tainted data used as array index without validation: \
validate the index before using it to access an array"
.into(),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
self.check_expr_inner(base, span, errors);
self.check_expr_inner(index, span, errors);
}
Expr::Call { func, args } => {
if let Expr::Ident(name) = &func.as_ref().node {
if is_alloc_function(name) {
for arg in args {
if self.infer_taint(arg) == TaintLabel::Untrusted {
errors.push(TypeError {
code: "A09102".into(),
message: format!(
"tainted data used as allocation size without \
validation: argument to `{name}` is untrusted"
),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
}
}
if let Some(param_labels) = self.trusted_sinks.get(name) {
for (i, arg) in args.iter().enumerate() {
let arg_taint = self.infer_taint(arg);
if let Some(Some(required)) = param_labels.get(i)
&& arg_taint < *required
{
errors.push(TypeError {
code: "A09103".into(),
message: format!(
"tainted data flows to trusted sink: \
argument {i} to `{name}` is `{arg_taint}` \
but parameter requires `{required}`"
),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
}
}
}
self.check_expr_inner(func, span, errors);
for arg in args {
self.check_expr_inner(arg, span, errors);
}
}
Expr::BinOp { lhs, rhs, .. } => {
self.check_expr_inner(lhs, span, errors);
self.check_expr_inner(rhs, span, errors);
}
Expr::UnaryOp { expr: inner, .. }
| Expr::Old(inner)
| Expr::Cast { expr: inner, .. }
| Expr::Ghost(inner) => {
self.check_expr_inner(inner, span, errors);
}
Expr::Field(receiver, _) => {
self.check_expr_inner(receiver, span, errors);
}
Expr::MethodCall { receiver, args, .. } => {
self.check_expr_inner(receiver, span, errors);
for arg in args {
self.check_expr_inner(arg, span, errors);
}
}
Expr::If {
cond,
then_branch,
else_branch,
} => {
self.check_expr_inner(cond, span, errors);
self.check_expr_inner(then_branch, span, errors);
if let Some(else_br) = else_branch {
self.check_expr_inner(else_br, span, errors);
}
}
Expr::List(items) => {
for item in items {
self.check_expr_inner(item, span, errors);
}
}
Expr::Block(exprs) => {
for e in exprs {
self.check_expr_inner(e, span, errors);
}
}
Expr::Forall { domain, body, .. } | Expr::Exists { domain, body, .. } => {
self.check_expr_inner(domain, span, errors);
self.check_expr_inner(body, span, errors);
}
Expr::Apply { args, .. } => {
for arg in args {
self.check_expr_inner(arg, span, errors);
}
}
Expr::Match { scrutinee, arms } => {
self.check_expr_inner(scrutinee, span, errors);
for arm in arms {
self.check_expr_inner(&arm.body, span, errors);
}
}
Expr::Let { value, body, .. } => {
self.check_expr_inner(value, span, errors);
self.check_expr_inner(body, span, errors);
}
Expr::Tuple(elems) => {
for e in elems {
self.check_expr_inner(e, span, errors);
}
}
Expr::Ident(_) | Expr::Literal(_) | Expr::Raw(_) => {}
}
}
pub fn check_file(source: &assura_parser::ast::SourceFile) -> Vec<TypeError> {
let mut checker = TaintChecker::new();
let mut has_taint_annotations = false;
for decl in &source.decls {
if !matches!(&decl.node, Decl::FnDef(_) | Decl::Extern(_)) {
continue;
}
let name = decl
.node
.name()
.expect("FnDef and Extern always have names");
if let Some(TaintLabel::Validated) = decl
.node
.return_ty()
.and_then(|ty| extract_taint_label_from_tokens(&ty.to_tokens()))
{
checker.register_validator(name.to_string());
has_taint_annotations = true;
}
let param_labels: Vec<Option<TaintLabel>> = decl
.node
.params()
.iter()
.map(|p| extract_taint_label(&p.ty))
.collect();
if param_labels
.iter()
.any(|l| matches!(l, Some(TaintLabel::Validated | TaintLabel::Trusted)))
{
checker.register_trusted_sink(name.to_string(), param_labels.clone());
has_taint_annotations = true;
}
if param_labels.iter().any(|l| l.is_some()) {
has_taint_annotations = true;
}
}
if !has_taint_annotations {
return Vec::new();
}
let mut errors = Vec::new();
for decl in &source.decls {
match &decl.node {
Decl::FnDef(_) | Decl::Extern(_) | Decl::Bind(_) => {
let mut fn_checker = checker.clone();
for param in decl.node.params() {
if let Some(label) = extract_taint_label(¶m.ty) {
fn_checker.declare(param.name.clone(), label);
}
}
if fn_checker.has_taint_info() {
for clause in decl.node.clauses() {
errors.extend(fn_checker.check_expr(&clause.body, &decl.span));
}
}
}
Decl::Contract(c) => {
if checker.has_taint_info() {
for clause in &c.clauses {
errors.extend(checker.check_expr(&clause.body, &decl.span));
}
}
}
Decl::Service(s) => {
for item in &s.items {
match item {
ServiceItem::Operation { clauses, .. }
| ServiceItem::Query { clauses, .. } => {
for clause in clauses {
errors.extend(checker.check_expr(&clause.body, &decl.span));
}
}
_ => {}
}
}
}
Decl::Block { body, .. } => {
for clause in body {
errors.extend(checker.check_expr(&clause.body, &decl.span));
}
}
Decl::Prophecy(_)
| Decl::CodecRegistry(_)
| Decl::TypeDef(_)
| Decl::EnumDef(_) => {}
}
}
errors
}
}
impl Default for TaintChecker {
fn default() -> Self {
Self::new()
}
}
fn is_alloc_function(name: &str) -> bool {
matches!(
name,
"alloc" | "allocate" | "malloc" | "realloc" | "reserve" | "resize"
)
}
#[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()))
}
fn int_lit(n: i64) -> SpExpr {
Spanned::no_span(Expr::Literal(Literal::Int(n.to_string())))
}
#[test]
fn taint_label_ordering() {
assert!(TaintLabel::Untrusted < TaintLabel::Validated);
assert!(TaintLabel::Validated < TaintLabel::Trusted);
}
#[test]
fn taint_label_display() {
assert_eq!(TaintLabel::Untrusted.to_string(), "untrusted");
assert_eq!(TaintLabel::Validated.to_string(), "validated");
assert_eq!(TaintLabel::Trusted.to_string(), "trusted");
}
#[test]
fn extract_taint_long_form() {
let tokens = vec!["@".into(), "taint".into(), ":".into(), "untrusted".into()];
assert_eq!(
extract_taint_label_from_tokens(&tokens),
Some(TaintLabel::Untrusted)
);
}
#[test]
fn extract_taint_short_form() {
let tokens = vec!["@".into(), "validated".into()];
assert_eq!(
extract_taint_label_from_tokens(&tokens),
Some(TaintLabel::Validated)
);
}
#[test]
fn extract_taint_none() {
let tokens: Vec<String> = vec!["Int".into()];
assert_eq!(extract_taint_label_from_tokens(&tokens), None);
}
#[test]
fn tc_infer_literal_trusted() {
let checker = TaintChecker::new();
assert_eq!(checker.infer_taint(&int_lit(42)), TaintLabel::Trusted);
}
#[test]
fn tc_infer_untrusted_ident() {
let mut checker = TaintChecker::new();
checker.declare("user_input".into(), TaintLabel::Untrusted);
assert_eq!(
checker.infer_taint(&ident("user_input")),
TaintLabel::Untrusted
);
}
#[test]
fn tc_infer_binop_propagates_taint() {
let mut checker = TaintChecker::new();
checker.declare("tainted".into(), TaintLabel::Untrusted);
let expr = Spanned::no_span(Expr::BinOp {
lhs: Box::new(ident("tainted")),
op: BinOp::Add,
rhs: Box::new(int_lit(1)),
});
assert_eq!(checker.infer_taint(&expr), TaintLabel::Untrusted);
}
#[test]
fn tc_infer_validation_fn_produces_validated() {
let checker = TaintChecker::new();
let expr = Spanned::no_span(Expr::Call {
func: Box::new(ident("validate")),
args: vec![ident("raw")],
});
assert_eq!(checker.infer_taint(&expr), TaintLabel::Validated);
}
#[test]
fn tc_check_untrusted_array_index() {
let mut checker = TaintChecker::new();
checker.declare("idx".into(), TaintLabel::Untrusted);
let expr = Spanned::no_span(Expr::Index {
expr: Box::new(ident("arr")),
index: Box::new(ident("idx")),
});
let errs = checker.check_expr(&expr, &span());
assert!(!errs.is_empty());
assert!(errs.iter().any(|e| e.code.as_ref() == "A09101"));
}
#[test]
fn tc_check_validated_array_index_ok() {
let mut checker = TaintChecker::new();
checker.declare("idx".into(), TaintLabel::Validated);
let expr = Spanned::no_span(Expr::Index {
expr: Box::new(ident("arr")),
index: Box::new(ident("idx")),
});
let errs = checker.check_expr(&expr, &span());
assert!(errs.is_empty());
}
#[test]
fn tc_check_untrusted_alloc_size() {
let mut checker = TaintChecker::new();
checker.declare("sz".into(), TaintLabel::Untrusted);
let expr = Spanned::no_span(Expr::Call {
func: Box::new(ident("malloc")),
args: vec![ident("sz")],
});
let errs = checker.check_expr(&expr, &span());
assert!(!errs.is_empty());
assert!(errs.iter().any(|e| e.code.as_ref() == "A09102"));
}
#[test]
fn tc_check_trusted_sink_violation() {
let mut checker = TaintChecker::new();
checker.declare("raw".into(), TaintLabel::Untrusted);
checker.register_trusted_sink("exec_query".into(), vec![Some(TaintLabel::Validated)]);
let expr = Spanned::no_span(Expr::Call {
func: Box::new(ident("exec_query")),
args: vec![ident("raw")],
});
let errs = checker.check_expr(&expr, &span());
assert!(!errs.is_empty());
assert!(errs.iter().any(|e| e.code.as_ref() == "A09103"));
}
#[test]
fn tc_check_trusted_sink_ok() {
let mut checker = TaintChecker::new();
checker.declare("safe".into(), TaintLabel::Validated);
checker.register_trusted_sink("exec_query".into(), vec![Some(TaintLabel::Validated)]);
let expr = Spanned::no_span(Expr::Call {
func: Box::new(ident("exec_query")),
args: vec![ident("safe")],
});
let errs = checker.check_expr(&expr, &span());
assert!(errs.is_empty());
}
#[test]
fn tc_has_taint_info() {
let mut checker = TaintChecker::new();
assert!(!checker.has_taint_info());
checker.declare("x".into(), TaintLabel::Untrusted);
assert!(checker.has_taint_info());
}
#[test]
fn tc_register_custom_validator() {
let mut checker = TaintChecker::new();
checker.register_validator("my_sanitize".into());
let expr = Spanned::no_span(Expr::Call {
func: Box::new(ident("my_sanitize")),
args: vec![ident("raw")],
});
assert_eq!(checker.infer_taint(&expr), TaintLabel::Validated);
}
#[test]
fn is_alloc_fn_known() {
assert!(is_alloc_function("malloc"));
assert!(is_alloc_function("realloc"));
assert!(is_alloc_function("reserve"));
assert!(!is_alloc_function("free"));
}
}