use edb_common::types::EDB_STATE_VAR_FLAG;
use foundry_compilers::artifacts::{ContractKind, Mutability, StorageLocation, TypeName};
use semver::Version;
use crate::{
analysis::{VariableRef, USID, UVID},
contains_function_type, contains_mapping_type, contains_user_defined_type, VersionRef,
MAGIC_SNAPSHOT_NUMBER, MAGIC_VARIABLE_UPDATE_NUMBER,
};
pub fn generate_step_hook(version: &VersionRef, usid: USID) -> Option<String> {
if **version < Version::parse("0.5.0").unwrap() {
Some(format!(
"require(keccak256(uint256({}), uint256({})) != bytes32(uint256(0x2333)));",
MAGIC_SNAPSHOT_NUMBER,
u64::from(usid)
))
} else {
Some(format!(
"require(keccak256(abi.encode(uint256({}), uint256({}))) != bytes32(uint256(0x2333)));",
MAGIC_SNAPSHOT_NUMBER,
u64::from(usid)
))
}
}
pub fn generate_variable_update_hook(
version: &VersionRef,
uvid: UVID,
variable: &VariableRef,
) -> Option<String> {
if **version < Version::parse("0.4.24").unwrap() {
return None;
}
let declaration = variable.declaration();
let base_type = &declaration.type_name;
let is_state_variable = declaration.state_variable;
let is_calldata_variable = declaration.storage_location == StorageLocation::Calldata;
let is_storage_variable = declaration.storage_location == StorageLocation::Storage;
if base_type.as_ref().is_some_and(|ty| {
(contains_user_defined_type(ty) && **version < Version::parse("0.8.0").unwrap())
|| contains_function_type(ty)
|| contains_mapping_type(ty)
|| is_state_variable
|| is_calldata_variable
|| is_storage_variable
}) {
return None;
}
let base_var = variable.base();
let base_name = &base_var.declaration().name;
Some(format!(
"require(keccak256(abi.encode(uint256({}), uint256({}), abi.encode({}))) != bytes32(uint256(0x2333)));",
MAGIC_VARIABLE_UPDATE_NUMBER,
u64::from(uvid),
base_name
))
}
pub fn generate_view_method(state_variable: &VariableRef) -> Option<String> {
if state_variable
.read()
.contract()
.map(|c| c.definition().kind != ContractKind::Contract)
.unwrap_or(true)
{
return None;
}
let declaration = state_variable.declaration();
if declaration.mutability == Some(Mutability::Constant) {
return None;
}
let var_name = &declaration.name;
let type_name = declaration.type_name.as_ref()?;
let (params, return_type) = analyze_type_for_view_method(type_name)?;
let params_str = if params.is_empty() { String::new() } else { params.join(", ") };
let body = generate_view_body(var_name, ¶ms);
let uvid = state_variable.id();
Some(format!(
" function {var_name}{EDB_STATE_VAR_FLAG}{uvid}({params_str}) public view returns ({return_type}) {{\n return {body};\n }}"
))
}
fn analyze_type_for_view_method(type_name: &TypeName) -> Option<(Vec<String>, String)> {
analyze_type_recursive(type_name, 0)
}
fn analyze_type_recursive(type_name: &TypeName, depth: usize) -> Option<(Vec<String>, String)> {
match type_name {
TypeName::ElementaryTypeName(elementary) => {
let return_type = format_return_type(&elementary.name);
Some((Vec::new(), return_type))
}
TypeName::Mapping(mapping) => {
let key_type = &mapping.key_type;
let value_type = &mapping.value_type;
let key_type_str = match key_type {
TypeName::ElementaryTypeName(elem) => elem.name.clone(),
_ => return None, };
let (mut sub_params, return_type) = analyze_type_recursive(value_type, depth + 1)?;
let param_name = if depth == 0 { "key".to_string() } else { format!("key{depth}") };
let key_param = format!("{key_type_str} {param_name}");
let mut params = vec![key_param];
params.append(&mut sub_params);
Some((params, return_type))
}
TypeName::ArrayTypeName(array) => {
let base_type = &array.base_type;
let (mut sub_params, return_type) = analyze_type_recursive(base_type, depth + 1)?;
let param_name = if depth == 0 { "index".to_string() } else { format!("index{depth}") };
let index_param = format!("uint256 {param_name}");
let mut params = vec![index_param];
params.append(&mut sub_params);
Some((params, return_type))
}
TypeName::UserDefinedTypeName(_) => {
None
}
TypeName::FunctionTypeName(_) => {
None
}
}
}
fn generate_view_body(var_name: &str, params: &[String]) -> String {
if params.is_empty() {
var_name.to_string()
} else {
let param_names: Vec<String> = params
.iter()
.map(|p| {
p.split_whitespace().last().unwrap_or("").to_string()
})
.collect();
let mut body = var_name.to_string();
for param_name in param_names {
body = format!("{body}[{param_name}]");
}
body
}
}
fn format_return_type(type_name: &str) -> String {
match type_name {
t if t.starts_with("string") => format!("{t} memory"),
t if t.ends_with("[]") => format!("{t} memory"),
t if t.contains('[') && t.contains(']') => format!("{t} memory"),
_ => type_name.to_string(),
}
}
#[cfg(test)]
mod tests {
use crate::analysis;
#[test]
fn test_generate_view_method_primitive_types() {
let source = r#"
contract C {
uint256 private myValue;
address private owner;
bool private isActive;
}
"#;
let (_sources, analysis) = analysis::tests::compile_and_analyze(source);
assert_eq!(analysis.private_state_variables.len(), 3);
for private_var in &analysis.private_state_variables {
let result = super::generate_view_method(private_var);
assert!(
result.is_some(),
"Should generate view method for primitive type: {}",
private_var.declaration().name
);
let code = result.unwrap();
let var_name = &private_var.declaration().name;
assert!(code.contains(&format!("function {var_name}_edb_state_var_")));
assert!(code.contains("public view returns"));
assert!(code.contains(&format!("return {var_name};")));
}
}
#[test]
fn test_generate_view_method_mapping_types() {
let source = r#"
contract C {
mapping(address => uint256) private balances;
mapping(uint256 => bool) private permissions;
}
"#;
let (_sources, analysis) = analysis::tests::compile_and_analyze(source);
assert_eq!(analysis.private_state_variables.len(), 2);
for private_var in &analysis.private_state_variables {
let result = super::generate_view_method(private_var);
assert!(
result.is_some(),
"Should generate view method for mapping: {}",
private_var.declaration().name
);
let code = result.unwrap();
let var_name = &private_var.declaration().name;
assert!(code.contains(&format!("function {var_name}_edb_state_var_")));
assert!(code.contains("key) public view returns"));
assert!(code.contains(&format!("return {var_name}[key];")));
}
}
#[test]
fn test_generate_view_method_array_types() {
let source = r#"
contract C {
uint256[] private numbers;
address[] private addresses;
}
"#;
let (_sources, analysis) = analysis::tests::compile_and_analyze(source);
assert_eq!(analysis.private_state_variables.len(), 2);
for private_var in &analysis.private_state_variables {
let result = super::generate_view_method(private_var);
assert!(
result.is_some(),
"Should generate view method for array: {}",
private_var.declaration().name
);
let code = result.unwrap();
let var_name = &private_var.declaration().name;
assert!(code.contains(&format!("function {var_name}_edb_state_var_")));
assert!(code.contains("(uint256 index)"));
assert!(code.contains("public view returns"));
assert!(code.contains(&format!("return {var_name}[index];")));
}
}
#[test]
fn test_generate_view_method_nested_types() {
let source = r#"
contract C {
mapping(address => uint256[]) private userTokens;
mapping(uint256 => mapping(address => bool)) private permissions;
}
"#;
let (_sources, analysis) = analysis::tests::compile_and_analyze(source);
assert_eq!(analysis.private_state_variables.len(), 2);
let user_tokens =
analysis.private_state_variables.iter().find(|v| v.declaration().name == "userTokens");
let permissions =
analysis.private_state_variables.iter().find(|v| v.declaration().name == "permissions");
if let Some(user_tokens_var) = user_tokens {
let result = super::generate_view_method(user_tokens_var);
assert!(result.is_some(), "Should generate view method for nested mapping->array");
let code = result.unwrap();
assert!(code.contains("function userTokens_edb_state_var_"));
assert!(code.contains("(address key, uint256 index1)"));
assert!(code.contains("return userTokens[key][index1];"));
}
if let Some(permissions_var) = permissions {
let result = super::generate_view_method(permissions_var);
assert!(result.is_some(), "Should generate view method for nested mapping->mapping");
let code = result.unwrap();
assert!(code.contains("function permissions_edb_state_var_"));
assert!(code.contains("(uint256 key, address key1)"));
assert!(code.contains("return permissions[key][key1];"));
}
}
#[test]
fn test_generate_view_method_user_defined_types() {
let source = r#"
contract C {
struct User {
uint256 balance;
address addr;
}
User private userData;
User[] private users;
}
"#;
let (_sources, analysis) = analysis::tests::compile_and_analyze(source);
for private_var in &analysis.private_state_variables {
let result = super::generate_view_method(private_var);
assert!(
result.is_none(),
"Should not generate view method for user-defined type: {}",
private_var.declaration().name
);
}
}
#[test]
fn test_generate_view_method_reference_types() {
let source = r#"
contract C {
string private message;
string[] private messages;
}
"#;
let (_sources, analysis) = analysis::tests::compile_and_analyze(source);
assert_eq!(analysis.private_state_variables.len(), 2);
for private_var in &analysis.private_state_variables {
let result = super::generate_view_method(private_var);
assert!(
result.is_some(),
"Should generate view method for reference type: {}",
private_var.declaration().name
);
let code = result.unwrap();
let var_name = &private_var.declaration().name;
match var_name.as_str() {
"message" => {
assert!(code.contains("function message_edb_state_var_"));
assert!(code.contains("() public view returns (string memory)"));
assert!(code.contains("return message;"));
}
"messages" => {
assert!(code.contains("function messages_edb_state_var_"));
assert!(code.contains("(uint256 index) public view returns (string memory)"));
assert!(code.contains("return messages[index];"));
}
_ => panic!("Unexpected variable name: {var_name}"),
}
}
}
}