1use alloc::{
2 format,
3 string::{String, ToString},
4 vec,
5 vec::Vec,
6};
7use std::{fs, path::Path};
8
9use miden_core::Word;
10use miden_precompiles::{
11 CurveId, CurvePrecompile, Limbs, ONE_LIMBS, TWO_LIMBS, UintDomain, UintPrecompile, ZERO_LIMBS,
12};
13
14const UINT_TEMPLATE_PATH: &str = "crates/lib/core/codegen/src/templates/uint.masm.tpl";
15const UINT_TEMPLATE: &str = include_str!("templates/uint.masm.tpl");
16const U256_CONSTANTS_TEMPLATE: &str = include_str!("templates/u256_constants.masm.tpl");
17const FIELD_CONSTANTS_TEMPLATE: &str = include_str!("templates/field_constants.masm.tpl");
18const FIELD_EXTRA_OPS_TEMPLATE: &str = include_str!("templates/field_extra_ops.masm.tpl");
19const CURVE_TEMPLATE_PATH: &str = "crates/lib/core/codegen/src/templates/curve.masm.tpl";
20const CURVE_TEMPLATE: &str = include_str!("templates/curve.masm.tpl");
21const REGENERATE_COMMAND: &str =
22 "cargo run -p miden-core-lib-codegen -- --out target/miden-core-lib-generated-asm";
23const ASM_PATH_PREFIX: &str = "asm/";
24
25fn generated_files() -> Result<Vec<GeneratedFile>, String> {
26 let mut generated = Vec::with_capacity(UintDomain::ALL.len() + CurveId::ALL.len());
27
28 for domain in UintDomain::ALL {
29 let config = UintMasmConfig::new(domain);
30 generated.push(GeneratedFile {
31 path: config.path,
32 contents: render_uint(&config)?,
33 });
34 }
35
36 for curve in CurveId::ALL {
37 let config = CurveMasmConfig::new(curve);
38 generated.push(GeneratedFile {
39 path: config.path,
40 contents: render_curve(&config)?,
41 });
42 }
43
44 Ok(generated)
45}
46
47fn render_uint(config: &UintMasmConfig) -> Result<String, String> {
48 let domain = config.domain;
49 let zero = constant(ZERO_LIMBS, domain);
50 let one = constant(ONE_LIMBS, domain);
51 let two = constant(TWO_LIMBS, domain);
52 let domain_constants = render_uint_constants(config)?;
53 let domain_extra_procs = render_uint_extra_procs(config)?;
54 let op_tag = |op_id| word_literal(tag_word(UintPrecompile::op_tag(op_id)));
55
56 let replacements = vec![
57 ("TEMPLATE_PATH", UINT_TEMPLATE_PATH.to_string()),
58 ("REGENERATE_COMMAND", REGENERATE_COMMAND.to_string()),
59 ("TITLE", config.title.to_string()),
60 ("DOMAIN_KIND", domain_kind(domain).to_string()),
61 ("VALUE_KIND", value_kind(domain).to_string()),
62 ("ENCODED_MODULUS_NOTE", encoded_modulus_note(domain).to_string()),
63 ("BOUND_PTR", domain.bound_ptr().to_string()),
64 ("ENCODED_MODULUS_LIMBS", limbs_literal(domain.encoded_modulus())),
65 ("PRECOMPILE_ID", UintPrecompile::id().as_canonical_u64().to_string()),
66 ("VALUE_TAG", word_literal(tag_word(UintPrecompile::value_tag(domain)))),
67 ("ADD_TAG", op_tag(UintPrecompile::ADD_OP_ID)),
68 ("SUB_TAG", op_tag(UintPrecompile::SUB_OP_ID)),
69 ("MUL_TAG", op_tag(UintPrecompile::MUL_OP_ID)),
70 ("EQ_TAG", op_tag(UintPrecompile::EQ_OP_ID)),
71 ("ZERO_DIGEST", zero.digest),
72 ("ZERO_LO_WORD", zero.lo_word),
73 ("ZERO_HI_WORD", zero.hi_word),
74 ("ONE_DIGEST", one.digest),
75 ("ONE_LO_WORD", one.lo_word),
76 ("ONE_HI_WORD", one.hi_word),
77 ("TWO_DIGEST", two.digest),
78 ("TWO_LO_WORD", two.lo_word),
79 ("TWO_HI_WORD", two.hi_word),
80 ("DOMAIN_CONSTANT_PROCS", domain_constants),
81 ("DOMAIN_EXTRA_PROCS", domain_extra_procs),
82 ];
83
84 render_template(UINT_TEMPLATE, &replacements)
85}
86
87fn render_uint_constants(config: &UintMasmConfig) -> Result<String, String> {
88 let domain = config.domain;
89 if domain == UintDomain::U256 {
90 let max = constant(domain.max().expect("u256 max is defined"), domain);
91 return render_template(
92 U256_CONSTANTS_TEMPLATE,
93 &[
94 ("MAX_DIGEST", max.digest),
95 ("MAX_LO_WORD", max.lo_word),
96 ("MAX_HI_WORD", max.hi_word),
97 ],
98 );
99 }
100
101 if !domain.is_prime_field() {
102 return Err(format!("{} must be marked as a prime-field domain", config.title));
103 }
104
105 let minus_one = constant(domain.minus_one(), domain);
106 let half = constant(domain.half().expect("field half is defined"), domain);
107 let pow_2_128 = constant(domain.pow2_mod(128).expect("field 2^128 is defined"), domain);
108 let pow_2_256 = constant(domain.pow2_mod(256).expect("field 2^256 is defined"), domain);
109 let pow_2_384 = constant(domain.pow2_mod(384).expect("field 2^384 is defined"), domain);
110 render_template(
111 FIELD_CONSTANTS_TEMPLATE,
112 &[
113 ("MINUS_ONE_DIGEST", minus_one.digest),
114 ("MINUS_ONE_LO_WORD", minus_one.lo_word),
115 ("MINUS_ONE_HI_WORD", minus_one.hi_word),
116 ("HALF_DIGEST", half.digest),
117 ("HALF_LO_WORD", half.lo_word),
118 ("HALF_HI_WORD", half.hi_word),
119 ("POW_2_128_DIGEST", pow_2_128.digest),
120 ("POW_2_128_LO_WORD", pow_2_128.lo_word),
121 ("POW_2_128_HI_WORD", pow_2_128.hi_word),
122 ("POW_2_256_DIGEST", pow_2_256.digest),
123 ("POW_2_256_LO_WORD", pow_2_256.lo_word),
124 ("POW_2_256_HI_WORD", pow_2_256.hi_word),
125 ("POW_2_384_DIGEST", pow_2_384.digest),
126 ("POW_2_384_LO_WORD", pow_2_384.lo_word),
127 ("POW_2_384_HI_WORD", pow_2_384.hi_word),
128 ],
129 )
130}
131
132fn render_uint_extra_procs(config: &UintMasmConfig) -> Result<String, String> {
133 if config.domain == UintDomain::U256 {
134 Ok(String::new())
135 } else {
136 render_template(FIELD_EXTRA_OPS_TEMPLATE, &[])
137 }
138}
139
140fn constant(value: Limbs, domain: UintDomain) -> ConstantMasm {
141 let digest = UintPrecompile::value_node(domain, value).digest();
142 let [lo, hi] = value_words(value);
143 ConstantMasm {
144 digest: word_literal(digest_word(digest)),
145 lo_word: limb_word_literal(lo),
146 hi_word: limb_word_literal(hi),
147 }
148}
149
150fn render_curve(config: &CurveMasmConfig) -> Result<String, String> {
151 let curve = config.curve;
152 let op_tag = |op_id| word_literal(tag_word(CurvePrecompile::op_tag(op_id)));
153 let replacements = vec![
154 ("TEMPLATE_PATH", CURVE_TEMPLATE_PATH.to_string()),
155 ("REGENERATE_COMMAND", REGENERATE_COMMAND.to_string()),
156 ("TITLE", config.title.to_string()),
157 ("BASE_FIELD_MODULE", config.base_field_module.to_string()),
158 ("BASE_FIELD_DESCRIPTION", config.base_field_description.to_string()),
159 ("PRECOMPILE_ID", CurvePrecompile::id().as_canonical_u64().to_string()),
160 ("GROUP_PTR", curve.group_ptr().to_string()),
161 ("VALUE_OP_ID", CurvePrecompile::VALUE_OP_ID.to_string()),
162 ("ADD_OP_ID", CurvePrecompile::ADD_OP_ID.to_string()),
163 ("SUB_OP_ID", CurvePrecompile::SUB_OP_ID.to_string()),
164 ("EQ_OP_ID", CurvePrecompile::EQ_OP_ID.to_string()),
165 ("MSM_OP_ID", CurvePrecompile::MSM_OP_ID.to_string()),
166 ("VALUE_TAG", word_literal(tag_word(CurvePrecompile::value_tag(curve)))),
167 ("ADD_TAG", op_tag(CurvePrecompile::ADD_OP_ID)),
168 ("SUB_TAG", op_tag(CurvePrecompile::SUB_OP_ID)),
169 ("EQ_TAG", op_tag(CurvePrecompile::EQ_OP_ID)),
170 ("MSM_TAG", word_literal(tag_word(CurvePrecompile::msm_tag()))),
171 (
172 "IDENTITY_DIGEST",
173 word_literal(digest_word(CurvePrecompile::identity_node(curve).digest())),
174 ),
175 (
176 "GENERATOR_DIGEST",
177 word_literal(digest_word(CurvePrecompile::generator_node(curve).digest())),
178 ),
179 ];
180
181 render_template(CURVE_TEMPLATE, &replacements)
182}
183
184fn render_template(template: &str, replacements: &[(&str, String)]) -> Result<String, String> {
185 let mut rendered = template.to_string();
186 apply_template_replacements(&mut rendered, replacements)?;
187 ensure_no_template_placeholders(&rendered)?;
188 Ok(rendered)
189}
190
191fn apply_template_replacements(
192 rendered: &mut String,
193 replacements: &[(&str, String)],
194) -> Result<(), String> {
195 for (name, value) in replacements {
196 let placeholder = format!("{{{{{name}}}}}");
197 if rendered.contains(&placeholder) {
198 *rendered = rendered.replace(&placeholder, value);
199 } else {
200 return Err(format!("template placeholder {placeholder} not found"));
201 }
202 }
203
204 Ok(())
205}
206
207pub fn write_math_masm(asm_dir: impl AsRef<Path>) -> Result<(), String> {
209 let asm_dir = asm_dir.as_ref();
210 for file in generated_files()? {
211 let relative_path = file.path.strip_prefix(ASM_PATH_PREFIX).ok_or_else(|| {
212 format!("generated path {} does not start with {ASM_PATH_PREFIX}", file.path)
213 })?;
214 write_file_if_changed(&asm_dir.join(relative_path), &file.contents)?;
215 }
216 Ok(())
217}
218
219pub fn render_u256_masm() -> Result<String, String> {
221 render_uint(&UintMasmConfig::new(UintDomain::U256))
222}
223
224pub fn write_to_dir(out_dir: impl AsRef<Path>) -> Result<(), String> {
226 let out_dir = out_dir.as_ref();
227 for file in generated_files()? {
228 write_file_if_changed(&out_dir.join(file.path), &file.contents)?;
229 }
230 Ok(())
231}
232
233fn ensure_no_template_placeholders(rendered: &str) -> Result<(), String> {
234 if let Some(start) = rendered.find("{{") {
235 let end = rendered[start..]
236 .find("}}")
237 .map(|offset| start + offset + 2)
238 .unwrap_or_else(|| (start + 40).min(rendered.len()));
239 return Err(format!("unreplaced template placeholder remains: {}", &rendered[start..end]));
240 }
241
242 if let Some(start) = rendered.find("}}") {
243 let end = (start + 40).min(rendered.len());
244 return Err(format!(
245 "unmatched template placeholder terminator remains: {}",
246 &rendered[start..end]
247 ));
248 }
249
250 Ok(())
251}
252
253fn tag_word(tag: miden_core::deferred::Tag) -> [u64; 4] {
254 let word = tag.as_word();
255 core::array::from_fn(|i| word[i].as_canonical_u64())
256}
257
258fn digest_word(digest: Word) -> [u64; 4] {
259 let elements = digest.as_elements();
260 core::array::from_fn(|i| elements[i].as_canonical_u64())
261}
262
263fn word_literal(word: [u64; 4]) -> String {
264 format!("[{}, {}, {}, {}]", word[0], word[1], word[2], word[3])
265}
266
267fn value_words(limbs: [u32; 8]) -> [[u32; 4]; 2] {
268 [
269 [limbs[0], limbs[1], limbs[2], limbs[3]],
270 [limbs[4], limbs[5], limbs[6], limbs[7]],
271 ]
272}
273
274fn limb_word_literal(word: [u32; 4]) -> String {
275 format!("[{}, {}, {}, {}]", word[0], word[1], word[2], word[3])
276}
277
278fn limbs_literal(limbs: [u32; 8]) -> String {
279 let limbs: Vec<String> = limbs.iter().map(|limb| format!("0x{limb:08x}")).collect();
280 format!("[{}]", limbs.join(", "))
281}
282
283fn write_file_if_changed(path: &Path, contents: &str) -> Result<(), String> {
284 use std::io::Write;
285
286 if let Some(parent) = path.parent() {
287 fs::create_dir_all(parent)
288 .map_err(|error| format!("failed to create {}: {error}", parent.display()))?;
289 }
290 match fs::read_to_string(path) {
291 Ok(actual) if actual == contents => Ok(()),
292 Ok(_) | Err(_) => {
293 let parent = path.parent().unwrap();
294 let name = path.file_stem().unwrap();
295 let mut tmpfile =
296 tempfile::NamedTempFile::with_prefix_in(name, parent).map_err(|error| {
297 format!("failed to create temporary file for {}: {error}", path.display())
298 })?;
299 tmpfile.write_all(contents.as_bytes()).map_err(|error| {
300 format!(
301 "failed to write contents to temporary file for {}: {error}",
302 path.display()
303 )
304 })?;
305 tmpfile
306 .persist(path)
307 .map_err(|error| format!("failed to persist {}: {error}", path.display()))?;
308
309 Ok(())
310 },
311 }
312}
313
314struct ConstantMasm {
315 digest: String,
316 lo_word: String,
317 hi_word: String,
318}
319
320struct GeneratedFile {
321 path: &'static str,
322 contents: String,
323}
324
325fn domain_kind(domain: UintDomain) -> &'static str {
326 if domain == UintDomain::U256 { "UINT" } else { "FIELD" }
327}
328
329fn value_kind(domain: UintDomain) -> &'static str {
330 if domain == UintDomain::U256 { "uint" } else { "field" }
331}
332
333fn encoded_modulus_note(domain: UintDomain) -> &'static str {
334 if domain == UintDomain::U256 {
335 ", all-zero means 2^256"
336 } else {
337 ""
338 }
339}
340
341#[derive(Clone, Copy)]
342struct UintMasmConfig {
343 path: &'static str,
344 title: &'static str,
345 domain: UintDomain,
346}
347
348impl UintMasmConfig {
349 const fn new(domain: UintDomain) -> Self {
350 match domain {
351 UintDomain::U256 => Self {
352 path: "asm/u256.masm",
353 title: "U256",
354 domain,
355 },
356 UintDomain::K1Base => Self {
357 path: "asm/fields/k1_base.masm",
358 title: "SECP256K1 BASE-FIELD",
359 domain,
360 },
361 UintDomain::K1Scalar => Self {
362 path: "asm/fields/k1_scalar.masm",
363 title: "SECP256K1 SCALAR-FIELD",
364 domain,
365 },
366 }
367 }
368}
369
370#[derive(Clone, Copy)]
371struct CurveMasmConfig {
372 path: &'static str,
373 title: &'static str,
374 base_field_module: &'static str,
375 base_field_description: &'static str,
376 curve: CurveId,
377}
378
379impl CurveMasmConfig {
380 const fn new(curve: CurveId) -> Self {
381 match curve {
382 CurveId::Secp256k1 => Self {
383 path: "asm/curves/secp256k1.masm",
384 title: "SECP256K1",
385 base_field_module: "k1_base",
386 base_field_description: "secp256k1 base-field",
387 curve,
388 },
389 }
390 }
391}
392
393#[cfg(test)]
394mod tests {
395 use std::{
396 env, fs,
397 path::PathBuf,
398 time::{SystemTime, UNIX_EPOCH},
399 };
400
401 use super::*;
402
403 #[test]
404 fn write_to_dir_preserves_unrelated_files() {
405 let out_dir = unique_temp_dir("miden-core-lib-codegen-write-to-dir");
406 let unrelated_file = out_dir.join("notes.txt");
407 fs::create_dir_all(&out_dir).unwrap();
408 fs::write(&unrelated_file, "keep me").unwrap();
409
410 write_to_dir(&out_dir).unwrap();
411
412 assert_eq!(fs::read_to_string(&unrelated_file).unwrap(), "keep me");
413 assert!(out_dir.join("asm/u256.masm").exists());
414 assert!(out_dir.join("asm/fields/k1_base.masm").exists());
415 assert!(out_dir.join("asm/curves/secp256k1.masm").exists());
416
417 fs::remove_dir_all(&out_dir).unwrap();
418 }
419
420 fn unique_temp_dir(prefix: &str) -> PathBuf {
421 let nanos = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_nanos();
422 env::temp_dir().join(format!("{prefix}-{}-{nanos}", std::process::id()))
423 }
424}