Skip to main content

miden_core_lib_codegen/
masm.rs

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
207/// Writes generated MASM files into an assembled MASM project root.
208pub 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
219/// Renders the generated U256 MASM module source.
220pub fn render_u256_masm() -> Result<String, String> {
221    render_uint(&UintMasmConfig::new(UintDomain::U256))
222}
223
224/// Writes generated MASM files into a developer preview directory.
225pub 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}