tracing-calltree 0.1.2

Always-on hierarchical profiling for Rust tracing spans with rolling latency statistics.
Documentation
use crate::sample::{RollingSamples, TimingSample};
use crate::snapshot::{CallTreeSnapshot, NodeSnapshot};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use tracing::Metadata;

fn reserve_node(node_count: &AtomicUsize, max_nodes: usize) -> bool {
    let mut current = node_count.load(Ordering::Relaxed);

    loop {
        if current >= max_nodes {
            return false;
        }

        match node_count.compare_exchange_weak(
            current,
            current + 1,
            Ordering::Relaxed,
            Ordering::Relaxed,
        ) {
            Ok(_) => return true,
            Err(actual) => current = actual,
        }
    }
}

#[derive(Clone, Debug, Hash, PartialEq, Eq, PartialOrd, Ord)]
pub(crate) struct NodeKey {
    pub(crate) name: String,
    pub(crate) target: String,
    pub(crate) module_path: Option<String>,
    pub(crate) line: Option<u32>,
}

impl NodeKey {
    pub(crate) fn from_metadata(metadata: &Metadata<'_>) -> Self {
        Self {
            name: metadata.name().to_string(),
            target: metadata.target().to_string(),
            module_path: metadata.module_path().map(str::to_string),
            line: metadata.line(),
        }
    }
}

#[derive(Debug)]
pub(crate) struct CallTreeState {
    roots: Mutex<HashMap<NodeKey, Arc<CallTreeNode>>>,
    node_count: AtomicUsize,
    rejected_nodes: AtomicU64,
    dropped_samples: AtomicU64,
}

impl Default for CallTreeState {
    fn default() -> Self {
        Self {
            roots: Mutex::new(HashMap::new()),
            node_count: AtomicUsize::new(0),
            rejected_nodes: AtomicU64::new(0),
            dropped_samples: AtomicU64::new(0),
        }
    }
}

impl CallTreeState {
    pub(crate) fn get_or_create_node(
        &self,
        parent: Option<&Arc<CallTreeNode>>,
        key: NodeKey,
        depth: usize,
        max_nodes: usize,
    ) -> Option<Arc<CallTreeNode>> {
        match parent {
            Some(parent) => parent.get_or_create_child(
                key,
                depth,
                max_nodes,
                &self.node_count,
                &self.rejected_nodes,
            ),
            None => {
                let mut roots = self
                    .roots
                    .lock()
                    .unwrap_or_else(|poisoned| poisoned.into_inner());

                if let Some(existing) = roots.get(&key) {
                    return Some(existing.clone());
                }

                if !reserve_node(&self.node_count, max_nodes) {
                    self.rejected_nodes.fetch_add(1, Ordering::Relaxed);
                    return None;
                }

                let node = Arc::new(CallTreeNode::new(key.clone(), depth));
                roots.insert(key, node.clone());
                Some(node)
            }
        }
    }

    pub(crate) fn snapshot(&self) -> CallTreeSnapshot {
        let mut roots: Vec<_> = {
            let guard = self
                .roots
                .lock()
                .unwrap_or_else(|poisoned| poisoned.into_inner());
            guard.values().cloned().collect()
        };

        roots.sort_by(|left, right| left.key.cmp(&right.key));

        CallTreeSnapshot {
            roots: roots.into_iter().map(|node| node.snapshot()).collect(),
        }
    }

    pub(crate) fn node_count(&self) -> usize {
        self.node_count.load(Ordering::Relaxed)
    }

    pub(crate) fn rejected_nodes(&self) -> u64 {
        self.rejected_nodes.load(Ordering::Relaxed)
    }

    pub(crate) fn dropped_samples(&self) -> u64 {
        self.dropped_samples.load(Ordering::Relaxed)
    }

    pub(crate) fn dropped_samples_counter(&self) -> &AtomicU64 {
        &self.dropped_samples
    }
}

#[derive(Debug)]
pub(crate) struct CallTreeNode {
    pub(crate) key: NodeKey,
    depth: usize,
    samples: Mutex<RollingSamples>,
    children: Mutex<HashMap<NodeKey, Arc<CallTreeNode>>>,
}

