1use anyhow::{Result, ensure};
2use cairo_lang_compiler::diagnostics::DiagnosticsReporter;
3use cairo_lang_compiler::{ensure_diagnostics, get_sierra_program_for_functions};
4use cairo_lang_debug::DebugWithDb;
5use cairo_lang_defs::db::DefsGroup;
6use cairo_lang_defs::ids::{FreeFunctionId, FunctionWithBodyId, ModuleItemId};
7use cairo_lang_filesystem::db::FilesGroup;
8use cairo_lang_filesystem::ids::{CrateId, CrateInput};
9use cairo_lang_lowering::ids::ConcreteFunctionWithBodyId;
10use cairo_lang_semantic::items::function_with_body::FunctionWithBodySemantic;
11use cairo_lang_semantic::items::functions::GenericFunctionId;
12use cairo_lang_semantic::plugin::PluginSuite;
13use cairo_lang_semantic::{ConcreteFunction, FunctionLongId};
14use cairo_lang_sierra::debug_info::{Annotations, DebugInfo};
15use cairo_lang_sierra::extensions::gas::{CostTokenMap, CostTokenType};
16use cairo_lang_sierra::ids::FunctionId;
17use cairo_lang_sierra::program::ProgramArtifact;
18use cairo_lang_sierra_generator::db::SierraGenGroup;
19use cairo_lang_sierra_generator::debug_info::{
20 SerializableTypeNamesDebugInfo, StatementsLocations,
21};
22use cairo_lang_sierra_generator::executables::{collect_executables, find_executable_function_ids};
23use cairo_lang_sierra_generator::program_generator::SierraProgramWithDebug;
24use cairo_lang_sierra_generator::replace_ids::{DebugReplacer, SierraIdReplacer};
25use cairo_lang_starknet::contract::{
26 ContractDeclaration, ContractInfo, find_contracts, get_contract_abi_functions,
27 get_contracts_info,
28};
29use cairo_lang_starknet::plugin::consts::{CONSTRUCTOR_MODULE, EXTERNAL_MODULE, L1_HANDLER_MODULE};
30use cairo_lang_starknet_classes::casm_contract_class::ENTRY_POINT_COST;
31use cairo_lang_utils::CloneableDatabase;
32use cairo_lang_utils::ordered_hash_map::{
33 OrderedHashMap, deserialize_ordered_hashmap_vec, serialize_ordered_hashmap_vec,
34};
35use itertools::{Itertools, chain};
36pub use plugin::TestPlugin;
37use salsa::Database;
38use serde::{Deserialize, Serialize};
39use starknet_types_core::felt::Felt as Felt252;
40pub use test_config::{TestConfig, try_extract_test_config};
41
42mod inline_macros;
43pub mod plugin;
44pub mod test_config;
45
46const TEST_ATTR: &str = "test";
47const SHOULD_PANIC_ATTR: &str = "should_panic";
48const IGNORE_ATTR: &str = "ignore";
49const AVAILABLE_GAS_ATTR: &str = "available_gas";
50const STATIC_GAS_ARG: &str = "static";
51
52#[derive(Clone)]
54pub struct TestsCompilationConfig<'db> {
55 pub starknet: bool,
57
58 pub contract_declarations: Option<Vec<ContractDeclaration<'db>>>,
62
63 pub contract_crate_ids: Option<&'db [CrateId<'db>]>,
66
67 pub executable_crate_ids: Option<Vec<CrateId<'db>>>,
70
71 pub add_statements_functions: bool,
74
75 pub add_statements_code_locations: bool,
78
79 pub add_functions_debug_info: bool,
82
83 pub add_type_names: bool,
85
86 pub replace_ids: bool,
88}
89
90pub fn compile_test_prepared_db<'db>(
102 db: &'db dyn CloneableDatabase,
103 tests_compilation_config: TestsCompilationConfig<'db>,
104 test_crate_ids: Vec<CrateInput>,
105 mut diagnostics_reporter: DiagnosticsReporter<'_>,
106) -> Result<TestCompilation<'db>> {
107 ensure!(
108 tests_compilation_config.starknet
109 || tests_compilation_config.contract_declarations.is_none(),
110 "Contract declarations can be provided only when starknet is enabled."
111 );
112 ensure!(
113 tests_compilation_config.starknet || tests_compilation_config.contract_crate_ids.is_none(),
114 "Contract crate ids can be provided only when starknet is enabled."
115 );
116
117 ensure_diagnostics(db, &mut diagnostics_reporter)?;
118
119 let contracts = tests_compilation_config.contract_declarations.unwrap_or_else(|| {
120 find_contracts(
121 db,
122 tests_compilation_config.contract_crate_ids.unwrap_or_else(|| db.crates()),
123 )
124 });
125 let all_entry_points = if tests_compilation_config.starknet {
126 contracts
127 .iter()
128 .flat_map(|contract| {
129 chain!(
130 get_contract_abi_functions(db, contract, EXTERNAL_MODULE).unwrap_or_default(),
131 get_contract_abi_functions(db, contract, CONSTRUCTOR_MODULE)
132 .unwrap_or_default(),
133 get_contract_abi_functions(db, contract, L1_HANDLER_MODULE).unwrap_or_default(),
134 )
135 })
136 .map(|func| ConcreteFunctionWithBodyId::from_semantic(db, func.value))
137 .collect()
138 } else {
139 vec![]
140 };
141
142 let test_crate_ids = CrateInput::into_crate_ids(db, test_crate_ids);
143 let executable_functions = find_executable_function_ids(
144 db,
145 tests_compilation_config.executable_crate_ids.unwrap_or_else(|| test_crate_ids.clone()),
146 );
147 let all_tests = find_all_tests(db, test_crate_ids);
148
149 let func_ids = chain!(
150 executable_functions.keys().cloned(),
151 all_entry_points.iter().cloned(),
152 all_tests.iter().flat_map(|(func_id, _cfg)| {
154 ConcreteFunctionWithBodyId::from_no_generics_free(db, *func_id)
155 })
156 )
157 .collect();
158
159 let SierraProgramWithDebug { program: sierra_program, debug_info } =
160 get_sierra_program_for_functions(db, func_ids)?;
161
162 let function_set_costs: OrderedHashMap<FunctionId, CostTokenMap<i32>> = all_entry_points
163 .iter()
164 .map(|func_id| {
165 (
166 db.function_with_body_sierra(*func_id).unwrap().id.clone(),
167 CostTokenMap::from_iter([(CostTokenType::Const, ENTRY_POINT_COST)]),
168 )
169 })
170 .collect();
171
172 let replacer = DebugReplacer { db };
173
174 let sierra_program = if tests_compilation_config.replace_ids {
175 replacer.apply(sierra_program)
176 } else {
177 let mut sierra_program = sierra_program.clone();
178 replacer.enrich_function_names(&mut sierra_program);
179 sierra_program
180 };
181
182 let mut annotations = Annotations::default();
183 if tests_compilation_config.add_statements_functions {
184 annotations.extend(Annotations::from(
185 debug_info.statements_locations.extract_statements_functions(db),
186 ))
187 }
188 if tests_compilation_config.add_statements_code_locations {
189 annotations.extend(Annotations::from(
190 debug_info.statements_locations.extract_statements_source_code_locations(db),
191 ))
192 }
193
194 if tests_compilation_config.add_functions_debug_info {
195 annotations.extend(Annotations::from(
196 debug_info.functions_info.extract_serializable_debug_info(db),
197 ))
198 }
199
200 if tests_compilation_config.add_type_names {
201 annotations.extend(Annotations::from(SerializableTypeNamesDebugInfo::extract_type_names(
202 db,
203 &sierra_program,
204 )))
205 }
206
207 let executables = collect_executables(db, executable_functions, &sierra_program);
208 let named_tests = all_tests
209 .into_iter()
210 .map(|(func_id, test)| {
211 (
212 format!(
213 "{:?}",
214 FunctionLongId {
215 function: ConcreteFunction {
216 generic_function: GenericFunctionId::Free(func_id),
217 generic_args: vec![]
218 }
219 }
220 .debug(db)
221 ),
222 test,
223 )
224 })
225 .collect_vec();
226 let contracts_info = get_contracts_info(db, contracts, &replacer)?;
227 let sierra_program = ProgramArtifact::stripped(sierra_program).with_debug_info(DebugInfo {
228 executables,
229 annotations,
230 ..DebugInfo::default()
231 });
232
233 Ok(TestCompilation {
234 sierra_program,
235 metadata: TestCompilationMetadata {
236 named_tests,
237 function_set_costs,
238 contracts_info,
239 statements_locations: Some(debug_info.statements_locations.clone()),
240 },
241 })
242}
243
244#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
250pub struct TestCompilation<'db> {
251 pub sierra_program: ProgramArtifact,
252 #[serde(flatten)]
253 pub metadata: TestCompilationMetadata<'db>,
254}
255
256#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
261pub struct TestCompilationMetadata<'db> {
262 #[serde(
263 serialize_with = "serialize_ordered_hashmap_vec",
264 deserialize_with = "deserialize_ordered_hashmap_vec"
265 )]
266 pub contracts_info: OrderedHashMap<Felt252, ContractInfo>,
267 #[serde(
268 serialize_with = "serialize_ordered_hashmap_vec",
269 deserialize_with = "deserialize_ordered_hashmap_vec"
270 )]
271 pub function_set_costs: OrderedHashMap<FunctionId, CostTokenMap<i32>>,
272 pub named_tests: Vec<(String, TestConfig)>,
273 #[serde(skip)]
277 pub statements_locations: Option<StatementsLocations<'db>>,
278}
279
280fn find_all_tests<'db>(
282 db: &'db dyn Database,
283 main_crates: Vec<CrateId<'db>>,
284) -> Vec<(FreeFunctionId<'db>, TestConfig)> {
285 let mut tests = vec![];
286 for crate_id in main_crates {
287 let modules = db.crate_modules(crate_id);
288 for module_id in modules.iter() {
289 let Ok(module_data) = module_id.module_data(db) else {
290 continue;
291 };
292 tests.extend(module_data.items(db).iter().filter_map(|item| {
293 let ModuleItemId::FreeFunction(func_id) = item else { return None };
294 let Ok(attrs) =
295 db.function_with_body_attributes(FunctionWithBodyId::Free(*func_id))
296 else {
297 return None;
298 };
299 Some((*func_id, try_extract_test_config(db, attrs).ok()??))
300 }));
301 }
302 }
303 tests
304}
305
306pub fn test_assert_suite() -> PluginSuite {
308 let mut suite = PluginSuite::default();
309 suite
310 .add_inline_macro_plugin::<inline_macros::assert::AssertEqMacro>()
311 .add_inline_macro_plugin::<inline_macros::assert::AssertNeMacro>()
312 .add_inline_macro_plugin::<inline_macros::assert::AssertLtMacro>()
313 .add_inline_macro_plugin::<inline_macros::assert::AssertLeMacro>()
314 .add_inline_macro_plugin::<inline_macros::assert::AssertGtMacro>()
315 .add_inline_macro_plugin::<inline_macros::assert::AssertGeMacro>();
316 suite
317}
318
319pub fn test_plugin_suite() -> PluginSuite {
321 let mut suite = PluginSuite::default();
322 suite.add_plugin::<TestPlugin>().add(test_assert_suite());
323 suite
324}