use std::fs;
use std::io::{self, Write};
use std::path::{Path, PathBuf};
use crate::sync::topology::{SyncNodeId, SyncTopology};
use super::{CONFIG_FILE, DatabaseError};
pub const ON_DISK_FORMAT_VERSION: u32 = 1;
#[derive(Clone, Debug, serde::Deserialize, serde::Serialize)]
pub struct DatabaseConfig {
pub data_dir: PathBuf,
pub shard_count: usize,
pub distributed: Option<DistributedDatabaseConfig>,
}
#[derive(Clone, Debug, serde::Deserialize, serde::Serialize)]
pub struct DistributedDatabaseConfig {
pub local_node: SyncNodeId,
pub nodes: Vec<SyncNodeId>,
pub topology: Option<SyncTopology>,
pub sync_interval: u64,
}
#[derive(serde::Serialize)]
struct StampedConfig<'config> {
format_version: u32,
#[serde(flatten)]
config: &'config DatabaseConfig,
}
#[derive(serde::Deserialize)]
struct StoredConfig {
format_version: Option<u32>,
#[serde(flatten)]
config: DatabaseConfig,
}
pub(super) fn write_config(config: &DatabaseConfig) -> Result<(), DatabaseError> {
let stamped = StampedConfig {
format_version: ON_DISK_FORMAT_VERSION,
config,
};
let bytes = serde_json::to_vec_pretty(&stamped).map_err(|error| {
DatabaseError::ConfigWrite(io::Error::new(io::ErrorKind::InvalidData, error))
})?;
install_config_atomic(&config.data_dir.join(CONFIG_FILE), &bytes)
}
pub(super) fn ensure_uninitialised(data_dir: &Path) -> Result<(), DatabaseError> {
let config_path = data_dir.join(CONFIG_FILE);
if config_path.exists() {
read_config(data_dir)?;
return Err(DatabaseError::DataDirAlreadyInitialised { config_path });
}
Ok(())
}
pub(super) fn read_config(path: &Path) -> Result<DatabaseConfig, DatabaseError> {
let bytes = fs::read(path.join(CONFIG_FILE)).map_err(DatabaseError::ConfigRead)?;
let value: serde_json::Value = serde_json::from_slice(&bytes)
.map_err(|error| DatabaseError::ConfigParse(error.to_string()))?;
validate_legacy_ttl_cadence(&value)?;
let stored: StoredConfig = serde_json::from_value(value)
.map_err(|error| DatabaseError::ConfigParse(error.to_string()))?;
validate_format_version(stored.format_version)?;
Ok(stored.config)
}
fn validate_legacy_ttl_cadence(value: &serde_json::Value) -> Result<(), DatabaseError> {
const LEGACY_KEY: &str = "sweep_interval";
let Some(value) = value.as_object().and_then(|object| object.get(LEGACY_KEY)) else {
return Ok(());
};
match value {
serde_json::Value::Null => Ok(()),
serde_json::Value::Number(number) if number.as_u64() == Some(0) => {
Err(DatabaseError::InvalidSweepInterval)
}
serde_json::Value::Number(number) if number.as_u64().is_some() => Ok(()),
_ => Err(DatabaseError::ConfigParse(
"legacy TTL cadence must be null or an unsigned integer".to_owned(),
)),
}
}
fn validate_format_version(stamp: Option<u32>) -> Result<(), DatabaseError> {
match stamp {
None | Some(ON_DISK_FORMAT_VERSION) => Ok(()),
Some(found) if found > ON_DISK_FORMAT_VERSION => Err(DatabaseError::FormatVersionTooNew {
found,
supported: ON_DISK_FORMAT_VERSION,
}),
Some(found) => Err(DatabaseError::ConfigParse(format!(
"on-disk format_version {found} has no read arm in this binary (reads \
{ON_DISK_FORMAT_VERSION}; unstamped legacy reads as 1; 0 was never issued — \
stamping begins at 1)"
))),
}
}
fn install_config_atomic(path: &Path, bytes: &[u8]) -> Result<(), DatabaseError> {
let parent = path.parent().ok_or_else(|| {
DatabaseError::ConfigWrite(io::Error::new(
io::ErrorKind::InvalidInput,
"config path has no parent directory",
))
})?;
let mut temp_file = tempfile::Builder::new()
.prefix(".config-")
.suffix(".tmp")
.tempfile_in(parent)
.map_err(DatabaseError::ConfigWrite)?;
temp_file
.write_all(bytes)
.map_err(DatabaseError::ConfigWrite)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
temp_file
.as_file()
.set_permissions(fs::Permissions::from_mode(0o644))
.map_err(DatabaseError::ConfigWrite)?;
}
temp_file
.as_file_mut()
.sync_all()
.map_err(DatabaseError::ConfigWrite)?;
temp_file
.persist(path)
.map(drop)
.map_err(|error| DatabaseError::ConfigWrite(error.error))?;
sync_parent_dir(parent)
}
#[cfg(unix)]
fn sync_parent_dir(parent: &Path) -> Result<(), DatabaseError> {
fs::File::open(parent)
.and_then(|dir| dir.sync_all())
.map_err(DatabaseError::ConfigWrite)
}
#[cfg(not(unix))]
fn sync_parent_dir(_parent: &Path) -> Result<(), DatabaseError> {
Ok(())
}
pub(super) fn validate_database_config(config: &DatabaseConfig) -> Result<(), DatabaseError> {
validate_shard_count(config.shard_count)?;
validate_distributed_config(config)?;
Ok(())
}
const fn validate_shard_count(shard_count: usize) -> Result<(), DatabaseError> {
if shard_count == 0 {
Err(DatabaseError::InvalidShardCount)
} else {
Ok(())
}
}
fn validate_distributed_config(config: &DatabaseConfig) -> Result<(), DatabaseError> {
let Some(distributed) = &config.distributed else {
return Ok(());
};
let Some(topology) = &distributed.topology else {
return Err(DatabaseError::MissingSyncTopology);
};
if distributed.sync_interval == 0 {
return Err(DatabaseError::InvalidSyncInterval);
}
topology
.partners_for(&distributed.local_node, &distributed.nodes)
.map_err(|error| DatabaseError::SyncSchedulerError(error.to_string()))?;
Ok(())
}