1use aion_package::{ContentHash, Package};
4use chrono::Utc;
5
6use super::{CatalogEntry, CatalogSnapshot, PackageGroup, WorkflowCatalog};
7use crate::loader::load::{LoadOutcome, StagedLoad, load_error, rollback_registered};
8use crate::{error::EngineError, runtime::RuntimeHandle};
9
10impl WorkflowCatalog {
11 pub async fn load_package(
31 &self,
32 runtime: &RuntimeHandle,
33 package: &Package,
34 ) -> Result<LoadOutcome, EngineError> {
35 self.load_package_mode(runtime, package, true).await
36 }
37
38 pub(crate) async fn stage_package(
40 &self,
41 runtime: &RuntimeHandle,
42 package: &Package,
43 ) -> Result<LoadOutcome, EngineError> {
44 self.load_package_mode(runtime, package, false).await
45 }
46
47 async fn load_package_mode(
48 &self,
49 runtime: &RuntimeHandle,
50 package: &Package,
51 publish_routes: bool,
52 ) -> Result<LoadOutcome, EngineError> {
53 let hash = package.content_hash();
54 let nif_modules = runtime.registered_nif_modules();
55
56 let originals: Vec<&str> = package
57 .beams()
58 .iter()
59 .map(aion_package::BeamModule::name)
60 .filter(|name| !nif_modules.contains(&(*name).to_owned()))
61 .collect();
62 let deployed: Vec<String> = originals
63 .iter()
64 .map(|name| aion_package::deployed_name(name, hash))
65 .collect();
66 let deployed_refs: Vec<&str> = deployed.iter().map(String::as_str).collect();
67 let rename_map = runtime.package_rename_map(&originals, &deployed_refs);
68
69 let nif_set: std::collections::HashSet<&str> =
70 nif_modules.iter().map(String::as_str).collect();
71 let is_nif = |name: &str| {
72 let original = name.split('$').next().unwrap_or(name);
73 nif_set.contains(original)
74 };
75
76 self.load_package_with_mode(
77 package,
78 |name, bytes| {
79 if is_nif(name) {
80 return Ok(());
81 }
82 runtime.register_module_with_renames(name, bytes, &rename_map)
83 },
84 |name| {
85 if is_nif(name) {
86 return Ok(());
87 }
88 runtime.unregister_module(name)
89 },
90 |entry_module, entry_function| {
91 if runtime.module_exports_function(entry_module, entry_function) {
92 Ok(())
93 } else {
94 Err(load_error(format!(
95 "deployed entry module `{entry_module}` does not export entry function `{entry_function}`"
96 )))
97 }
98 },
99 publish_routes,
100 )
101 .await
102 }
103
104 #[cfg(test)]
106 pub(crate) async fn load_package_with<F, R, V>(
107 &self,
108 package: &Package,
109 register: F,
110 rollback: R,
111 verify_entry: V,
112 ) -> Result<LoadOutcome, EngineError>
113 where
114 F: FnMut(&str, &[u8]) -> Result<(), EngineError>,
115 R: FnMut(&str) -> Result<(), EngineError>,
116 V: FnMut(&str, &str) -> Result<(), EngineError>,
117 {
118 self.load_package_with_mode(package, register, rollback, verify_entry, true)
119 .await
120 }
121
122 #[cfg(test)]
124 pub(crate) async fn stage_package_with<F, R, V>(
125 &self,
126 package: &Package,
127 register: F,
128 rollback: R,
129 verify_entry: V,
130 ) -> Result<LoadOutcome, EngineError>
131 where
132 F: FnMut(&str, &[u8]) -> Result<(), EngineError>,
133 R: FnMut(&str) -> Result<(), EngineError>,
134 V: FnMut(&str, &str) -> Result<(), EngineError>,
135 {
136 self.load_package_with_mode(package, register, rollback, verify_entry, false)
137 .await
138 }
139
140 async fn load_package_with_mode<F, R, V>(
141 &self,
142 package: &Package,
143 mut register: F,
144 mut rollback: R,
145 verify_entry: V,
146 publish_routes: bool,
147 ) -> Result<LoadOutcome, EngineError>
148 where
149 F: FnMut(&str, &[u8]) -> Result<(), EngineError>,
150 R: FnMut(&str) -> Result<(), EngineError>,
151 V: FnMut(&str, &str) -> Result<(), EngineError>,
152 {
153 let mutation_guard = self.mutations.lock().await;
154 let staged = StagedLoad::new(package)?;
155 let snapshot = self.current()?;
156
157 validate_staged_module_names(&staged, &snapshot)?;
158 if let Some(outcome) = self.resolve_existing_load(&staged, &snapshot, publish_routes)? {
159 drop(mutation_guard);
160 return Ok(outcome);
161 }
162
163 let registered_now =
164 register_staged_modules(&staged, &snapshot, &mut register, &mut rollback)?;
165 verify_staged_entries(&staged, verify_entry, &mut rollback, ®istered_now)?;
166
167 let records = staged.records();
168 let Some(record) = records.first().cloned() else {
169 return Err(load_error("package staged no workflow entries".to_owned()));
170 };
171 let mut next = (*snapshot).clone();
172 for module in &staged.modules {
173 next.registered_modules
174 .entry(module.deployed_name.clone())
175 .or_insert_with(|| staged.version.clone());
176 }
177 let loaded_at = Utc::now();
178 for workflow in &records {
179 next.by_version.insert(
180 (workflow.workflow_type().to_owned(), staged.version.clone()),
181 CatalogEntry {
182 workflow: workflow.clone(),
183 manifest_version: staged.manifest_version.clone(),
184 manifest_digest: staged.manifest_digest.clone(),
185 loaded_at,
186 },
187 );
188 }
189 next.package_groups.insert(
190 (
191 package.manifest().entry_module.clone(),
192 staged.version.clone(),
193 ),
194 PackageGroup {
195 primary_workflow_type: package.manifest().entry_module.clone(),
196 workflow_types: records
197 .iter()
198 .map(|workflow| workflow.workflow_type().to_owned())
199 .collect(),
200 },
201 );
202 let route_changed = records
205 .iter()
206 .any(|workflow| snapshot.routed.get(workflow.workflow_type()) != Some(&staged.version));
207 if publish_routes {
208 for workflow in &records {
209 next.routed
210 .insert(workflow.workflow_type().to_owned(), staged.version.clone());
211 }
212 }
213 self.install(next)?;
214 drop(mutation_guard);
215 Ok(LoadOutcome {
216 record,
217 freshly_loaded: true,
218 route_changed,
219 })
220 }
221
222 fn resolve_existing_load(
223 &self,
224 staged: &StagedLoad<'_>,
225 snapshot: &CatalogSnapshot,
226 publish_routes: bool,
227 ) -> Result<Option<LoadOutcome>, EngineError> {
228 let existing: Vec<_> = staged
229 .workflows
230 .iter()
231 .filter_map(|workflow| {
232 snapshot
233 .by_version
234 .get(&(workflow.workflow_type.clone(), staged.version.clone()))
235 })
236 .collect();
237 if existing.is_empty() {
238 return Ok(None);
239 }
240 if existing.len() != staged.workflows.len() {
241 return Err(load_error(format!(
242 "package version `{}` is only partially registered ({}/{} workflow entries)",
243 staged.version,
244 existing.len(),
245 staged.workflows.len()
246 )));
247 }
248
249 let first = existing[0];
250 if existing
253 .iter()
254 .any(|entry| entry.manifest_digest != staged.manifest_digest)
255 {
256 return Err(EngineError::ManifestMismatch {
257 workflow_type: first.workflow.workflow_type().to_owned(),
258 version: staged.version.clone(),
259 resident_digest: first.manifest_digest.to_string(),
260 incoming_digest: staged.manifest_digest.to_string(),
261 });
262 }
263
264 let route_changed = staged
265 .workflows
266 .iter()
267 .any(|workflow| snapshot.routed.get(&workflow.workflow_type) != Some(&staged.version));
268 if publish_routes && route_changed {
269 let mut next = snapshot.clone();
270 for workflow in &staged.workflows {
271 next.routed
272 .insert(workflow.workflow_type.clone(), staged.version.clone());
273 }
274 self.install(next)?;
275 }
276 Ok(Some(LoadOutcome {
277 record: first.workflow.clone(),
278 freshly_loaded: false,
279 route_changed,
280 }))
281 }
282
283 pub(crate) async fn publish_package_routes(
285 &self,
286 primary_workflow_type: &str,
287 version: &ContentHash,
288 ) -> Result<(), EngineError> {
289 let mutation_guard = self.mutations.lock().await;
290 let snapshot = self.current()?;
291 let group = snapshot
292 .package_groups
293 .get(&(primary_workflow_type.to_owned(), version.clone()))
294 .ok_or_else(|| load_error(format!("staged package group `{version}` is not loaded")))?;
295 for workflow_type in &group.workflow_types {
296 if !snapshot
297 .by_version
298 .contains_key(&(workflow_type.clone(), version.clone()))
299 {
300 return Err(load_error(format!(
301 "staged package group `{version}` is missing workflow `{workflow_type}`"
302 )));
303 }
304 }
305 let mut next = (*snapshot).clone();
306 for workflow_type in &group.workflow_types {
307 next.routed.insert(workflow_type.clone(), version.clone());
308 }
309 self.install(next)?;
310 drop(mutation_guard);
311 Ok(())
312 }
313}
314
315fn validate_staged_module_names(
316 staged: &StagedLoad<'_>,
317 snapshot: &CatalogSnapshot,
318) -> Result<(), EngineError> {
319 for module in &staged.modules {
320 if let Some(existing) = snapshot.registered_modules.get(&module.deployed_name) {
321 if existing != &staged.version {
322 return Err(load_error(format!(
323 "deployed module `{}` is already registered for content hash `{existing}`, not `{}`",
324 module.deployed_name, staged.version
325 )));
326 }
327 }
328 }
329 Ok(())
330}
331
332fn register_staged_modules<F, R>(
333 staged: &StagedLoad<'_>,
334 snapshot: &CatalogSnapshot,
335 register: &mut F,
336 rollback: &mut R,
337) -> Result<Vec<String>, EngineError>
338where
339 F: FnMut(&str, &[u8]) -> Result<(), EngineError>,
340 R: FnMut(&str) -> Result<(), EngineError>,
341{
342 let mut registered = Vec::new();
343 for module in &staged.modules {
344 if snapshot
345 .registered_modules
346 .contains_key(&module.deployed_name)
347 {
348 continue;
349 }
350 if let Err(error) = register(&module.deployed_name, module.bytes) {
351 let rollback_errors = rollback_registered(rollback, ®istered);
352 return Err(load_error(format!(
353 "runtime rejected deployed module `{}` after {} staged registrations: {error}{rollback_errors}",
354 module.deployed_name,
355 registered.len()
356 )));
357 }
358 registered.push(module.deployed_name.clone());
359 }
360 Ok(registered)
361}
362
363fn verify_staged_entries<V, R>(
364 staged: &StagedLoad<'_>,
365 mut verify: V,
366 rollback: &mut R,
367 registered: &[String],
368) -> Result<(), EngineError>
369where
370 V: FnMut(&str, &str) -> Result<(), EngineError>,
371 R: FnMut(&str) -> Result<(), EngineError>,
372{
373 for workflow in &staged.workflows {
374 if let Err(error) = verify(&workflow.deployed_entry_module, &workflow.entry_function) {
375 let rollback_errors = rollback_registered(rollback, registered);
376 return Err(load_error(format!(
377 "entry verification failed for workflow `{}`: {error}{rollback_errors}",
378 workflow.workflow_type
379 )));
380 }
381 }
382 Ok(())
383}