use std::collections::{BTreeMap, BTreeSet};
use super::reader::{find_open_brace, match_brace};
use super::scan::{Token, tokenise};
struct ForbiddenCall {
qualifier: &'static str,
member: &'static str,
kind: ViolationKind,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum ViolationKind {
WallClock,
Entropy,
}
impl ViolationKind {
#[must_use]
pub fn remedy(self) -> &'static str {
match self {
ViolationKind::WallClock => {
"use the recorded `workflow.now()` instead of reading the wall clock"
}
ViolationKind::Entropy => {
"use the seeded `workflow.random()` / `workflow.random_int(..)` \
instead of drawing entropy"
}
}
}
}
const FORBIDDEN_CALLS: &[ForbiddenCall] = &[
ForbiddenCall {
qualifier: "erlang",
member: "system_time",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "erlang",
member: "monotonic_time",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "erlang",
member: "now",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "erlang",
member: "timestamp",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "erlang",
member: "unique_integer",
kind: ViolationKind::Entropy,
},
ForbiddenCall {
qualifier: "os",
member: "system_time",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "os",
member: "timestamp",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "os",
member: "perf_counter",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "os",
member: "erlang_timestamp",
kind: ViolationKind::WallClock,
},
ForbiddenCall {
qualifier: "rand",
member: "uniform",
kind: ViolationKind::Entropy,
},
ForbiddenCall {
qualifier: "rand",
member: "uniform_real",
kind: ViolationKind::Entropy,
},
ForbiddenCall {
qualifier: "rand",
member: "bytes",
kind: ViolationKind::Entropy,
},
ForbiddenCall {
qualifier: "crypto",
member: "strong_rand_bytes",
kind: ViolationKind::Entropy,
},
ForbiddenCall {
qualifier: "float",
member: "random",
kind: ViolationKind::Entropy,
},
ForbiddenCall {
qualifier: "int",
member: "random",
kind: ViolationKind::Entropy,
},
];
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub struct Violation {
pub function: String,
pub call: String,
pub kind: ViolationKind,
}
#[derive(thiserror::Error, Debug, PartialEq, Eq)]
pub enum DeterminismError {
#[error(
"entry function `{function}` is not defined in the workflow source; the determinism \
analysis requires its body to walk from"
)]
EntryFunctionNotFound {
function: String,
},
}
#[derive(Clone, Copy)]
struct FnBody {
start: usize,
end: usize,
}
pub fn analyze_determinism(
source: &str,
entry_function: &str,
) -> Result<Vec<Violation>, DeterminismError> {
let tokens = tokenise(source);
let functions = map_functions(&tokens);
if !functions.contains_key(entry_function) {
return Err(DeterminismError::EntryFunctionNotFound {
function: entry_function.to_owned(),
});
}
let mut violations: BTreeSet<Violation> = BTreeSet::new();
let mut visited: BTreeSet<String> = BTreeSet::new();
walk(
&tokens,
&functions,
entry_function,
&mut visited,
&mut violations,
);
Ok(violations.into_iter().collect())
}
fn walk(
tokens: &[Token],
functions: &BTreeMap<String, FnBody>,
function: &str,
visited: &mut BTreeSet<String>,
violations: &mut BTreeSet<Violation>,
) {
if !visited.insert(function.to_owned()) {
return;
}
let Some(body) = functions.get(function).copied() else {
return;
};
let mut callees: Vec<String> = Vec::new();
let upper = body.end.min(tokens.len());
let mut index = body.start;
let mut depth: usize = 0;
while index < upper {
match &tokens[index] {
Token::OpenParen => depth += 1,
Token::CloseParen => depth = depth.saturating_sub(1),
Token::Qualified { left, right } => {
if let Some(forbidden) = match_forbidden(left, right) {
violations.insert(Violation {
function: function.to_owned(),
call: format!("{left}.{right}"),
kind: forbidden.kind,
});
}
}
Token::Ident(name) if functions.contains_key(name) => {
let applied = matches!(tokens.get(index + 1), Some(Token::OpenParen));
if applied || depth >= 1 {
callees.push(name.clone());
}
}
_ => {}
}
index += 1;
}
for callee in callees {
walk(tokens, functions, &callee, visited, violations);
}
}
fn match_forbidden(qualifier: &str, member: &str) -> Option<&'static ForbiddenCall> {
FORBIDDEN_CALLS
.iter()
.find(|call| call.qualifier == qualifier && call.member == member)
}
fn map_functions(tokens: &[Token]) -> BTreeMap<String, FnBody> {
let mut functions = BTreeMap::new();
let mut index = 0;
while index < tokens.len() {
if matches!(&tokens[index], Token::Ident(word) if word == "fn") {
if let Some(Token::Ident(name)) = tokens.get(index + 1) {
if let Some(open) = find_open_brace(tokens, index + 2, tokens.len()) {
if let Some(close) = match_brace(tokens, open, tokens.len()) {
functions.insert(
name.clone(),
FnBody {
start: open + 1,
end: close,
},
);
index = close + 1;
continue;
}
}
}
}
index += 1;
}
functions
}
#[cfg(test)]
mod tests {
use super::{DeterminismError, ViolationKind, analyze_determinism};
#[test]
fn clean_workflow_has_no_violations() -> Result<(), Box<dyn std::error::Error>> {
let source = "import aion/workflow\n\
pub fn run(input) {\n \
let assert Ok(at) = workflow.now()\n \
let assert Ok(seed) = workflow.random()\n \
workflow.run(wrappers.charge_activity(input))\n}\n";
assert!(analyze_determinism(source, "run")?.is_empty());
Ok(())
}
#[test]
fn direct_wall_clock_call_is_flagged() -> Result<(), Box<dyn std::error::Error>> {
let source = "pub fn run(input) {\n \
let now = erlang.system_time(1000)\n \
workflow.run(wrappers.charge_activity(input))\n}\n";
let violations = analyze_determinism(source, "run")?;
assert_eq!(violations.len(), 1);
assert_eq!(violations[0].call, "erlang.system_time");
assert_eq!(violations[0].kind, ViolationKind::WallClock);
assert_eq!(violations[0].function, "run");
Ok(())
}
#[test]
fn entropy_in_a_reachable_helper_is_flagged() -> Result<(), Box<dyn std::error::Error>> {
let source = "pub fn run(input) {\n \
let id = make_id(input)\n \
workflow.run(wrappers.charge_activity(id))\n}\n\
fn make_id(input) {\n float.random()\n}\n";
let violations = analyze_determinism(source, "run")?;
assert_eq!(violations.len(), 1, "{violations:?}");
assert_eq!(violations[0].call, "float.random");
assert_eq!(violations[0].kind, ViolationKind::Entropy);
assert_eq!(violations[0].function, "make_id");
Ok(())
}
#[test]
fn entropy_in_a_helper_passed_as_a_value_is_flagged() -> Result<(), Box<dyn std::error::Error>>
{
let source = "pub fn run(input) {\n \
let _ = list.map(input, tainted)\n \
workflow.run(wrappers.charge_activity(input))\n}\n\
fn tainted(item) {\n float.random()\n}\n";
let violations = analyze_determinism(source, "run")?;
assert_eq!(violations.len(), 1, "{violations:?}");
assert_eq!(violations[0].call, "float.random");
assert_eq!(violations[0].kind, ViolationKind::Entropy);
assert_eq!(violations[0].function, "tainted");
Ok(())
}
#[test]
fn unreachable_helper_violation_is_not_flagged() -> Result<(), Box<dyn std::error::Error>> {
let source = "pub fn run(input) {\n \
workflow.run(wrappers.charge_activity(input))\n}\n\
fn tainted(input) {\n int.random()\n}\n";
assert!(analyze_determinism(source, "run")?.is_empty());
Ok(())
}
#[test]
fn forbidden_word_inside_a_string_literal_is_not_flagged()
-> Result<(), Box<dyn std::error::Error>> {
let source = "pub fn run(input) {\n \
log(\"erlang.system_time is forbidden here\")\n \
workflow.run(wrappers.charge_activity(input))\n}\n";
assert!(analyze_determinism(source, "run")?.is_empty());
Ok(())
}
#[test]
fn missing_entry_function_is_a_loud_error() {
let source = "fn helper() {\n Nil\n}\n";
let result = analyze_determinism(source, "run");
assert_eq!(
result,
Err(DeterminismError::EntryFunctionNotFound {
function: "run".to_owned(),
})
);
}
#[test]
fn mutually_recursive_helpers_terminate() -> Result<(), Box<dyn std::error::Error>> {
let source = "pub fn run(input) {\n ping(input)\n}\n\
fn ping(input) {\n pong(input)\n}\n\
fn pong(input) {\n ping(input)\n os.system_time(1)\n}\n";
let violations = analyze_determinism(source, "run")?;
assert_eq!(violations.len(), 1);
assert_eq!(violations[0].call, "os.system_time");
Ok(())
}
#[test]
fn multiple_distinct_calls_are_all_reported() -> Result<(), Box<dyn std::error::Error>> {
let source = "pub fn run(input) {\n \
let a = os.system_time(1)\n \
let b = crypto.strong_rand_bytes(16)\n \
let c = erlang.unique_integer([])\n}\n";
let violations = analyze_determinism(source, "run")?;
let calls: Vec<&str> = violations.iter().map(|v| v.call.as_str()).collect();
assert_eq!(
calls,
vec![
"crypto.strong_rand_bytes",
"erlang.unique_integer",
"os.system_time",
]
);
Ok(())
}
}