use std::collections::BTreeMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use serde::Serialize;
use tokio::sync::{watch, Notify};
use tokio::time::{interval, Duration};
use crate::devices::gpu::GpuManager;
use crate::power_streaming::{unix_timestamp_ms, PowerBroadcast, PowerBroadcasts, PowerPoller};
#[derive(Clone, Debug, Default, Serialize)]
pub struct GpuPowerSnapshot {
pub timestamp_ms: u64,
pub power_mw: BTreeMap<usize, u32>,
}
#[derive(Clone, Debug, Default, Serialize)]
pub struct GpuPowerSample {
pub timestamp_ms: u64,
pub gpu_id: usize,
pub power_mw: u32,
}
pub type GpuPowerBroadcast = PowerBroadcast<GpuPowerSample>;
pub type GpuPowerBroadcasts = PowerBroadcasts<GpuPowerSample>;
pub type GpuPowerPoller = PowerPoller<GpuPowerSample>;
pub fn start_gpu_poller<T: GpuManager + Send + 'static>(
gpus: Vec<(usize, T)>,
poll_hz: u32,
) -> GpuPowerBroadcasts {
let mut broadcasts = BTreeMap::new();
for (gpu_id, gpu) in gpus {
let poller = PowerPoller::start(move |tx, subscriber_count, wake| {
gpu_power_poll_task(gpu_id, gpu, tx, poll_hz, subscriber_count, wake)
});
broadcasts.insert(gpu_id, poller.broadcast());
}
PowerBroadcasts::new(broadcasts)
}
async fn gpu_power_poll_task<T: GpuManager>(
gpu_id: usize,
mut gpu: T,
tx: watch::Sender<GpuPowerSample>,
poll_hz: u32,
subscriber_count: Arc<AtomicUsize>,
wake: Arc<Notify>,
) {
let period_us = 1_000_000u64 / poll_hz.max(1) as u64;
let mut last_power: Option<u32> = None;
tracing::info!(
"GPU power poller ready for GPU {} at {} Hz when subscribers are present",
gpu_id,
poll_hz
);
loop {
while subscriber_count.load(Ordering::Relaxed) == 0 {
wake.notified().await;
}
tracing::info!("GPU power poller starting for GPU {}", gpu_id);
let mut tick = interval(Duration::from_micros(period_us));
while subscriber_count.load(Ordering::Relaxed) > 0 {
tick.tick().await;
match gpu.get_instant_power_mw() {
Ok(power_mw) => {
if last_power == Some(power_mw) {
continue;
}
last_power = Some(power_mw);
let _ = tx.send(GpuPowerSample {
timestamp_ms: unix_timestamp_ms(),
gpu_id,
power_mw,
});
}
Err(e) => {
tracing::warn!("Failed to read power for GPU {}: {}", gpu_id, e);
}
}
}
last_power = None;
tracing::info!("GPU power poller pausing for GPU {}", gpu_id);
}
}