1use alloc::boxed::Box;
2use alloc::string::ToString;
3use alloc::sync::Arc;
4use alloc::vec::Vec;
5
6use miden_core::mast::{MastNodeExt, UntrustedMastForest};
7use miden_mast_package::Package;
8use miden_mast_package::debug_info::PackageDebugInfo;
9use miden_processor::LoadedMastForest;
10use thiserror::Error;
11
12use crate::assembly::Path;
13use crate::package::{loaded_mast_forest, package_debug_info};
14use crate::utils::create_external_node_forest;
15use crate::utils::serde::{
16 ByteReader,
17 ByteWriter,
18 Deserializable,
19 DeserializationError,
20 Serializable,
21};
22use crate::vm::AdviceMap;
23use crate::{MastForest, MastNodeId, Word};
24
25#[derive(Debug, Error)]
30pub enum MastForestScriptError {
31 #[error("entrypoint node {0} is not in the provided MAST forest")]
32 EntrypointNotInForest(MastNodeId),
33 #[error("package does not contain a procedure with '@{0}' attribute")]
34 NoProcedureWithAttribute(Box<str>),
35 #[error("package contains multiple procedures with '@{0}' attribute")]
36 MultipleProceduresWithAttribute(Box<str>),
37 #[error("procedure at path '{0}' not found in package")]
38 ProcedureNotFound(Box<str>),
39 #[error("procedure at path '{0}' does not have the specified attribute")]
40 ProcedureMissingAttribute(Box<str>),
41 #[error("expected a library package, but the provided package is an executable")]
42 ExecutablePackage,
43}
44
45#[derive(Debug, Clone)]
55pub(crate) struct MastForestScript {
56 mast: Arc<MastForest>,
57 entrypoint: MastNodeId,
58 package_debug_info: Option<Arc<PackageDebugInfo>>,
59}
60
61impl MastForestScript {
62 pub fn from_parts(
70 mast: Arc<MastForest>,
71 entrypoint: MastNodeId,
72 ) -> Result<Self, MastForestScriptError> {
73 if mast.get_node_by_id(entrypoint).is_none() {
74 return Err(MastForestScriptError::EntrypointNotInForest(entrypoint));
75 }
76 Ok(Self {
77 mast,
78 entrypoint,
79 package_debug_info: None,
80 })
81 }
82
83 pub(crate) fn from_package(
89 package: &Package,
90 attribute: &str,
91 ) -> Result<Self, MastForestScriptError> {
92 if package.is_program() {
93 return Err(MastForestScriptError::ExecutablePackage);
94 }
95
96 let mut entrypoint = None;
97
98 for export in package.manifest.exports() {
99 if let Some(proc_export) = export.as_procedure()
100 && proc_export.attributes.has(attribute)
101 {
102 if entrypoint.is_some() {
103 return Err(MastForestScriptError::MultipleProceduresWithAttribute(
104 attribute.into(),
105 ));
106 }
107 entrypoint = Some(proc_export.node.ok_or_else(|| {
108 MastForestScriptError::NoProcedureWithAttribute(attribute.into())
109 })?);
110 }
111 }
112
113 let entrypoint = entrypoint
114 .ok_or_else(|| MastForestScriptError::NoProcedureWithAttribute(attribute.into()))?;
115
116 Ok(Self::from_parts(package.mast_forest().clone(), entrypoint)?
117 .with_package_debug_info(package))
118 }
119
120 pub(crate) fn from_package_reference(
129 package: &Package,
130 path: &Path,
131 attribute: &str,
132 ) -> Result<Self, MastForestScriptError> {
133 let export = package
134 .manifest
135 .exports()
136 .find(|e| e.path().as_ref() == path)
137 .ok_or_else(|| MastForestScriptError::ProcedureNotFound(path.to_string().into()))?;
138
139 let proc_export = export
140 .as_procedure()
141 .ok_or_else(|| MastForestScriptError::ProcedureNotFound(path.to_string().into()))?;
142
143 if !proc_export.attributes.has(attribute) {
144 return Err(MastForestScriptError::ProcedureMissingAttribute(path.to_string().into()));
145 }
146
147 let digest = proc_export.digest;
148
149 let (mast, entrypoint) = create_external_node_forest(digest);
150
151 Ok(Self::from_parts(Arc::new(mast), entrypoint)?.with_package_debug_info(package))
152 }
153
154 pub fn mast(&self) -> Arc<MastForest> {
159 self.mast.clone()
160 }
161
162 pub fn loaded_mast_forest(&self) -> LoadedMastForest {
164 loaded_mast_forest(self.mast.clone(), self.package_debug_info.clone())
165 }
166
167 pub fn digest(&self) -> Word {
169 self.mast[self.entrypoint].digest()
170 }
171
172 pub fn entrypoint(&self) -> MastNodeId {
174 self.entrypoint
175 }
176
177 pub fn clear_debug_info(&mut self) {
179 self.package_debug_info = None;
180 }
181
182 pub fn with_package_debug_info(mut self, package: &Package) -> Self {
185 self.package_debug_info = package_debug_info(package);
186 self
187 }
188
189 pub fn with_advice_map(mut self, advice_map: AdviceMap) -> Self {
195 if advice_map.is_empty() {
196 return self;
197 }
198
199 let mast = (*self.mast).clone().with_advice_map(advice_map);
200 self.mast = Arc::new(mast);
201 self
202 }
203}
204
205impl PartialEq for MastForestScript {
206 fn eq(&self, other: &Self) -> bool {
207 self.mast == other.mast && self.entrypoint == other.entrypoint
208 }
209}
210
211impl Eq for MastForestScript {}
212
213impl Serializable for MastForestScript {
217 fn write_into<W: ByteWriter>(&self, target: &mut W) {
218 self.mast.write_hashless(target);
219 target.write_u32(u32::from(self.entrypoint));
220 }
221
222 fn get_size_hint(&self) -> usize {
223 let mut mast_target = Vec::new();
227 self.mast.write_hashless(&mut mast_target);
228 let mast_size = mast_target.len();
229 let u32_size = 0u32.get_size_hint();
230
231 mast_size + u32_size
232 }
233}
234
235impl Deserializable for MastForestScript {
236 fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
237 let mast = UntrustedMastForest::read_from_reader(source)?
239 .validate()
240 .map_err(|err| DeserializationError::InvalidValue(err.to_string()))?;
241 let entrypoint = MastNodeId::from_u32_safe(source.read_u32()?, &mast)?;
242
243 Self::from_parts(Arc::new(mast), entrypoint)
244 .map_err(|e| DeserializationError::InvalidValue(e.to_string()))
245 }
246}