impl CallTreeNode {
    fn new(key: NodeKey, depth: usize) -> Self {
        Self {
            key,
            depth,
            samples: Mutex::new(RollingSamples::default()),
            children: Mutex::new(HashMap::new()),
        }
    }

    pub(crate) fn depth(&self) -> usize {
        self.depth
    }

    pub(crate) fn record_sample(
        &self,
        sample: TimingSample,
        window_size: usize,
        dropped_samples: &AtomicU64,
    ) {
        let mut guard = self
            .samples
            .lock()
            .unwrap_or_else(|poisoned| poisoned.into_inner());
        guard.record(sample, window_size, dropped_samples);
    }

    fn get_or_create_child(
        &self,
        key: NodeKey,
        depth: usize,
        max_nodes: usize,
        node_count: &AtomicUsize,
        rejected_nodes: &AtomicU64,
    ) -> Option<Arc<CallTreeNode>> {
        let mut children = self
            .children
            .lock()
            .unwrap_or_else(|poisoned| poisoned.into_inner());

        if let Some(existing) = children.get(&key) {
            return Some(existing.clone());
        }

        if !reserve_node(node_count, max_nodes) {
            rejected_nodes.fetch_add(1, Ordering::Relaxed);
            return None;
        }

        let node = Arc::new(CallTreeNode::new(key.clone(), depth));
        children.insert(key, node.clone());
        Some(node)
    }

    pub(crate) fn snapshot(&self) -> NodeSnapshot {
        let sample_snapshot = {
            let guard = self
                .samples
                .lock()
                .unwrap_or_else(|poisoned| poisoned.into_inner());
            guard.snapshot()
        };

        let mut children: Vec<_> = {
            let guard = self
                .children
                .lock()
                .unwrap_or_else(|poisoned| poisoned.into_inner());
            guard.values().cloned().collect()
        };
        children.sort_by(|left, right| left.key.cmp(&right.key));

        NodeSnapshot::from_samples(
            self.key.name.clone(),
            self.key.target.clone(),
            self.key.module_path.clone(),
            self.key.line,
            sample_snapshot.total_calls,
            &sample_snapshot.values,
            children.into_iter().map(|child| child.snapshot()).collect(),
        )
    }
}

#[cfg(test)]
mod tests {
    use super::{CallTreeState, NodeKey};
    use std::sync::Arc;

    #[test]
    fn same_callsite_under_different_parents_stays_distinct() {
        let state = CallTreeState::default();

        let request = state
            .get_or_create_node(
                None,
                NodeKey {
                    name: "request".to_string(),
                    target: "app".to_string(),
                    module_path: Some("app".to_string()),
                    line: Some(1),
                },
                1,
                10,
            )
            .unwrap();
        let maintenance = state
            .get_or_create_node(
                None,
                NodeKey {
                    name: "maintenance".to_string(),
                    target: "app".to_string(),
                    module_path: Some("app".to_string()),
                    line: Some(2),
                },
                1,
                10,
            )
            .unwrap();

        let request_db = state
            .get_or_create_node(
                Some(&request),
                NodeKey {
                    name: "database".to_string(),
                    target: "app".to_string(),
                    module_path: Some("app".to_string()),
                    line: Some(3),
                },
                2,
                10,
            )
            .unwrap();
        let maintenance_db = state
            .get_or_create_node(
                Some(&maintenance),
                NodeKey {
                    name: "database".to_string(),
                    target: "app".to_string(),
                    module_path: Some("app".to_string()),
                    line: Some(3),
                },
                2,
                10,
            )
            .unwrap();

        assert!(!Arc::ptr_eq(&request_db, &maintenance_db));
        let snapshot = state.snapshot();
        assert_eq!(snapshot.roots.len(), 2);
        assert_eq!(snapshot.roots[0].children.len(), 1);
        assert_eq!(snapshot.roots[1].children.len(), 1);
        assert_eq!(snapshot.roots[0].children[0].name, "database");
        assert_eq!(snapshot.roots[1].children[0].name, "database");
    }
}