use sha2::{Digest, Sha256};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use thiserror::Error;
pub const RNG_RECORDS_METADATA_KEY: &str = "rng_records";
const REPLICATE_SEED_METHOD: &str = "sha256-domain-separated";
const REPLICATE_SEED_VERSION: &str = "scientific-workflow.replicate-seed.v1";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ReplicateSeedDeriver {
base_seed: u64,
replicate_index: u64,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct DerivedSeed {
value: u64,
record: RngRecord,
}
impl ReplicateSeedDeriver {
pub const fn new(base_seed: u64, replicate_index: u64) -> Self {
Self {
base_seed,
replicate_index,
}
}
pub const fn base_seed(self) -> u64 {
self.base_seed
}
pub const fn replicate_index(self) -> u64 {
self.replicate_index
}
pub fn derive(self, namespace: &str) -> Result<DerivedSeed, RngRecordError> {
if namespace.trim().is_empty() {
return Err(RngRecordError::EmptySeedNamespace);
}
let namespace_length =
u64::try_from(namespace.len()).map_err(|_| RngRecordError::SeedNamespaceTooLong)?;
let mut hasher = Sha256::new();
hasher.update(REPLICATE_SEED_VERSION.as_bytes());
hasher.update([0]);
hasher.update(self.base_seed.to_be_bytes());
hasher.update(self.replicate_index.to_be_bytes());
hasher.update(namespace_length.to_be_bytes());
hasher.update(namespace.as_bytes());
let digest = hasher.finalize();
let value = u64::from_be_bytes(
digest[..8]
.try_into()
.expect("SHA-256 always contains at least eight bytes"),
);
let parameters = Map::from_iter([
("base_seed".to_owned(), Value::from(self.base_seed)),
(
"replicate_index".to_owned(),
Value::from(self.replicate_index),
),
]);
let record = RngRecord::new(
namespace,
REPLICATE_SEED_METHOD,
REPLICATE_SEED_VERSION,
"u64-decimal",
value.to_string(),
Some(parameters),
)?;
Ok(DerivedSeed { value, record })
}
}
impl DerivedSeed {
pub const fn value(&self) -> u64 {
self.value
}
pub const fn record(&self) -> &RngRecord {
&self.record
}
pub fn into_parts(self) -> (u64, RngRecord) {
(self.value, self.record)
}
}
#[derive(Clone, Debug, Eq, PartialEq, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RngRecord {
namespace: String,
method: String,
version: String,
key_encoding: String,
key: String,
#[serde(default, skip_serializing_if = "Map::is_empty")]
parameters: Map<String, Value>,
}
impl RngRecord {
pub fn new(
namespace: impl Into<String>,
method: impl Into<String>,
version: impl Into<String>,
key_encoding: impl Into<String>,
key: impl Into<String>,
parameters: Option<Map<String, Value>>,
) -> Result<Self, RngRecordError> {
let record = Self {
namespace: namespace.into(),
method: method.into(),
version: version.into(),
key_encoding: key_encoding.into(),
key: key.into(),
parameters: parameters.unwrap_or_default(),
};
record.validate()?;
Ok(record)
}
pub fn namespace(&self) -> &str {
&self.namespace
}
pub fn method(&self) -> &str {
&self.method
}
pub fn version(&self) -> &str {
&self.version
}
pub fn key_encoding(&self) -> &str {
&self.key_encoding
}
pub fn key(&self) -> &str {
&self.key
}
pub const fn parameters(&self) -> &Map<String, Value> {
&self.parameters
}
pub fn insert_into_metadata(
&self,
metadata: &mut Map<String, Value>,
) -> Result<(), RngRecordError> {
self.validate()?;
let records = metadata
.entry(RNG_RECORDS_METADATA_KEY.to_owned())
.or_insert_with(|| Value::Object(Map::new()))
.as_object_mut()
.ok_or(RngRecordError::InvalidMetadataShape)?;
if records.contains_key(&self.namespace) {
return Err(RngRecordError::DuplicateNamespace {
namespace: self.namespace.clone(),
});
}
records.insert(
self.namespace.clone(),
serde_json::to_value(self).expect("RNG records contain only JSON-compatible values"),
);
Ok(())
}
pub fn from_metadata(
metadata: &Map<String, Value>,
namespace: &str,
) -> Result<Option<Self>, RngRecordError> {
let Some(value) = metadata.get(RNG_RECORDS_METADATA_KEY) else {
return Ok(None);
};
let records = value
.as_object()
.ok_or(RngRecordError::InvalidMetadataShape)?;
let Some(value) = records.get(namespace) else {
return Ok(None);
};
let record: Self = serde_json::from_value(value.clone()).map_err(|source| {
RngRecordError::InvalidStoredRecord {
namespace: namespace.to_owned(),
source,
}
})?;
record.validate()?;
if record.namespace != namespace {
return Err(RngRecordError::NamespaceMismatch {
index: namespace.to_owned(),
record: record.namespace,
});
}
Ok(Some(record))
}
fn validate(&self) -> Result<(), RngRecordError> {
for (field, value) in [
("namespace", self.namespace.as_str()),
("method", self.method.as_str()),
("version", self.version.as_str()),
("key_encoding", self.key_encoding.as_str()),
("key", self.key.as_str()),
] {
if value.trim().is_empty() {
return Err(RngRecordError::EmptyField { field });
}
}
Ok(())
}
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum RngRecordError {
#[error("replicate seed namespace must not be empty or whitespace-only")]
EmptySeedNamespace,
#[error("replicate seed namespace is too long")]
SeedNamespaceTooLong,
#[error("RNG record field `{field}` must not be empty")]
EmptyField {
field: &'static str,
},
#[error("RNG namespace `{namespace}` is recorded more than once")]
DuplicateNamespace {
namespace: String,
},
#[error("user metadata `{RNG_RECORDS_METADATA_KEY}` entry must be an object")]
InvalidMetadataShape,
#[error("invalid RNG record for namespace `{namespace}`")]
InvalidStoredRecord {
namespace: String,
#[source]
source: serde_json::Error,
},
#[error("RNG metadata index `{index}` contains record namespace `{record}`")]
NamespaceMismatch {
index: String,
record: String,
},
}