use crate::model::artifacts::{ChecksumError, canonical_hash};
use serde::{Deserialize, Deserializer, Serialize, de};
use std::{
collections::BTreeSet,
fmt,
path::{Path, PathBuf},
};
use thiserror::Error;
pub const MAX_RESTORE_REFERENCES: usize = 1024;
pub const MAX_JOURNAL_PATH_BYTES: usize = 4096;
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(try_from = "ReferenceFields")]
pub struct RestoreReferenceRecord {
journal: PathBuf,
authority: String,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct ReferenceFields {
journal: PathBuf,
authority: String,
}
impl TryFrom<ReferenceFields> for RestoreReferenceRecord {
type Error = RestoreReferenceError;
fn try_from(fields: ReferenceFields) -> Result<Self, Self::Error> {
Self::new(fields.journal, &fields.authority)
}
}
impl RestoreReferenceRecord {
pub fn new(journal: PathBuf, authority: &str) -> Result<Self, RestoreReferenceError> {
if !super::journal_path::is_canonical(&journal, MAX_JOURNAL_PATH_BYTES) {
return Err(RestoreReferenceError::InvalidJournal { journal });
}
let authority = canonical_hash(authority)?;
Ok(Self { journal, authority })
}
#[must_use]
pub fn journal(&self) -> &Path {
&self.journal
}
#[must_use]
pub fn authority(&self) -> &str {
&self.authority
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(try_from = "ReferencesFields")]
pub struct RestoreReferencesRecord {
version: u16,
restores: Vec<RestoreReferenceRecord>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct ReferencesFields {
version: u16,
#[serde(deserialize_with = "read_bounded_references")]
restores: Vec<RestoreReferenceRecord>,
}
impl TryFrom<ReferencesFields> for RestoreReferencesRecord {
type Error = RestoreReferenceError;
fn try_from(mut fields: ReferencesFields) -> Result<Self, Self::Error> {
if fields.version != 1 {
return Err(RestoreReferenceError::UnsupportedVersion(fields.version));
}
let mut journals = BTreeSet::new();
for entry in &fields.restores {
if !journals.insert(entry.journal()) {
return Err(RestoreReferenceError::DuplicateJournal {
journal: entry.journal.clone(),
});
}
}
fields
.restores
.sort_by(|left, right| left.journal.cmp(&right.journal));
Ok(Self {
version: 1,
restores: fields.restores,
})
}
}
impl RestoreReferencesRecord {
#[must_use]
pub const fn empty() -> Self {
Self {
version: 1,
restores: Vec::new(),
}
}
#[must_use]
pub fn entries(&self) -> &[RestoreReferenceRecord] {
&self.restores
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.restores.is_empty()
}
pub(crate) fn retain(
&mut self,
reference: RestoreReferenceRecord,
) -> Result<bool, RestoreReferenceError> {
if let Some(existing) = self
.restores
.iter()
.find(|entry| entry.journal == reference.journal)
{
if existing != &reference {
return Err(RestoreReferenceError::AuthorityConflict {
journal: reference.journal,
});
}
return Ok(false);
}
if self.restores.len() == MAX_RESTORE_REFERENCES {
return Err(RestoreReferenceError::TooManyReferences {
limit: MAX_RESTORE_REFERENCES,
});
}
self.restores.push(reference);
self.restores
.sort_by(|left, right| left.journal.cmp(&right.journal));
Ok(true)
}
}
fn read_bounded_references<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Vec<RestoreReferenceRecord>, D::Error> {
struct ReferencesVisitor;
impl<'de> de::Visitor<'de> for ReferencesVisitor {
type Value = Vec<RestoreReferenceRecord>;
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("a bounded list of restore references")
}
fn visit_seq<A: de::SeqAccess<'de>>(
self,
mut sequence: A,
) -> Result<Self::Value, A::Error> {
let mut entries = Vec::new();
while let Some(entry) = sequence.next_element()? {
if entries.len() == MAX_RESTORE_REFERENCES {
return Err(de::Error::custom("restore reference count exceeds limit"));
}
entries.push(entry);
}
Ok(entries)
}
}
deserializer.deserialize_seq(ReferencesVisitor)
}
#[derive(Debug, Error)]
pub enum RestoreReferenceError {
#[error("invalid restore journal location: {journal:?}")]
InvalidJournal {
journal: PathBuf,
},
#[error(transparent)]
Checksum(#[from] ChecksumError),
#[error("unsupported restore references version {0}")]
UnsupportedVersion(u16),
#[error("duplicate restore journal location: {journal:?}")]
DuplicateJournal {
journal: PathBuf,
},
#[error("restore authority conflict at {journal:?}")]
AuthorityConflict {
journal: PathBuf,
},
#[error("restore reference count exceeds {limit}")]
TooManyReferences {
limit: usize,
},
}
#[cfg(test)]
mod tests;