use std::{
collections::BTreeMap,
fs::File,
io::{Cursor, Read, Seek},
path::Path,
};
use zip::{ZipArchive, result::ZipError};
use crate::{
AwlSource, BeamModule, BeamSet, ContentHash, ContractIdentityError, ExtractionLimits, Manifest,
PackageContract, PackageError,
awl::{AWL_DOCUMENT_PREFIX, AWL_SCHEMA_PREFIX},
builder::is_safe_logical_name,
extraction::ExtractionBudget,
hash::{has_contract_identity, verified_content_hash_with_contract},
namespace::deployed_name,
version::WorkflowVersion,
};
const MANIFEST_ENTRY: &str = "manifest.json";
const CONTRACT_ENTRY: &str = "contract.json";
const BEAM_PREFIX: &str = "beam/";
const BEAM_SUFFIX: &str = ".beam";
const SOURCE_PREFIX: &str = "src/";
const SOURCE_SUFFIX: &str = ".gleam";
const AWL_PREFIX: &str = "awl/";
#[derive(Clone, Debug, PartialEq)]
pub struct Package {
manifest: Manifest,
contract: Option<PackageContract>,
beams: BeamSet,
source: BTreeMap<String, Vec<u8>>,
awl: Option<AwlSource>,
content_hash: ContentHash,
}
struct ArchiveEntries {
beams: BeamSet,
source: BTreeMap<String, Vec<u8>>,
awl: Option<AwlSource>,
}
impl Package {
pub fn load_from_path(
path: impl AsRef<Path>,
limits: ExtractionLimits,
) -> Result<Self, PackageError> {
let file =
File::open(path).map_err(|source| PackageError::ArchiveRead(ZipError::Io(source)))?;
Self::load_from_reader(file, limits)
}
pub fn load_from_bytes(
bytes: impl AsRef<[u8]>,
limits: ExtractionLimits,
) -> Result<Self, PackageError> {
Self::load_from_reader(Cursor::new(bytes.as_ref()), limits)
}
fn load_from_reader<R>(reader: R, limits: ExtractionLimits) -> Result<Self, PackageError>
where
R: Read + Seek,
{
let mut archive = ZipArchive::new(reader).map_err(PackageError::ArchiveRead)?;
let mut budget = limits.budget();
let manifest = read_manifest(&mut archive, &mut budget)?;
manifest.check_format_version()?;
let contract = read_contract(&mut archive, &mut budget)?;
let entries = read_archive_entries(&mut archive, &mut budget)?;
let ArchiveEntries { beams, source, awl } = entries;
let content_hash =
verified_content_hash_with_contract(&beams, &manifest, contract.as_ref())?;
if beams.get(&manifest.entry_module).is_none() {
return Err(PackageError::MissingEntryModule {
module: manifest.entry_module.clone(),
});
}
Ok(Self {
manifest,
contract,
beams,
source,
awl,
content_hash,
})
}
#[must_use]
pub const fn manifest(&self) -> &Manifest {
&self.manifest
}
#[must_use]
pub const fn beams(&self) -> &BeamSet {
&self.beams
}
#[must_use]
pub const fn source(&self) -> &BTreeMap<String, Vec<u8>> {
&self.source
}
#[must_use]
pub const fn awl(&self) -> Option<&AwlSource> {
self.awl.as_ref()
}
#[must_use]
pub const fn content_hash(&self) -> &ContentHash {
&self.content_hash
}
pub fn contract(&self) -> Result<&PackageContract, ContractIdentityError> {
if has_contract_identity(
&self.beams,
&self.manifest,
self.contract.as_ref(),
&self.content_hash,
) {
self.contract
.as_ref()
.ok_or_else(|| ContractIdentityError::RedeployRequired {
stored_version: self.content_hash.to_string(),
})
} else {
Err(ContractIdentityError::RedeployRequired {
stored_version: self.content_hash.to_string(),
})
}
}
#[must_use]
pub fn has_declared_timeout(&self) -> bool {
self.manifest.timeout.is_some()
&& has_contract_identity(
&self.beams,
&self.manifest,
self.contract.as_ref(),
&self.content_hash,
)
}
#[must_use]
pub fn declared_timeout(&self) -> Option<std::time::Duration> {
self.declared_entry_timeout(self.manifest.timeout)
}
#[must_use]
pub fn declared_entry_timeout(
&self,
entry_timeout: Option<std::time::Duration>,
) -> Option<std::time::Duration> {
if self.has_declared_timeout() {
entry_timeout
} else {
None
}
}
#[must_use]
pub fn version_record(&self) -> WorkflowVersion {
WorkflowVersion {
entry_module: self.manifest.entry_module.clone(),
content_hash: self.content_hash.clone(),
activities: self.manifest.activities.clone(),
input_schema: self.manifest.input_schema.clone(),
output_schema: self.manifest.output_schema.clone(),
}
}
#[must_use]
pub fn deployed_modules(&self) -> Vec<(String, &[u8])> {
self.beams
.iter()
.map(|module| {
(
deployed_name(module.name(), &self.content_hash),
module.bytes(),
)
})
.collect()
}
#[must_use]
pub fn deployed_entry_module(&self) -> String {
deployed_name(&self.manifest.entry_module, &self.content_hash)
}
pub fn to_archive_bytes(&self) -> Result<Vec<u8>, PackageError> {
let mut builder = crate::PackageBuilder::with_source(
self.manifest.clone(),
self.beams.clone(),
self.source.clone(),
);
if let Some(awl) = self.awl.clone() {
builder = builder.with_awl_source(awl);
}
builder
.preserving_loaded_identity(self.content_hash.clone(), self.contract.clone())
.write_to_bytes()
}
#[cfg(any(test, feature = "test-support"))]
#[doc(hidden)]
#[must_use]
pub fn from_validated_parts_for_test(
manifest: Manifest,
beams: BeamSet,
source: BTreeMap<String, Vec<u8>>,
content_hash: ContentHash,
) -> Self {
Self {
manifest,
contract: None,
beams,
source,
awl: None,
content_hash,
}
}
}
fn read_manifest<R>(
archive: &mut ZipArchive<R>,
budget: &mut ExtractionBudget,
) -> Result<Manifest, PackageError>
where
R: Read + Seek,
{
let mut manifest_file = match archive.by_name(MANIFEST_ENTRY) {
Ok(file) => file,
Err(ZipError::FileNotFound) => return Err(PackageError::MissingManifest),
Err(error) => return Err(PackageError::ArchiveRead(error)),
};
let manifest_bytes = budget.read_entry(&mut manifest_file)?;
serde_json::from_slice(&manifest_bytes).map_err(|source| PackageError::ManifestParse { source })
}
fn read_contract<R>(
archive: &mut ZipArchive<R>,
budget: &mut ExtractionBudget,
) -> Result<Option<PackageContract>, PackageError>
where
R: Read + Seek,
{
let mut contract_file = match archive.by_name(CONTRACT_ENTRY) {
Ok(file) => file,
Err(ZipError::FileNotFound) => return Ok(None),
Err(error) => return Err(PackageError::ArchiveRead(error)),
};
let contract_bytes = budget.read_entry(&mut contract_file)?;
let contract = serde_json::from_slice(&contract_bytes)
.map_err(|source| PackageError::ContractParse { source })?;
Ok(Some(contract))
}
fn read_archive_entries<R>(
archive: &mut ZipArchive<R>,
budget: &mut ExtractionBudget,
) -> Result<ArchiveEntries, PackageError>
where
R: Read + Seek,
{
let mut modules = Vec::new();
let mut source = BTreeMap::new();
let mut document: Option<(String, String)> = None;
let mut schemas = BTreeMap::new();
for index in 0..archive.len() {
let mut file = archive.by_index(index).map_err(PackageError::ArchiveRead)?;
if file.is_dir() {
continue;
}
let entry = file.name().to_owned();
if entry == MANIFEST_ENTRY || entry == CONTRACT_ENTRY {
continue;
}
if entry.starts_with(BEAM_PREFIX) {
let logical = logical_name_from_entry(&entry, BEAM_PREFIX, BEAM_SUFFIX)?;
let bytes = budget.read_entry(&mut file)?;
modules.push(BeamModule::new(logical, bytes));
} else if entry.starts_with(SOURCE_PREFIX) {
let logical = logical_name_from_entry(&entry, SOURCE_PREFIX, SOURCE_SUFFIX)?;
let bytes = budget.read_entry(&mut file)?;
if source.insert(logical, bytes).is_some() {
return Err(PackageError::MalformedBeamEntry { entry });
}
} else if let Some(name) = entry.strip_prefix(AWL_DOCUMENT_PREFIX) {
if name.contains('/') {
return Err(PackageError::MalformedAwlEntry { entry });
}
let name = awl_relative_path(&entry, name)?;
let bytes = budget.read_entry(&mut file)?;
let text =
String::from_utf8(bytes).map_err(|source| PackageError::AwlDocumentNotUtf8 {
entry: entry.clone(),
source,
})?;
if document.replace((name, text)).is_some() {
return Err(PackageError::MalformedAwlEntry { entry });
}
} else if let Some(path) = entry.strip_prefix(AWL_SCHEMA_PREFIX) {
let path = awl_relative_path(&entry, path)?;
let bytes = budget.read_entry(&mut file)?;
if schemas.insert(path, bytes).is_some() {
return Err(PackageError::MalformedAwlEntry { entry });
}
} else if entry.starts_with(AWL_PREFIX) {
return Err(PackageError::MalformedAwlEntry { entry });
}
}
let awl = match document {
Some((name, text)) => Some(AwlSource::new(name, text, schemas)),
None if schemas.is_empty() => None,
None => return Err(PackageError::MissingAwlDocument),
};
let beams = BeamSet::new(modules)?;
Ok(ArchiveEntries { beams, source, awl })
}
fn awl_relative_path(entry: &str, relative_path: &str) -> Result<String, PackageError> {
if is_safe_logical_name(relative_path) {
Ok(relative_path.to_owned())
} else {
Err(PackageError::MalformedAwlEntry {
entry: entry.to_owned(),
})
}
}
fn logical_name_from_entry(
entry: &str,
prefix: &str,
suffix: &str,
) -> Result<String, PackageError> {
let Some(without_prefix) = entry.strip_prefix(prefix) else {
return Err(PackageError::MalformedBeamEntry {
entry: entry.to_owned(),
});
};
let Some(logical) = without_prefix.strip_suffix(suffix) else {
return Err(PackageError::MalformedBeamEntry {
entry: entry.to_owned(),
});
};
if is_safe_logical_name(logical) {
Ok(logical.to_owned())
} else {
Err(PackageError::MalformedBeamEntry {
entry: entry.to_owned(),
})
}
}
#[cfg(test)]
#[path = "package_tests.rs"]
mod tests;