use std::cell::RefCell;
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
use std::time::Instant;
use lora_compiler::physical::PhysicalNodeId;
use crate::errors::ExecResult;
use crate::pull::RowSource;
use crate::value::Row;
#[derive(Debug, Clone, Default)]
pub struct OperatorProfile {
pub rows: u64,
pub elapsed_ns: u64,
pub next_calls: u64,
}
#[derive(Debug, Default)]
pub struct MetricsCollector {
inner: Mutex<BTreeMap<PhysicalNodeId, OperatorProfile>>,
}
impl MetricsCollector {
pub fn new() -> Self {
Self::default()
}
pub fn record(&self, op: PhysicalNodeId, elapsed_ns: u64, produced_row: bool) {
if let Ok(mut map) = self.inner.lock() {
let entry = map.entry(op).or_default();
entry.next_calls += 1;
entry.elapsed_ns = entry.elapsed_ns.saturating_add(elapsed_ns);
if produced_row {
entry.rows += 1;
}
}
}
pub fn snapshot(&self) -> BTreeMap<PhysicalNodeId, OperatorProfile> {
self.inner.lock().map(|m| m.clone()).unwrap_or_default()
}
}
thread_local! {
static CURRENT: RefCell<Option<Arc<MetricsCollector>>> = const { RefCell::new(None) };
}
pub struct CollectorGuard {
_private: (),
}
impl CollectorGuard {
pub fn install(collector: Arc<MetricsCollector>) -> Self {
CURRENT.with(|cell| {
*cell.borrow_mut() = Some(collector);
});
Self { _private: () }
}
}
impl Drop for CollectorGuard {
fn drop(&mut self) {
CURRENT.with(|cell| {
*cell.borrow_mut() = None;
});
}
}
pub(crate) fn wrap_metered<'a>(
op_id: PhysicalNodeId,
inner: Box<dyn RowSource + 'a>,
) -> Box<dyn RowSource + 'a> {
let collector = CURRENT.with(|cell| cell.borrow().clone());
match collector {
Some(c) => Box::new(MeteredSource {
inner,
op_id,
collector: c,
}),
None => inner,
}
}
struct MeteredSource<'a> {
inner: Box<dyn RowSource + 'a>,
op_id: PhysicalNodeId,
collector: Arc<MetricsCollector>,
}
impl<'a> RowSource for MeteredSource<'a> {
fn next_row(&mut self) -> ExecResult<Option<Row>> {
let t0 = Instant::now();
let result = self.inner.next_row();
let elapsed_ns = t0.elapsed().as_nanos().min(u128::from(u64::MAX)) as u64;
let produced = matches!(&result, Ok(Some(_)));
self.collector.record(self.op_id, elapsed_ns, produced);
result
}
}