#![allow(clippy::mutable_key_type)]
use crate::cmd::LocalSwarmCmd;
use crate::driver::MAX_PACKET_SIZE;
use crate::send_local_swarm_cmd;
use crate::target_arch::{spawn, Instant};
use crate::{event::NetworkEvent, log_markers::Marker};
use aes_gcm_siv::{
aead::{Aead, KeyInit},
Aes256GcmSiv, Key as AesKey, Nonce,
};
use hkdf::Hkdf;
use itertools::Itertools;
use libp2p::{
identity::PeerId,
kad::{
store::{Error, RecordStore, Result},
KBucketDistance as Distance, ProviderRecord, Record, RecordKey as Key,
},
};
#[cfg(feature = "open-metrics")]
use prometheus_client::metrics::gauge::Gauge;
use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
use serde::{Deserialize, Serialize};
use sha2::Sha256;
use sn_evm::{AttoTokens, QuotingMetrics};
use sn_protocol::{
storage::{RecordHeader, RecordKind, RecordType},
NetworkAddress, PrettyPrintRecordKey,
};
use std::{
borrow::Cow,
collections::{HashMap, HashSet},
fs,
path::{Path, PathBuf},
time::SystemTime,
vec,
};
use tokio::sync::mpsc;
use walkdir::{DirEntry, WalkDir};
use xor_name::XorName;
const MAX_RECORDS_COUNT: usize = 16 * 1024;
const MAX_RECORDS_CACHE_SIZE: usize = 25;
const HISTORICAL_QUOTING_METRICS_FILENAME: &str = "historic_quoting_metrics";
const MAX_STORE_COST: u64 = 1_000_000;
const MIN_STORE_COST: u64 = 1;
fn derive_aes256gcm_siv_from_seed(seed: &[u8; 16]) -> (Aes256GcmSiv, [u8; 4]) {
let salt = b"autonomi_record_store";
let hk = Hkdf::<Sha256>::new(Some(salt), seed);
let mut okm = [0u8; 32];
hk.expand(b"", &mut okm)
.expect("32 bytes is a valid length for HKDF output");
let seeded_key = AesKey::<Aes256GcmSiv>::from_slice(&okm);
let mut nonce_starter = [0u8; 4];
let bytes_to_copy = seed.len().min(nonce_starter.len());
nonce_starter[..bytes_to_copy].copy_from_slice(&seed[..bytes_to_copy]);
trace!("seeded_key is {seeded_key:?} nonce_starter is {nonce_starter:?}");
(Aes256GcmSiv::new(seeded_key), nonce_starter)
}
struct RecordCache {
records_cache: HashMap<Key, (Record, SystemTime)>,
cache_size: usize,
}
impl RecordCache {
fn new(cache_size: usize) -> Self {
RecordCache {
records_cache: HashMap::new(),
cache_size,
}
}
fn remove(&mut self, key: &Key) -> Option<(Record, SystemTime)> {
self.records_cache.remove(key)
}
fn get(&self, key: &Key) -> Option<&(Record, SystemTime)> {
self.records_cache.get(key)
}
fn push_back(&mut self, key: Key, record: Record) {
self.free_up_space();
let _ = self.records_cache.insert(key, (record, SystemTime::now()));
}
fn free_up_space(&mut self) {
while self.records_cache.len() >= self.cache_size {
self.remove_oldest_entry()
}
}
fn remove_oldest_entry(&mut self) {
let mut oldest_timestamp = SystemTime::now();
for (_record, timestamp) in self.records_cache.values() {
if *timestamp < oldest_timestamp {
oldest_timestamp = *timestamp;
}
}
self.records_cache
.retain(|_key, (_record, timestamp)| *timestamp != oldest_timestamp);
}
}
pub struct NodeRecordStore {
local_address: NetworkAddress,
config: NodeRecordStoreConfig,
records: HashMap<Key, (NetworkAddress, RecordType)>,
records_by_bucket: HashMap<u32, HashSet<Key>>,
records_cache: RecordCache,
network_event_sender: mpsc::Sender<NetworkEvent>,
local_swarm_cmd_sender: mpsc::Sender<LocalSwarmCmd>,
responsible_distance_range: Option<u32>,
#[cfg(feature = "open-metrics")]
record_count_metric: Option<Gauge>,
received_payment_count: usize,
encryption_details: (Aes256GcmSiv, [u8; 4]),
timestamp: SystemTime,
farthest_record: Option<(Key, Distance)>,
}
#[derive(Debug, Clone)]
pub struct NodeRecordStoreConfig {
pub storage_dir: PathBuf,
pub historic_quote_dir: PathBuf,
pub max_records: usize,
pub max_value_bytes: usize,
pub records_cache_size: usize,
pub encryption_seed: [u8; 16],
}
impl Default for NodeRecordStoreConfig {
fn default() -> Self {
let historic_quote_dir = std::env::temp_dir();
Self {
storage_dir: historic_quote_dir.clone(),
historic_quote_dir,
max_records: MAX_RECORDS_COUNT,
max_value_bytes: MAX_PACKET_SIZE,
records_cache_size: MAX_RECORDS_CACHE_SIZE,
encryption_seed: [0u8; 16],
}
}
}
fn generate_nonce_for_record(nonce_starter: &[u8; 4], key: &Key) -> Nonce {
let mut nonce_bytes = nonce_starter.to_vec();
nonce_bytes.extend_from_slice(key.as_ref());
nonce_bytes.resize(12, 0); Nonce::from_iter(nonce_bytes)
}
#[derive(Clone, Serialize, Deserialize)]
struct HistoricQuotingMetrics {
received_payment_count: usize,
timestamp: SystemTime,
}
impl NodeRecordStore {
fn update_records_from_an_existing_store(
config: &NodeRecordStoreConfig,
encryption_details: &(Aes256GcmSiv, [u8; 4]),
) -> HashMap<Key, (NetworkAddress, RecordType)> {
let process_entry = |entry: &DirEntry| -> _ {
let path = entry.path();
if path.is_file() {
debug!("Existing record found: {path:?}");
let filename = match path.file_name().and_then(|n| n.to_str()) {
Some(file_name) => file_name,
None => {
warn!(
"Found a file in the storage dir that is not a valid record: {:?}",
path
);
if let Err(e) = fs::remove_file(path) {
warn!(
"Failed to remove invalid record file from storage dir: {:?}",
e
);
}
return None;
}
};
let key = Self::get_data_from_filename(filename)?;
let record = match fs::read(path) {
Ok(bytes) => {
if let Some(record) =
Self::get_record_from_bytes(bytes, &key, encryption_details)
{
record
} else {
info!("Failed to decrypt record from file {filename:?}, clean it up.");
if let Err(e) = fs::remove_file(path) {
warn!(
"Failed to remove outdated record file {filename:?} from storage dir: {:?}",
e
);
}
return None;
}
}
Err(err) => {
error!("Error while reading file. filename: {filename}, error: {err:?}");
return None;
}
};
let record_type = match RecordHeader::is_record_of_type_chunk(&record) {
Ok(true) => RecordType::Chunk,
Ok(false) => {
let xorname_hash = XorName::from_content(&record.value);
RecordType::NonChunk(xorname_hash)
}
Err(error) => {
warn!(
"Failed to parse record type of record {filename:?}: {:?}",
error
);
if let Err(e) = fs::remove_file(path) {
warn!(
"Failed to remove invalid record file {filename:?} from storage dir: {:?}",
e
);
}
return None;
}
};
let address = NetworkAddress::from_record_key(&key);
info!("Existing record loaded: {path:?}");
return Some((key, (address, record_type)));
}
None
};
info!("Attempting to repopulate records from existing store...");
let records = WalkDir::new(&config.storage_dir)
.into_iter()
.filter_map(|e| e.ok())
.collect_vec()
.par_iter()
.filter_map(process_entry)
.collect();
records
}
fn restore_quoting_metrics(storage_dir: &Path) -> Option<HistoricQuotingMetrics> {
let file_path = storage_dir.join(HISTORICAL_QUOTING_METRICS_FILENAME);
if let Ok(file) = fs::File::open(file_path) {
if let Ok(quoting_metrics) = rmp_serde::from_read(&file) {
return Some(quoting_metrics);
}
}
None
}
fn flush_historic_quoting_metrics(&self) {
let file_path = self
.config
.historic_quote_dir
.join(HISTORICAL_QUOTING_METRICS_FILENAME);
let historic_quoting_metrics = HistoricQuotingMetrics {
received_payment_count: self.received_payment_count,
timestamp: self.timestamp,
};
spawn(async move {
if let Ok(mut file) = fs::File::create(file_path) {
let mut serialiser = rmp_serde::encode::Serializer::new(&mut file);
let _ = historic_quoting_metrics.serialize(&mut serialiser);
}
});
}
pub fn with_config(
local_id: PeerId,
config: NodeRecordStoreConfig,
network_event_sender: mpsc::Sender<NetworkEvent>,
swarm_cmd_sender: mpsc::Sender<LocalSwarmCmd>,
) -> Self {
info!("Using encryption_seed of {:?}", config.encryption_seed);
let encryption_details = derive_aes256gcm_siv_from_seed(&config.encryption_seed);
let (received_payment_count, timestamp) = if let Some(historic_quoting_metrics) =
Self::restore_quoting_metrics(&config.historic_quote_dir)
{
(
historic_quoting_metrics.received_payment_count,
historic_quoting_metrics.timestamp,
)
} else {
(0, SystemTime::now())
};
let records = Self::update_records_from_an_existing_store(&config, &encryption_details);
let local_address = NetworkAddress::from_peer(local_id);
let mut records_by_bucket: HashMap<u32, HashSet<Key>> = HashMap::new();
for (key, (addr, _record_type)) in records.iter() {
let distance = local_address.distance(addr);
let bucket = distance.ilog2().unwrap_or_default();
records_by_bucket
.entry(bucket)
.or_default()
.insert(key.clone());
}
let cache_size = config.records_cache_size;
let mut record_store = NodeRecordStore {
local_address,
config,
records,
records_by_bucket,
records_cache: RecordCache::new(cache_size),
network_event_sender,
local_swarm_cmd_sender: swarm_cmd_sender,
responsible_distance_range: None,
#[cfg(feature = "open-metrics")]
record_count_metric: None,
received_payment_count,
encryption_details,
timestamp,
farthest_record: None,
};
record_store.farthest_record = record_store.calculate_farthest();
record_store.flush_historic_quoting_metrics();
record_store
}
#[cfg(feature = "open-metrics")]
pub fn set_record_count_metric(mut self, metric: Gauge) -> Self {
self.record_count_metric = Some(metric);
self
}
pub fn get_responsible_distance_range(&self) -> Option<u32> {
self.responsible_distance_range
}
fn generate_filename(key: &Key) -> String {
hex::encode(key.as_ref())
}
fn get_data_from_filename(hex_str: &str) -> Option<Key> {
match hex::decode(hex_str) {
Ok(bytes) => Some(Key::from(bytes)),
Err(error) => {
error!("Error decoding hex string: {:?}", error);
None
}
}
}
fn get_record_from_bytes<'a>(
bytes: Vec<u8>,
key: &Key,
encryption_details: &(Aes256GcmSiv, [u8; 4]),
) -> Option<Cow<'a, Record>> {
let mut record = Record {
key: key.clone(),
value: bytes,
publisher: None,
expires: None,
};
if !cfg!(feature = "encrypt-records") {
return Some(Cow::Owned(record));
}
let (cipher, nonce_starter) = encryption_details;
let nonce = generate_nonce_for_record(nonce_starter, key);
match cipher.decrypt(&nonce, record.value.as_ref()) {
Ok(value) => {
record.value = value;
return Some(Cow::Owned(record));
}
Err(error) => {
error!("Error while decrypting record. key: {key:?}: {error:?}");
None
}
}
}
fn read_from_disk<'a>(
encryption_details: &(Aes256GcmSiv, [u8; 4]),
key: &Key,
storage_dir: &Path,
) -> Option<Cow<'a, Record>> {
let start = Instant::now();
let filename = Self::generate_filename(key);
let file_path = storage_dir.join(&filename);
match fs::read(file_path) {
Ok(bytes) => {
info!(
"Retrieved record from disk! filename: {filename} after {:?}",
start.elapsed()
);
Self::get_record_from_bytes(bytes, key, encryption_details)
}
Err(err) => {
error!("Error while reading file. filename: {filename}, error: {err:?}");
None
}
}
}
pub fn get_farthest(&self) -> Option<Key> {
if let Some((ref key, _distance)) = self.farthest_record {
Some(key.clone())
} else {
None
}
}
fn calculate_farthest(&self) -> Option<(Key, Distance)> {
let mut sorted_records: Vec<_> = self.records.keys().collect();
sorted_records.sort_by_key(|key| {
let addr = NetworkAddress::from_record_key(key);
self.local_address.distance(&addr)
});
if let Some(key) = sorted_records.last() {
let addr = NetworkAddress::from_record_key(key);
Some(((*key).clone(), self.local_address.distance(&addr)))
} else {
None
}
}
fn prune_records_if_needed(&mut self, incoming_record_key: &Key) -> Result<()> {
if self.records.len() < self.config.max_records {
return Ok(());
}
if let Some((farthest_record, farthest_record_distance)) = self.farthest_record.clone() {
if farthest_record_distance
< self
.local_address
.distance(&NetworkAddress::from_record_key(incoming_record_key))
{
return Err(Error::MaxRecords);
}
info!(
"Record {:?} will be pruned to free up space for new records",
PrettyPrintRecordKey::from(&farthest_record)
);
self.remove(&farthest_record);
}
Ok(())
}
pub fn cleanup_irrelevant_records(&mut self) {
let accumulated_records = self.records.len();
if accumulated_records < MAX_RECORDS_COUNT / 10 {
return;
}
let max_bucket = if let Some(range) = self.responsible_distance_range {
if range == 0 {
return;
}
range
} else {
return;
};
let keys_to_remove: Vec<Key> = self
.records_by_bucket
.iter()
.filter(|(&bucket, _)| bucket > max_bucket)
.flat_map(|(_, keys)| keys.iter().cloned())
.collect();
let keys_to_remove_len = keys_to_remove.len();
for key in keys_to_remove {
self.remove(&key);
}
info!("Cleaned up {} unrelevant records, among the original {accumulated_records} accumulated_records",
keys_to_remove_len);
}
}
impl NodeRecordStore {
pub(crate) fn contains(&self, key: &Key) -> bool {
self.records.contains_key(key)
}
pub(crate) fn record_addresses(&self) -> HashMap<NetworkAddress, RecordType> {
self.records
.iter()
.map(|(_record_key, (addr, record_type))| (addr.clone(), record_type.clone()))
.collect()
}
pub(crate) fn record_addresses_ref(&self) -> &HashMap<Key, (NetworkAddress, RecordType)> {
&self.records
}
pub(crate) fn mark_as_stored(&mut self, key: Key, record_type: RecordType) {
let addr = NetworkAddress::from_record_key(&key);
let distance = self.local_address.distance(&addr);
let bucket = distance.ilog2().unwrap_or_default();
self.records
.insert(key.clone(), (addr.clone(), record_type));
self.records_by_bucket
.entry(bucket)
.or_default()
.insert(key.clone());
if let Some((_farthest_record, farthest_record_distance)) = self.farthest_record.clone() {
if distance > farthest_record_distance {
self.farthest_record = Some((key, distance));
}
} else {
self.farthest_record = Some((key, distance));
}
}
fn prepare_record_bytes(
record: Record,
encryption_details: (Aes256GcmSiv, [u8; 4]),
) -> Option<Vec<u8>> {
if !cfg!(feature = "encrypt-records") {
return Some(record.value);
}
let (cipher, nonce_starter) = encryption_details;
let nonce = generate_nonce_for_record(&nonce_starter, &record.key);
match cipher.encrypt(&nonce, record.value.as_ref()) {
Ok(value) => Some(value),
Err(error) => {
warn!(
"Failed to encrypt record {:?} : {error:?}",
PrettyPrintRecordKey::from(&record.key),
);
None
}
}
}
pub(crate) fn put_verified(&mut self, r: Record, record_type: RecordType) -> Result<()> {
let key = &r.key;
let record_key = PrettyPrintRecordKey::from(&r.key).into_owned();
debug!("PUTting a verified Record: {record_key:?}");
if let Some((existing_record, _timestamp)) = self.records_cache.remove(key) {
if existing_record.value == r.value {
self.records_cache.push_back(key.clone(), existing_record);
return Ok(());
}
}
self.records_cache.push_back(key.clone(), r.clone());
self.prune_records_if_needed(key)?;
let filename = Self::generate_filename(key);
let file_path = self.config.storage_dir.join(&filename);
#[cfg(feature = "open-metrics")]
if let Some(metric) = &self.record_count_metric {
let _ = metric.set(self.records.len() as i64);
}
let encryption_details = self.encryption_details.clone();
let cloned_cmd_sender = self.local_swarm_cmd_sender.clone();
let record_key2 = record_key.clone();
spawn(async move {
let key = r.key.clone();
if let Some(bytes) = Self::prepare_record_bytes(r, encryption_details) {
let cmd = match fs::write(&file_path, bytes) {
Ok(_) => {
info!("Wrote record {record_key2:?} to disk! filename: {filename}");
LocalSwarmCmd::AddLocalRecordAsStored { key, record_type }
}
Err(err) => {
error!(
"Error writing record {record_key2:?} filename: {filename}, error: {err:?}"
);
LocalSwarmCmd::RemoveFailedLocalRecord { key }
}
};
send_local_swarm_cmd(cloned_cmd_sender, cmd);
}
});
Ok(())
}
pub(crate) fn store_cost(&self, key: &Key) -> (AttoTokens, QuotingMetrics) {
let records_stored = self.records.len();
let record_keys_as_hashset: HashSet<&Key> = self.records.keys().collect();
let live_time = if let Ok(elapsed) = self.timestamp.elapsed() {
elapsed.as_secs()
} else {
0
};
let mut quoting_metrics = QuotingMetrics {
close_records_stored: records_stored,
max_records: self.config.max_records,
received_payment_count: self.received_payment_count,
live_time,
};
if let Some(distance_range) = self.responsible_distance_range {
let relevant_records =
self.get_records_within_distance_range(record_keys_as_hashset, distance_range);
quoting_metrics.close_records_stored = relevant_records;
} else {
info!("Basing cost of _total_ records stored.");
};
let cost = if self.contains(key) {
0
} else {
calculate_cost_for_records(quoting_metrics.close_records_stored)
};
info!("Cost is now {cost:?} for quoting_metrics {quoting_metrics:?}");
(AttoTokens::from_u64(cost), quoting_metrics)
}
pub(crate) fn payment_received(&mut self) {
self.received_payment_count = self.received_payment_count.saturating_add(1);
self.flush_historic_quoting_metrics();
}
pub fn get_records_within_distance_range(
&self,
_records: HashSet<&Key>,
max_bucket: u32,
) -> usize {
let within_range = self
.records_by_bucket
.iter()
.filter(|(&bucket, _)| bucket <= max_bucket)
.map(|(_, keys)| keys.len())
.sum();
Marker::CloseRecordsLen(within_range).log();
within_range
}
pub(crate) fn set_responsible_distance_range(&mut self, farthest_responsible_bucket: u32) {
self.responsible_distance_range = Some(farthest_responsible_bucket);
}
}
impl RecordStore for NodeRecordStore {
type RecordsIter<'a> = vec::IntoIter<Cow<'a, Record>>;
type ProvidedIter<'a> = vec::IntoIter<Cow<'a, ProviderRecord>>;
fn get(&self, k: &Key) -> Option<Cow<'_, Record>> {
let key = PrettyPrintRecordKey::from(k);
let cached_record = self.records_cache.get(k);
if let Some((record, _timestamp)) = cached_record {
return Some(Cow::Borrowed(record));
}
if !self.records.contains_key(k) {
debug!("Record not found locally: {key:?}");
return None;
}
debug!("GET request for Record key: {key}");
Self::read_from_disk(&self.encryption_details, k, &self.config.storage_dir)
}
fn put(&mut self, record: Record) -> Result<()> {
let record_key = PrettyPrintRecordKey::from(&record.key);
if record.value.len() >= self.config.max_value_bytes {
warn!(
"Record {record_key:?} not stored. Value too large: {} bytes",
record.value.len()
);
return Err(Error::ValueTooLarge);
}
match RecordHeader::from_record(&record) {
Ok(record_header) => {
match record_header.kind {
RecordKind::ChunkWithPayment | RecordKind::RegisterWithPayment => {
debug!("Record {record_key:?} with payment shall always be processed.");
}
_ => {
match self.records.get(&record.key) {
Some((_addr, RecordType::Chunk)) => {
debug!("Chunk {record_key:?} already exists.");
return Ok(());
}
Some((_addr, RecordType::NonChunk(existing_content_hash))) => {
let content_hash = XorName::from_content(&record.value);
if content_hash == *existing_content_hash {
debug!("A non-chunk record {record_key:?} with same content_hash {content_hash:?} already exists.");
return Ok(());
}
}
_ => {}
}
}
}
}
Err(err) => {
error!("For record {record_key:?}, failed to parse record_header {err:?}");
return Ok(());
}
}
debug!("Unverified Record {record_key:?} try to validate and store");
let event_sender = self.network_event_sender.clone();
let _handle = spawn(async move {
if let Err(error) = event_sender
.send(NetworkEvent::UnverifiedRecord(record))
.await
{
error!("SwarmDriver failed to send event: {}", error);
}
});
Ok(())
}
fn remove(&mut self, k: &Key) {
if let Some((addr, _)) = self.records.remove(k) {
let bucket = self
.local_address
.distance(&addr)
.ilog2()
.unwrap_or_default();
if let Some(bucket_keys) = self.records_by_bucket.get_mut(&bucket) {
bucket_keys.remove(k);
if bucket_keys.is_empty() {
self.records_by_bucket.remove(&bucket);
}
}
}
self.records_cache.remove(k);
#[cfg(feature = "open-metrics")]
if let Some(metric) = &self.record_count_metric {
let _ = metric.set(self.records.len() as i64);
}
if let Some((farthest_record, _)) = self.farthest_record.clone() {
if farthest_record == *k {
self.farthest_record = self.calculate_farthest();
}
}
let filename = Self::generate_filename(k);
let file_path = self.config.storage_dir.join(&filename);
let _handle = spawn(async move {
match fs::remove_file(file_path) {
Ok(_) => {
info!("Removed record from disk! filename: {filename}");
}
Err(err) => {
error!("Error while removing file. filename: {filename}, error: {err:?}");
}
}
});
}
fn records(&self) -> Self::RecordsIter<'_> {
vec![].into_iter()
}
fn add_provider(&mut self, _record: ProviderRecord) -> Result<()> {
Ok(())
}
fn providers(&self, _key: &Key) -> Vec<ProviderRecord> {
vec![]
}
fn provided(&self) -> Self::ProvidedIter<'_> {
vec![].into_iter()
}
fn remove_provider(&mut self, _key: &Key, _provider: &PeerId) {
}
}
#[derive(Default, Debug)]
pub struct ClientRecordStore {
empty_record_addresses: HashMap<Key, (NetworkAddress, RecordType)>,
}
impl ClientRecordStore {
pub(crate) fn contains(&self, _key: &Key) -> bool {
false
}
pub(crate) fn record_addresses(&self) -> HashMap<NetworkAddress, RecordType> {
HashMap::new()
}
pub(crate) fn record_addresses_ref(&self) -> &HashMap<Key, (NetworkAddress, RecordType)> {
&self.empty_record_addresses
}
pub(crate) fn put_verified(&mut self, _r: Record, _record_type: RecordType) -> Result<()> {
Ok(())
}
pub(crate) fn mark_as_stored(&mut self, _r: Key, _t: RecordType) {}
}
impl RecordStore for ClientRecordStore {
type RecordsIter<'a> = vec::IntoIter<Cow<'a, Record>>;
type ProvidedIter<'a> = vec::IntoIter<Cow<'a, ProviderRecord>>;
fn get(&self, _k: &Key) -> Option<Cow<'_, Record>> {
None
}
fn put(&mut self, _record: Record) -> Result<()> {
Ok(())
}
fn remove(&mut self, _k: &Key) {}
fn records(&self) -> Self::RecordsIter<'_> {
vec![].into_iter()
}
fn add_provider(&mut self, _record: ProviderRecord) -> Result<()> {
Ok(())
}
fn providers(&self, _key: &Key) -> Vec<ProviderRecord> {
vec![]
}
fn provided(&self) -> Self::ProvidedIter<'_> {
vec![].into_iter()
}
fn remove_provider(&mut self, _key: &Key, _provider: &PeerId) {}
}
pub fn calculate_cost_for_records(records_stored: usize) -> u64 {
use std::cmp::{max, min};
let max_records = MAX_RECORDS_COUNT;
let ori_cost = positive_input_0_1_sigmoid(records_stored as f64 / max_records as f64)
* MAX_STORE_COST as f64;
let charge = max(MIN_STORE_COST, ori_cost as u64);
min(MAX_STORE_COST, charge)
}
fn positive_input_0_1_sigmoid(x: f64) -> f64 {
1.0 / (1.0 + (-30.0 * (x - 0.5)).exp())
}
#[expect(trivial_casts)]
#[cfg(test)]
mod tests {
use crate::get_fees_from_store_cost_responses;
use super::*;
use bls::SecretKey;
use xor_name::XorName;
use assert_fs::TempDir;
use bytes::Bytes;
use eyre::{bail, ContextCompat};
use libp2p::kad::K_VALUE;
use libp2p::{core::multihash::Multihash, kad::RecordKey};
use quickcheck::*;
use sn_evm::utils::dummy_address;
use sn_evm::{PaymentQuote, RewardsAddress};
use sn_protocol::storage::{
try_deserialize_record, try_serialize_record, Chunk, ChunkAddress, Scratchpad,
};
use std::collections::BTreeMap;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use tokio::runtime::Runtime;
use tokio::time::{sleep, Duration};
const MULITHASH_CODE: u64 = 0x12;
#[derive(Clone, Debug)]
struct ArbitraryKey(Key);
#[derive(Clone, Debug)]
struct ArbitraryRecord(Record);
impl Arbitrary for ArbitraryKey {
fn arbitrary(g: &mut Gen) -> ArbitraryKey {
let hash: [u8; 32] = core::array::from_fn(|_| u8::arbitrary(g));
ArbitraryKey(Key::from(
Multihash::<64>::wrap(MULITHASH_CODE, &hash).expect("Failed to gen MultiHash"),
))
}
}
impl Arbitrary for ArbitraryRecord {
fn arbitrary(g: &mut Gen) -> ArbitraryRecord {
let value = match try_serialize_record(
&(0..50).map(|_| rand::random::<u8>()).collect::<Bytes>(),
RecordKind::Chunk,
) {
Ok(value) => value.to_vec(),
Err(err) => panic!("Cannot generate record value {err:?}"),
};
let record = Record {
key: ArbitraryKey::arbitrary(g).0,
value,
publisher: None,
expires: None,
};
ArbitraryRecord(record)
}
}
#[test]
fn test_calculate_max_cost_for_records() {
let sut = calculate_cost_for_records(MAX_RECORDS_COUNT + 1);
assert_eq!(sut, MAX_STORE_COST - 1);
}
#[test]
fn test_calculate_50_percent_cost_for_records() {
let percent = MAX_RECORDS_COUNT * 50 / 100;
let sut = calculate_cost_for_records(percent);
assert_eq!(sut, 500000);
}
#[test]
fn test_calculate_60_percent_cost_for_records() {
let percent = MAX_RECORDS_COUNT * 60 / 100;
let sut = calculate_cost_for_records(percent);
assert_eq!(sut, 952541);
}
#[test]
fn test_calculate_65_percent_cost_for_records() {
let percent = MAX_RECORDS_COUNT * 65 / 100;
let sut = calculate_cost_for_records(percent);
assert_eq!(sut, 989001);
}
#[test]
fn test_calculate_70_percent_cost_for_records() {
let percent = MAX_RECORDS_COUNT * 70 / 100;
let sut = calculate_cost_for_records(percent);
assert_eq!(sut, 997523);
}
#[test]
fn test_calculate_80_percent_cost_for_records() {
let percent = MAX_RECORDS_COUNT * 80 / 100;
let sut = calculate_cost_for_records(percent);
assert_eq!(sut, 999876);
}
#[test]
fn test_calculate_90_percent_cost_for_records() {
let percent = MAX_RECORDS_COUNT * 90 / 100;
let sut = calculate_cost_for_records(percent);
assert_eq!(sut, 999993);
}
#[test]
fn test_calculate_min_cost_for_records() {
let sut = calculate_cost_for_records(0);
assert_eq!(sut, MIN_STORE_COST);
}
#[test]
fn put_get_remove_record() {
fn prop(r: ArbitraryRecord) {
let rt = if let Ok(rt) = Runtime::new() {
rt
} else {
panic!("Cannot create runtime");
};
rt.block_on(testing_thread(r));
}
quickcheck(prop as fn(_))
}
async fn testing_thread(r: ArbitraryRecord) {
let r = r.0;
let (network_event_sender, mut network_event_receiver) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
PeerId::random(),
Default::default(),
network_event_sender,
swarm_cmd_sender,
);
let store_cost_before = store.store_cost(&r.key);
assert!(store.put(r.clone()).is_ok());
assert!(store.get(&r.key).is_none());
assert_eq!(
store.store_cost(&r.key).0,
store_cost_before.0,
"store cost should not change over unverified put"
);
let returned_record = if let Some(event) = network_event_receiver.recv().await {
if let NetworkEvent::UnverifiedRecord(record) = event {
record
} else {
panic!("Unexpected network event {event:?}");
}
} else {
panic!("Failed recevied the record for further verification");
};
let returned_record_key = returned_record.key.clone();
assert!(store
.put_verified(returned_record, RecordType::Chunk)
.is_ok());
store.mark_as_stored(returned_record_key, RecordType::Chunk);
let max_iterations = 10;
let mut iteration = 0;
while iteration < max_iterations {
if store
.get(&r.key)
.is_some_and(|record| Cow::Borrowed(&r) == record)
{
break;
}
sleep(Duration::from_millis(100)).await;
iteration += 1;
}
if iteration == max_iterations {
panic!("record_store test failed with stored record cann't be read back");
}
assert_eq!(
Some(Cow::Borrowed(&r)),
store.get(&r.key),
"record can be retrieved after put"
);
store.remove(&r.key);
assert!(store.get(&r.key).is_none());
}
#[tokio::test]
#[ignore = "fails on ci"]
async fn can_store_after_restart() -> eyre::Result<()> {
let temp_dir = TempDir::new().expect("Should be able to create a temp dir.");
let store_config = NodeRecordStoreConfig {
storage_dir: temp_dir.to_path_buf(),
encryption_seed: [1u8; 16],
..Default::default()
};
let self_id = PeerId::random();
let (network_event_sender, _) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
self_id,
store_config.clone(),
network_event_sender.clone(),
swarm_cmd_sender.clone(),
);
let chunk_data = Bytes::from_static(b"Test chunk data");
let chunk = Chunk::new(chunk_data);
let chunk_address = *chunk.address();
let record = Record {
key: NetworkAddress::ChunkAddress(chunk_address).to_record_key(),
value: try_serialize_record(&chunk, RecordKind::Chunk)?.to_vec(),
expires: None,
publisher: None,
};
assert!(store
.put_verified(record.clone(), RecordType::Chunk)
.is_ok());
store.mark_as_stored(record.key.clone(), RecordType::Chunk);
let stored_record = store.get(&record.key);
assert!(stored_record.is_some(), "Chunk should be stored");
sleep(Duration::from_secs(1)).await;
drop(store);
let store = NodeRecordStore::with_config(
self_id,
store_config,
network_event_sender.clone(),
swarm_cmd_sender.clone(),
);
sleep(Duration::from_secs(1)).await;
let stored_record = store.get(&record.key);
assert!(stored_record.is_some(), "Chunk should be stored");
let self_id_diff = PeerId::random();
let store_config_diff = NodeRecordStoreConfig {
storage_dir: temp_dir.to_path_buf(),
encryption_seed: [2u8; 16],
..Default::default()
};
let store_diff = NodeRecordStore::with_config(
self_id_diff,
store_config_diff,
network_event_sender,
swarm_cmd_sender,
);
sleep(Duration::from_secs(1)).await;
if cfg!(feature = "encrypt-records") {
assert!(
store_diff.get(&record.key).is_none(),
"Chunk should be gone"
);
} else {
assert!(
store_diff.get(&record.key).is_some(),
"Chunk shall persists without encryption"
);
}
Ok(())
}
#[tokio::test]
async fn can_store_and_retrieve_chunk() {
let temp_dir = std::env::temp_dir();
let store_config = NodeRecordStoreConfig {
storage_dir: temp_dir,
..Default::default()
};
let self_id = PeerId::random();
let (network_event_sender, _) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
self_id,
store_config,
network_event_sender,
swarm_cmd_sender,
);
let chunk_data = Bytes::from_static(b"Test chunk data");
let chunk = Chunk::new(chunk_data.clone());
let chunk_address = *chunk.address();
let record = Record {
key: NetworkAddress::ChunkAddress(chunk_address).to_record_key(),
value: chunk_data.to_vec(),
expires: None,
publisher: None,
};
assert!(store
.put_verified(record.clone(), RecordType::Chunk)
.is_ok());
store.mark_as_stored(record.key.clone(), RecordType::Chunk);
let stored_record = store.get(&record.key);
assert!(stored_record.is_some(), "Chunk should be stored");
if let Some(stored) = stored_record {
assert_eq!(
stored.value, chunk_data,
"Stored chunk data should match original"
);
let stored_address = ChunkAddress::new(XorName::from_content(&stored.value));
assert_eq!(
stored_address, chunk_address,
"Stored chunk address should match original"
);
}
store.remove(&record.key);
assert!(
store.get(&record.key).is_none(),
"Chunk should be removed after cleanup"
);
}
#[tokio::test]
async fn can_store_and_retrieve_scratchpad() -> eyre::Result<()> {
let temp_dir = std::env::temp_dir();
let store_config = NodeRecordStoreConfig {
storage_dir: temp_dir,
..Default::default()
};
let self_id = PeerId::random();
let (network_event_sender, _) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
self_id,
store_config,
network_event_sender,
swarm_cmd_sender,
);
let unencrypted_scratchpad_data = Bytes::from_static(b"Test scratchpad data");
let owner_sk = SecretKey::random();
let owner_pk = owner_sk.public_key();
let mut scratchpad = Scratchpad::new(owner_pk, 0);
let _next_version =
scratchpad.update_and_sign(unencrypted_scratchpad_data.clone(), &owner_sk);
let scratchpad_address = *scratchpad.address();
let record = Record {
key: NetworkAddress::ScratchpadAddress(scratchpad_address).to_record_key(),
value: try_serialize_record(&scratchpad, RecordKind::Scratchpad)?.to_vec(),
expires: None,
publisher: None,
};
assert!(store
.put_verified(
record.clone(),
RecordType::NonChunk(XorName::from_content(&record.value))
)
.is_ok());
store.mark_as_stored(
record.key.clone(),
RecordType::NonChunk(XorName::from_content(&record.value)),
);
let stored_record = store.get(&record.key);
assert!(stored_record.is_some(), "Scratchpad should be stored");
if let Some(stored) = stored_record {
let scratchpad = try_deserialize_record::<Scratchpad>(&stored)?;
let stored_address = scratchpad.address();
assert_eq!(
stored_address, &scratchpad_address,
"Stored scratchpad address should match original"
);
let decrypted_data = scratchpad.decrypt_data(&owner_sk)?;
assert_eq!(
decrypted_data, unencrypted_scratchpad_data,
"Stored scratchpad data should match original"
);
}
store.remove(&record.key);
assert!(
store.get(&record.key).is_none(),
"Scratchpad should be removed after cleanup"
);
Ok(())
}
#[tokio::test]
async fn pruning_on_full() -> Result<()> {
let max_iterations = 10;
let max_records = 50;
let temp_dir = std::env::temp_dir();
let unique_dir_name = uuid::Uuid::new_v4().to_string();
let storage_dir = temp_dir.join(unique_dir_name);
fs::create_dir_all(&storage_dir).expect("Failed to create directory");
let store_config = NodeRecordStoreConfig {
max_records,
storage_dir,
..Default::default()
};
let self_id = PeerId::random();
let (network_event_sender, _) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
self_id,
store_config.clone(),
network_event_sender,
swarm_cmd_sender,
);
let mut stored_records_at_some_point: Vec<RecordKey> = vec![];
let self_address = NetworkAddress::from_peer(self_id);
let mut failed_records = vec![];
for _ in 0..max_records * 2 {
let record_key = NetworkAddress::from_peer(PeerId::random()).to_record_key();
let value = match try_serialize_record(
&(0..50).map(|_| rand::random::<u8>()).collect::<Bytes>(),
RecordKind::Chunk,
) {
Ok(value) => value.to_vec(),
Err(err) => panic!("Cannot generate record value {err:?}"),
};
let record = Record {
key: record_key.clone(),
value,
publisher: None,
expires: None,
};
let succeeded = store.put_verified(record, RecordType::Chunk).is_ok();
if !succeeded {
failed_records.push(record_key.clone());
println!("failed {:?}", PrettyPrintRecordKey::from(&record_key));
} else {
store.mark_as_stored(record_key.clone(), RecordType::Chunk);
println!("success sotred len: {:?} ", store.record_addresses().len());
stored_records_at_some_point.push(record_key.clone());
if stored_records_at_some_point.len() <= max_records {
assert!(succeeded);
}
let mut iteration = 0;
while iteration < max_iterations {
if store.get(&record_key).is_some() {
break;
}
sleep(Duration::from_millis(100)).await;
iteration += 1;
}
if iteration == max_iterations {
panic!("record_store prune test failed with stored record {record_key:?} can't be read back");
}
}
}
let stored_data_at_end = store.record_addresses();
assert!(
stored_data_at_end.len() == max_records,
"Stored records ({:?}) should be max_records, {max_records:?}",
stored_data_at_end.len(),
);
assert!(
stored_records_at_some_point.len() >= max_records,
"we should have stored ata least max over time"
);
let mut sorted_stored_data = stored_data_at_end.iter().collect_vec();
sorted_stored_data
.sort_by(|(a, _), (b, _)| self_address.distance(a).cmp(&self_address.distance(b)));
if let Some((most_distant_data, _)) = sorted_stored_data.last() {
for failed_record in failed_records {
let failed_data = NetworkAddress::from_record_key(&failed_record);
assert!(
self_address.distance(&failed_data) > self_address.distance(most_distant_data),
"failed record {failed_data:?} should be farther than the farthest stored record {most_distant_data:?}"
);
}
for data in stored_records_at_some_point {
let data_addr = NetworkAddress::from_record_key(&data);
if !sorted_stored_data.contains(&(&data_addr, &RecordType::Chunk)) {
assert!(
self_address.distance(&data_addr)
> self_address.distance(most_distant_data),
"stored record should be farther than the farthest stored record"
);
}
}
}
Ok(())
}
#[tokio::test]
async fn get_records_within_bucket_range() -> eyre::Result<()> {
let max_records = 50;
let temp_dir = std::env::temp_dir();
let unique_dir_name = uuid::Uuid::new_v4().to_string();
let storage_dir = temp_dir.join(unique_dir_name);
let store_config = NodeRecordStoreConfig {
max_records,
storage_dir,
..Default::default()
};
let self_id = PeerId::random();
let (network_event_sender, _) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
self_id,
store_config,
network_event_sender,
swarm_cmd_sender,
);
let mut stored_records: Vec<RecordKey> = vec![];
let self_address = NetworkAddress::from_peer(self_id);
for _ in 0..max_records - 1 {
let record_key = NetworkAddress::from_peer(PeerId::random()).to_record_key();
let value = match try_serialize_record(
&(0..max_records)
.map(|_| rand::random::<u8>())
.collect::<Bytes>(),
RecordKind::Chunk,
) {
Ok(value) => value.to_vec(),
Err(err) => panic!("Cannot generate record value {err:?}"),
};
let record = Record {
key: record_key.clone(),
value,
publisher: None,
expires: None,
};
assert!(store.put_verified(record, RecordType::Chunk).is_ok());
store.mark_as_stored(record_key.clone(), RecordType::Chunk);
stored_records.push(record_key);
stored_records.sort_by(|a, b| {
let a = NetworkAddress::from_record_key(a);
let b = NetworkAddress::from_record_key(b);
self_address.distance(&a).cmp(&self_address.distance(&b))
});
}
let halfway_record_address = NetworkAddress::from_record_key(
stored_records
.get((stored_records.len() / 2) - 1)
.wrap_err("Could not parse record store key")?,
);
let distance = self_address
.distance(&halfway_record_address)
.ilog2()
.unwrap_or(0);
store.set_responsible_distance_range(distance);
let record_keys = store.records.keys().collect();
assert!(
store.get_records_within_distance_range(record_keys, distance)
>= stored_records.len() / 2
);
Ok(())
}
#[tokio::test]
async fn historic_quoting_metrics() -> Result<()> {
let temp_dir = std::env::temp_dir();
let unique_dir_name = uuid::Uuid::new_v4().to_string();
let storage_dir = temp_dir.join(unique_dir_name);
fs::create_dir_all(&storage_dir).expect("Failed to create directory");
let historic_quote_dir = storage_dir.clone();
let store_config = NodeRecordStoreConfig {
storage_dir,
historic_quote_dir,
..Default::default()
};
let self_id = PeerId::random();
let (network_event_sender, _) = mpsc::channel(1);
let (swarm_cmd_sender, _) = mpsc::channel(1);
let mut store = NodeRecordStore::with_config(
self_id,
store_config.clone(),
network_event_sender.clone(),
swarm_cmd_sender.clone(),
);
store.payment_received();
sleep(Duration::from_millis(5000)).await;
let new_store = NodeRecordStore::with_config(
self_id,
store_config,
network_event_sender,
swarm_cmd_sender,
);
assert_eq!(1, new_store.received_payment_count);
assert_eq!(store.timestamp, new_store.timestamp);
Ok(())
}
struct PeerStats {
address: NetworkAddress,
rewards_addr: RewardsAddress,
records_stored: AtomicUsize,
nanos_earned: AtomicU64,
payments_received: AtomicUsize,
}
#[ignore]
#[test]
fn address_distribution_sim() {
use rayon::prelude::*;
let num_of_peers = 5_000;
let num_of_chunks_per_hour = 1_000_000;
let max_hours = 50;
let k = K_VALUE.get();
let replication_group_size = k / 3;
let mut peers: Vec<PeerStats> = (0..num_of_peers)
.into_par_iter()
.map(|_| PeerStats {
address: NetworkAddress::from_peer(PeerId::random()),
records_stored: AtomicUsize::new(0),
nanos_earned: AtomicU64::new(0),
payments_received: AtomicUsize::new(0),
rewards_addr: dummy_address(),
})
.collect();
let mut hour = 0;
let mut total_received_payment_count = 0;
let peers_len = peers.len();
let sorting_target_address =
NetworkAddress::from_chunk_address(ChunkAddress::new(XorName::default()));
peers.par_sort_by(|a, b| {
sorting_target_address
.distance(&a.address)
.cmp(&sorting_target_address.distance(&b.address))
});
loop {
let _chunk_results: Vec<_> = (0..num_of_chunks_per_hour)
.into_par_iter()
.map(|_| {
let name = xor_name::rand::random();
let chunk_address = NetworkAddress::from_chunk_address(ChunkAddress::new(name));
let chunk_distance_to_sorting = sorting_target_address.distance(&chunk_address);
let partition_point = peers.partition_point(|peer| {
sorting_target_address.distance(&peer.address) < chunk_distance_to_sorting
});
let mut close_group = Vec::with_capacity(replication_group_size);
let mut left = partition_point;
let mut right = partition_point;
while close_group.len() < replication_group_size
&& (left > 0 || right < peers_len)
{
if left > 0 {
left -= 1;
close_group.push(left);
}
if close_group.len() < replication_group_size && right < peers_len {
close_group.push(right);
right += 1;
}
}
close_group.truncate(replication_group_size);
let Ok((payee_index, cost)) = pick_cheapest_payee(&peers, &close_group) else {
bail!("Failed to find a payee");
};
for &peer_index in &close_group {
let peer = &peers[peer_index];
peer.records_stored.fetch_add(1, Ordering::Relaxed);
if peer_index == payee_index {
peer.nanos_earned.fetch_add(
cost.as_atto().try_into().unwrap_or(u64::MAX),
Ordering::Relaxed,
);
peer.payments_received.fetch_add(1, Ordering::Relaxed);
}
}
Ok(())
})
.collect();
let (
received_payment_count,
empty_earned_nodes,
min_earned,
max_earned,
min_store_cost,
max_store_cost,
) = peers
.par_iter()
.map(|peer| {
let cost =
calculate_cost_for_records(peer.records_stored.load(Ordering::Relaxed));
let earned = peer.nanos_earned.load(Ordering::Relaxed);
(
peer.payments_received.load(Ordering::Relaxed),
if earned == 0 { 1 } else { 0 },
earned,
earned,
cost,
cost,
)
})
.reduce(
|| (0, 0, u64::MAX, 0, u64::MAX, 0),
|a, b| {
let (
a_received_payment_count,
a_empty_earned_nodes,
a_min_earned,
a_max_earned,
a_min_store_cost,
a_max_store_cost,
) = a;
let (
b_received_payment_count,
b_empty_earned_nodes,
b_min_earned,
b_max_earned,
b_min_store_cost,
b_max_store_cost,
) = b;
(
a_received_payment_count + b_received_payment_count,
a_empty_earned_nodes + b_empty_earned_nodes,
a_min_earned.min(b_min_earned),
a_max_earned.max(b_max_earned),
a_min_store_cost.min(b_min_store_cost),
a_max_store_cost.max(b_max_store_cost),
)
},
);
total_received_payment_count += num_of_chunks_per_hour;
assert_eq!(total_received_payment_count, received_payment_count);
println!("After the completion of hour {hour} with {num_of_chunks_per_hour} chunks put, there are {empty_earned_nodes} nodes which earned nothing");
println!("\t\t with storecost variation of (min {min_store_cost} - max {max_store_cost}), and earned variation of (min {min_earned} - max {max_earned})");
hour += 1;
if hour == max_hours {
let acceptable_percentage = 0.01;
let acceptable_empty_nodes =
(num_of_peers as f64 * acceptable_percentage).ceil() as usize;
assert!(
empty_earned_nodes <= acceptable_empty_nodes,
"More than {acceptable_percentage}% of nodes ({acceptable_empty_nodes}) still not earning: {empty_earned_nodes}"
);
assert!(
(max_store_cost / min_store_cost) < 1000000,
"store cost is not 'balanced', expected ratio max/min to be < 1000000, but was {}",
max_store_cost / min_store_cost
);
assert!(
(max_earned / min_earned) < 500000000,
"earning distribution is not balanced, expected to be < 500000000, but was {}",
max_earned / min_earned
);
break;
}
}
}
fn pick_cheapest_payee(
peers: &[PeerStats],
close_group: &[usize],
) -> eyre::Result<(usize, AttoTokens)> {
let mut costs_vec = Vec::with_capacity(close_group.len());
let mut address_to_index = BTreeMap::new();
for &i in close_group {
let peer = &peers[i];
address_to_index.insert(peer.address.clone(), i);
let close_records_stored = peer.records_stored.load(Ordering::Relaxed);
let cost = AttoTokens::from(calculate_cost_for_records(close_records_stored));
let quote = PaymentQuote {
content: XorName::default(), cost,
timestamp: std::time::SystemTime::now(),
quoting_metrics: QuotingMetrics {
close_records_stored: peer.records_stored.load(Ordering::Relaxed),
max_records: MAX_RECORDS_COUNT,
received_payment_count: 1, live_time: 0, },
bad_nodes: vec![],
pub_key: bls::SecretKey::random().public_key().to_bytes().to_vec(),
signature: vec![],
rewards_address: peer.rewards_addr, };
costs_vec.push((peer.address.clone(), peer.rewards_addr, quote));
}
costs_vec.sort_by(|(a_addr, _, _), (b_addr, _, _)| a_addr.cmp(b_addr));
let Ok((recip_id, _pk, q)) = get_fees_from_store_cost_responses(costs_vec) else {
bail!("Failed to get fees from store cost responses")
};
let Some(index) = address_to_index
.get(&NetworkAddress::from_peer(recip_id))
.copied()
else {
bail!("Cannot find the index for the cheapest payee");
};
Ok((index, q.cost))
}
}