use std::fs::{self, File, OpenOptions};
use std::io::{self, BufWriter, Read, Seek, SeekFrom, Write};
use std::path::{Path, PathBuf};
use std::time::Duration;
use super::energy::{EnergyKey, MAX_DEVICES, PowerIntegrator};
pub const RECORD_LEN: usize = 24;
pub const DEFAULT_FLUSH_INTERVAL: Duration = Duration::from_secs(60);
pub const MAX_REPLAY_RECORDS: usize = 1_000_000;
pub const WAL_MAX_BYTES: u64 = 16 * 1024 * 1024;
use crate::common::paths::{cache_dir, expand_tilde};
pub fn resolve_wal_path(configured: Option<&str>) -> Option<PathBuf> {
if let Some(s) = configured.map(str::trim).filter(|s| !s.is_empty()) {
return Some(expand_tilde(Path::new(s)));
}
cache_dir().map(|d| d.join("energy-wal.bin"))
}
pub fn replay_from_path(
path: &Path,
_integrator: &mut PowerIntegrator,
) -> io::Result<WalReplayIndex> {
let path = expand_tilde(path);
match fs::symlink_metadata(&path) {
Ok(meta) if meta.file_type().is_symlink() => {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!(
"refusing to replay energy WAL at {} — path is a symlink",
path.display()
),
));
}
Ok(_) => {}
Err(e) if e.kind() == io::ErrorKind::NotFound => {
return Ok(WalReplayIndex::default());
}
Err(e) => return Err(e),
}
let mut f = match open_secure_read(&path) {
Ok(f) => f,
Err(e) if e.kind() == io::ErrorKind::NotFound => return Ok(WalReplayIndex::default()),
Err(e) => return Err(e),
};
let size_u64 = f.metadata()?.len();
let size = usize::try_from(size_u64).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("energy WAL too large for this target (size {size_u64} bytes)"),
)
})?;
let usable = size - (size % RECORD_LEN);
let record_count = usable / RECORD_LEN;
if record_count > MAX_REPLAY_RECORDS {
tracing::warn!(
"energy WAL: {record_count} records at {} exceed MAX_REPLAY_RECORDS={MAX_REPLAY_RECORDS}; truncating replay",
path.display()
);
}
let to_read = record_count.min(MAX_REPLAY_RECORDS);
let mut index = WalReplayIndex::default();
let mut buf = [0u8; RECORD_LEN];
for _ in 0..to_read {
if let Err(e) = f.read_exact(&mut buf) {
if e.kind() == io::ErrorKind::UnexpectedEof {
break;
}
return Err(e);
}
let host_hash = u64::from_le_bytes(buf[0..8].try_into().unwrap());
let device_hash = u64::from_le_bytes(buf[8..16].try_into().unwrap());
let joules = f64::from_le_bytes(buf[16..24].try_into().unwrap());
if !joules.is_finite() || joules <= 0.0 {
continue;
}
index.accumulate(host_hash, device_hash, joules);
}
Ok(index)
}
#[derive(Clone, Debug, Default)]
pub struct WalReplayIndex {
entries: std::collections::HashMap<(u64, u64), f64>,
}
impl WalReplayIndex {
fn accumulate(&mut self, host_hash: u64, device_hash: u64, joules: f64) {
let pair = (host_hash, device_hash);
if self.entries.len() >= MAX_DEVICES && !self.entries.contains_key(&pair) {
return;
}
*self.entries.entry(pair).or_insert(0.0) += joules;
}
#[allow(dead_code)] pub fn lookup(&self, host_hash: u64, device_hash: u64) -> Option<f64> {
self.entries.get(&(host_hash, device_hash)).copied()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn seed_if_matches(&mut self, key: &EnergyKey, integrator: &mut PowerIntegrator) -> f64 {
let hash_pair = (key.host_hash(), key.device_hash());
if let Some(joules) = self.entries.remove(&hash_pair) {
integrator.seed_lifetime(key.clone(), joules);
return joules;
}
0.0
}
}
#[derive(Debug)]
pub struct WalWriter {
#[allow(dead_code)] path: PathBuf,
writer: Option<BufWriter<File>>,
}
impl WalWriter {
pub fn open(path: impl AsRef<Path>) -> io::Result<Self> {
let path = expand_tilde(path.as_ref());
if let Some(parent) = path.parent()
&& !parent.as_os_str().is_empty()
{
fs::create_dir_all(parent)?;
}
let file = open_secure_append(&path)?;
Ok(Self {
path,
writer: Some(BufWriter::new(file)),
})
}
pub fn write_record(
&mut self,
host_hash: u64,
device_hash: u64,
joules: f64,
) -> io::Result<()> {
if !joules.is_finite() || joules <= 0.0 {
return Ok(());
}
let writer = self
.writer
.as_mut()
.ok_or_else(|| io::Error::other("WAL writer already closed"))?;
let mut buf = [0u8; RECORD_LEN];
buf[0..8].copy_from_slice(&host_hash.to_le_bytes());
buf[8..16].copy_from_slice(&device_hash.to_le_bytes());
buf[16..24].copy_from_slice(&joules.to_le_bytes());
writer.write_all(&buf)
}
pub fn flush_and_fsync(&mut self) -> io::Result<()> {
let writer = match self.writer.as_mut() {
Some(w) => w,
None => return Ok(()),
};
writer.flush()?;
writer.get_ref().sync_data()?;
Ok(())
}
#[allow(dead_code)] pub fn path(&self) -> &Path {
&self.path
}
}
impl Drop for WalWriter {
fn drop(&mut self) {
if let Some(mut w) = self.writer.take() {
let _ = w.flush();
}
}
}
#[cfg(feature = "cli")]
pub struct WalFlushHandle {
pub join: tokio::task::JoinHandle<()>,
pub shutdown: tokio::sync::oneshot::Sender<()>,
}
#[cfg(feature = "cli")]
impl WalFlushHandle {
pub async fn shutdown(self) {
let _ = self.shutdown.send(());
if let Err(e) = self.join.await {
tracing::warn!("energy WAL flush task terminated abnormally: {e}");
}
}
}
#[cfg(feature = "cli")]
pub fn spawn_wal_flush_task(
shared_state: std::sync::Arc<tokio::sync::RwLock<crate::app_state::AppState>>,
wal_path: PathBuf,
flush_interval: std::time::Duration,
) -> WalFlushHandle {
let (shutdown_tx, mut shutdown_rx) = tokio::sync::oneshot::channel::<()>();
let join = tokio::spawn(async move {
let mut writer = match WalWriter::open(&wal_path) {
Ok(w) => w,
Err(e) => {
tracing::warn!(
"energy WAL: failed to open {} ({e}); counters are in-memory only",
wal_path.display()
);
return;
}
};
let mut ticker = tokio::time::interval(flush_interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
ticker.tick().await;
loop {
tokio::select! {
_ = ticker.tick() => {}
_ = &mut shutdown_rx => {
tracing::debug!("energy WAL: shutdown requested; performing final flush");
let _ = flush_cycle(writer, &wal_path, &shared_state).await;
return;
}
}
writer = flush_cycle(writer, &wal_path, &shared_state).await;
}
});
WalFlushHandle {
join,
shutdown: shutdown_tx,
}
}
#[cfg(feature = "cli")]
async fn flush_cycle(
writer: WalWriter,
wal_path: &Path,
shared_state: &std::sync::Arc<tokio::sync::RwLock<crate::app_state::AppState>>,
) -> WalWriter {
let (deltas, lifetime_snapshot) = {
let mut state = shared_state.write().await;
let deltas = state.energy.integrator_mut().drain_wal_deltas();
let lifetime: Vec<(EnergyKey, f64)> = state
.energy
.integrator()
.iter_stats()
.map(|s| (s.key.clone(), s.lifetime_joules))
.collect();
(deltas, lifetime)
};
let wal_path_owned = wal_path.to_path_buf();
let result = tokio::task::spawn_blocking(move || {
let mut w = writer;
for (key, joules) in deltas {
if let Err(e) = w.write_record(key.host_hash(), key.device_hash(), joules) {
tracing::warn!("energy WAL: write failed: {e}");
}
}
let fsync_result = w.flush_and_fsync();
let should_compact = match fs::metadata(&wal_path_owned) {
Ok(meta) => meta.len() > WAL_MAX_BYTES,
Err(_) => false,
};
let (w_out, compact_result) = if should_compact {
match compact_wal(w, &wal_path_owned, &lifetime_snapshot) {
Ok(new_w) => (new_w, Ok(())),
Err((old_w, e)) => (old_w, Err(e)),
}
} else {
(w, Ok(()))
};
(w_out, fsync_result, compact_result)
})
.await;
match result {
Ok((w, fsync_result, compact_result)) => {
if let Err(e) = fsync_result {
tracing::warn!("energy WAL: fsync failed: {e}");
}
if let Err(e) = compact_result {
tracing::warn!("energy WAL: compaction failed: {e}");
}
w
}
Err(e) => {
tracing::error!("energy WAL: blocking flush task panicked: {e}");
match WalWriter::open(wal_path) {
Ok(w) => w,
Err(open_err) => {
tracing::error!(
"energy WAL: reopen after panic failed: {open_err}; subsequent flushes are no-ops until restart"
);
WalWriter {
path: wal_path.to_path_buf(),
writer: None,
}
}
}
}
}
}
fn compact_wal(
old_writer: WalWriter,
wal_path: &Path,
lifetime_snapshot: &[(EnergyKey, f64)],
) -> Result<WalWriter, (WalWriter, io::Error)> {
let resolved = expand_tilde(wal_path);
let tmp_path = {
let mut tmp = resolved.clone();
let fname = resolved
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("energy-wal.bin");
tmp.set_file_name(format!("{fname}.tmp"));
tmp
};
if let Ok(meta) = fs::symlink_metadata(&tmp_path)
&& !meta.file_type().is_symlink()
{
let _ = fs::remove_file(&tmp_path);
}
let write_and_rename = || -> io::Result<()> {
let mut tmp = WalWriter::open(&tmp_path)?;
for (key, joules) in lifetime_snapshot {
tmp.write_record(key.host_hash(), key.device_hash(), *joules)?;
}
tmp.flush_and_fsync()?;
drop(tmp);
fs::rename(&tmp_path, &resolved)?;
Ok(())
};
drop(old_writer);
if let Err(e) = write_and_rename() {
let _ = fs::remove_file(&tmp_path);
let recovered = WalWriter::open(wal_path).unwrap_or_else(|reopen_err| {
tracing::error!(
"energy WAL: reopen after compaction failure also failed: {reopen_err}"
);
WalWriter {
path: wal_path.to_path_buf(),
writer: None,
}
});
return Err((recovered, e));
}
match WalWriter::open(wal_path) {
Ok(w) => Ok(w),
Err(e) => {
tracing::error!("energy WAL: post-compaction reopen failed: {e}");
Err((
WalWriter {
path: wal_path.to_path_buf(),
writer: None,
},
e,
))
}
}
}
fn open_secure_append(path: &Path) -> io::Result<File> {
match fs::symlink_metadata(path) {
Ok(meta) if meta.file_type().is_symlink() => {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!(
"refusing to open energy WAL at {} — path is a symlink",
path.display()
),
));
}
_ => {}
}
let mut file = {
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt;
OpenOptions::new()
.create(true)
.append(true)
.read(true)
.custom_flags(libc::O_NOFOLLOW)
.mode(0o600)
.open(path)?
}
#[cfg(windows)]
{
use std::os::windows::fs::OpenOptionsExt;
OpenOptions::new()
.create(true)
.append(true)
.read(true)
.share_mode(0)
.open(path)?
}
#[cfg(not(any(unix, windows)))]
{
OpenOptions::new()
.create(true)
.append(true)
.read(true)
.open(path)?
}
};
file.seek(SeekFrom::End(0))?;
Ok(file)
}
fn open_secure_read(path: &Path) -> io::Result<File> {
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt;
OpenOptions::new()
.read(true)
.custom_flags(libc::O_NOFOLLOW)
.open(path)
}
#[cfg(windows)]
{
use std::os::windows::fs::OpenOptionsExt;
OpenOptions::new().read(true).share_mode(0).open(path)
}
#[cfg(not(any(unix, windows)))]
{
OpenOptions::new().read(true).open(path)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use tempfile::tempdir;
#[test]
fn write_then_replay_round_trips() {
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
writer.write_record(1, 2, 100.0).unwrap();
writer.write_record(1, 2, 50.0).unwrap();
writer.write_record(3, 4, 200.0).unwrap();
writer.flush_and_fsync().unwrap();
}
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.lookup(1, 2), Some(150.0));
assert_eq!(index.lookup(3, 4), Some(200.0));
assert_eq!(index.len(), 2);
}
#[test]
fn seed_if_matches_migrates_replay_into_live_key() {
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
let live_key = EnergyKey::gpu("host-a", "uuid-0");
let host_hash = live_key.host_hash();
let device_hash = live_key.device_hash();
{
let mut writer = WalWriter::open(&path).unwrap();
writer
.write_record(host_hash, device_hash, 5_000.0)
.unwrap();
writer.flush_and_fsync().unwrap();
}
let mut integ = PowerIntegrator::default();
let mut index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.len(), 1);
assert_eq!(integ.lifetime_joules(&live_key), 0.0);
let seeded = index.seed_if_matches(&live_key, &mut integ);
assert_eq!(seeded, 5_000.0);
assert_eq!(integ.lifetime_joules(&live_key), 5_000.0);
assert_eq!(index.len(), 0);
let seeded2 = index.seed_if_matches(&live_key, &mut integ);
assert_eq!(seeded2, 0.0);
assert_eq!(integ.lifetime_joules(&live_key), 5_000.0);
}
#[test]
fn missing_wal_returns_empty_index() {
let dir = tempdir().unwrap();
let path = dir.path().join("does-not-exist.bin");
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert!(index.is_empty());
}
#[test]
fn torn_final_record_is_discarded() {
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
writer.write_record(1, 2, 100.0).unwrap();
writer.write_record(3, 4, 200.0).unwrap();
writer.flush_and_fsync().unwrap();
}
let metadata = fs::metadata(&path).unwrap();
assert_eq!(metadata.len(), (RECORD_LEN * 2) as u64);
let truncated = (RECORD_LEN + 12) as u64;
let f = OpenOptions::new().write(true).open(&path).unwrap();
f.set_len(truncated).unwrap();
drop(f);
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.len(), 1);
assert_eq!(index.lookup(1, 2), Some(100.0));
assert_eq!(index.lookup(3, 4), None);
}
#[cfg(unix)]
#[test]
fn wal_file_is_mode_0o600() {
use std::os::unix::fs::PermissionsExt;
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
writer.write_record(1, 2, 10.0).unwrap();
writer.flush_and_fsync().unwrap();
}
let mode = fs::metadata(&path).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600, "WAL file must be 0o600, got {mode:o}");
}
#[cfg(unix)]
#[test]
fn wal_refuses_symlink_path() {
use std::os::unix::fs::symlink;
let dir = tempdir().unwrap();
let target = dir.path().join("actual-target");
let link = dir.path().join("energy-wal.bin");
fs::write(&target, b"existing").unwrap();
symlink(&target, &link).unwrap();
let err = WalWriter::open(&link).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::PermissionDenied);
}
#[test]
fn non_positive_records_are_ignored_on_write_and_replay() {
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
writer.write_record(1, 2, 100.0).unwrap();
writer.write_record(1, 2, 0.0).unwrap();
writer.write_record(1, 2, f64::NAN).unwrap();
writer.write_record(1, 2, -5.0).unwrap();
writer.flush_and_fsync().unwrap();
}
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.lookup(1, 2), Some(100.0));
}
#[cfg(unix)]
#[test]
fn wal_replay_refuses_symlink_path() {
use std::os::unix::fs::symlink;
let dir = tempdir().unwrap();
let target = dir.path().join("actual-target");
let link = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&target).unwrap();
writer.write_record(1, 2, 100.0).unwrap();
writer.flush_and_fsync().unwrap();
}
symlink(&target, &link).unwrap();
let mut integ = PowerIntegrator::default();
let err = replay_from_path(&link, &mut integ).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::PermissionDenied);
}
#[test]
fn replay_truncates_to_max_records() {
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
writer.write_record(1, 2, 10.0).unwrap();
writer.flush_and_fsync().unwrap();
}
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.lookup(1, 2), Some(10.0));
const _: () = assert!(
MAX_REPLAY_RECORDS >= 1,
"MAX_REPLAY_RECORDS must be a real cap"
);
}
#[test]
fn wal_replay_drops_excess_device_cardinality() {
use crate::metrics::energy::MAX_DEVICES;
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
for i in 0..(MAX_DEVICES as u64 + 50) {
writer.write_record(i, i + 1, 1.0).unwrap();
}
writer.flush_and_fsync().unwrap();
}
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.len(), MAX_DEVICES);
}
#[test]
fn compaction_rewrites_wal_under_threshold() {
let dir = tempdir().unwrap();
let path = dir.path().join("energy-wal.bin");
{
let mut writer = WalWriter::open(&path).unwrap();
writer.write_record(1, 2, 50.0).unwrap();
writer.write_record(1, 2, 25.0).unwrap(); writer.flush_and_fsync().unwrap();
}
let writer = WalWriter::open(&path).unwrap();
let live_key = EnergyKey::gpu("host-a", "uuid-0");
let snapshot = vec![(live_key.clone(), 5_000.0)];
let new_writer = compact_wal(writer, &path, &snapshot).expect("compaction succeeds");
drop(new_writer);
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&path, &mut integ).unwrap();
assert_eq!(index.len(), 1);
assert_eq!(
index.lookup(live_key.host_hash(), live_key.device_hash()),
Some(5_000.0)
);
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mode = fs::metadata(&path).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600, "compacted WAL must be 0o600, got {mode:o}");
}
}
#[test]
fn expand_tilde_replaces_home_prefix() {
let Some(home) = dirs::home_dir() else {
return;
};
let expanded = expand_tilde(Path::new("~/.cache/all-smi/energy-wal.bin"));
assert_eq!(expanded, home.join(".cache/all-smi/energy-wal.bin"));
let unchanged = expand_tilde(Path::new("/absolute/path"));
assert_eq!(unchanged, PathBuf::from("/absolute/path"));
}
#[cfg(feature = "cli")]
#[tokio::test]
async fn wal_flush_handle_shutdown_persists_pending_deltas() {
use crate::app_state::AppState;
use crate::metrics::energy::{EnergyKey, PowerIntegrator};
use std::sync::Arc;
use tokio::sync::RwLock;
let dir = tempfile::tempdir().unwrap();
let wal_path = dir.path().join("shutdown-test.bin");
let state = Arc::new(RwLock::new(AppState::new()));
{
let mut s = state.write().await;
let key = EnergyKey::gpu("test-host", "uuid-shutdown");
let origin = std::time::Instant::now();
s.energy
.integrator_mut()
.record_sample(key.clone(), origin, 300.0);
s.energy.integrator_mut().record_sample(
key.clone(),
origin + std::time::Duration::from_secs(10),
300.0,
);
}
let handle = crate::metrics::energy_wal::spawn_wal_flush_task(
state.clone(),
wal_path.clone(),
std::time::Duration::from_secs(3600),
);
handle.shutdown().await;
let mut integ = PowerIntegrator::default();
let index = replay_from_path(&wal_path, &mut integ).unwrap();
assert!(
!index.is_empty(),
"shutdown must flush pending deltas to disk; WAL index is empty"
);
}
}