use std::collections::HashMap;
use std::convert::Infallible;
use axum::Extension;
use axum::extract::{RawQuery, State};
use axum::http::StatusCode;
use axum::response::sse::{Event, KeepAlive, Sse};
use axum_extra::extract::Query as QueryExtra;
use futures::stream::Stream;
use hex;
use serde::{Deserialize, Serialize};
use crate::ingest::aggregate::{
self, AggContext, AggSnapshot, Coverage, FACETS, FacetResult, FlamegraphAccum,
PollDurationBucket, SampleFilter, Scope,
};
use crate::ingest::refine::{self, FoldErrors, Folded, RefineOpts, Resolved};
use crate::server::AppState;
use crate::server::credentials::MaybeCreds;
use crate::server::fold_stream;
use crate::server::metrics::OperationMetrics;
#[derive(Deserialize)]
pub struct FlamegraphParams {
pub service: Option<String>,
pub from: Option<String>,
pub to: Option<String>,
#[serde(default)]
pub host: Vec<String>,
pub start_ns: Option<i64>,
pub end_ns: Option<i64>,
pub max_files: Option<usize>,
pub bucket: Option<String>,
pub prefix: Option<String>,
pub aws_region: Option<String>,
pub thread_class: Option<String>,
pub source: Option<String>,
pub phase: Option<String>,
pub spawn_location: Option<String>,
pub min_poll_ns: Option<i64>,
pub max_poll_ns: Option<i64>,
pub span_type_uid: Option<String>,
pub min_span_ns: Option<i64>,
pub max_span_ns: Option<i64>,
}
#[derive(Serialize)]
pub struct FlamegraphResponse {
pub tree: FlamegraphNode,
pub total_samples: usize,
#[serde(skip_serializing_if = "Option::is_none")]
pub coverage: Option<Coverage>,
pub metadata: FlamegraphMetadata,
}
#[derive(Serialize, Clone)]
pub struct FlamegraphNode {
pub name: String,
pub count: u64,
#[serde(rename = "self")]
pub self_count: u64,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub children: Vec<FlamegraphNode>,
}
#[derive(Serialize)]
pub struct FlamegraphMetadata {
pub service: Option<String>,
pub hosts: usize,
pub time_range: Option<String>,
pub min_timestamp_ns: Option<i64>,
pub max_timestamp_ns: Option<i64>,
pub facets: Vec<FacetResult>,
pub poll_duration_histogram: Vec<PollDurationBucket>,
pub scope: ScopeEcho,
}
#[derive(Serialize)]
pub struct ScopeEcho {
pub service: Option<String>,
pub hosts: Vec<String>,
pub start_ns: Option<i64>,
pub end_ns: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_poll_ns: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_poll_ns: Option<i64>,
pub filters: HashMap<String, String>,
}
fn build_flamegraph_tree(
stack_counts: &[(Vec<u8>, u64)],
stacks_dict: &HashMap<Vec<u8>, Vec<String>>,
) -> FlamegraphNode {
let mut root = TrieNode::new("(all)".to_string());
for (stack_id, count) in stack_counts {
let frames = match stacks_dict.get(stack_id) {
Some(f) => f,
None => continue,
};
root.count += count;
let mut node = &mut root;
for frame in frames.iter().rev() {
node = node.get_or_insert_child(frame.clone());
node.count += count;
}
node.self_count += count;
}
root.into_response()
}
struct TrieNode {
name: String,
count: u64,
self_count: u64,
children: HashMap<String, TrieNode>,
}
impl TrieNode {
fn new(name: String) -> Self {
Self {
name,
count: 0,
self_count: 0,
children: HashMap::new(),
}
}
fn get_or_insert_child(&mut self, name: String) -> &mut TrieNode {
self.children
.entry(name)
.or_insert_with_key(|name| TrieNode::new(name.clone()))
}
fn into_response(self) -> FlamegraphNode {
let mut children: Vec<_> = self
.children
.into_values()
.map(|c| c.into_response())
.collect();
children.sort_by_key(|c| std::cmp::Reverse(c.count));
FlamegraphNode {
name: self.name,
count: self.count,
self_count: self.self_count,
children,
}
}
}
pub async fn get_flamegraph(
State(state): State<AppState>,
creds: MaybeCreds,
QueryExtra(params): QueryExtra<FlamegraphParams>,
RawQuery(raw_query): RawQuery,
) -> Result<
(
Extension<OperationMetrics>,
Sse<impl Stream<Item = Result<Event, Infallible>>>,
),
(StatusCode, String),
> {
if raw_query.as_deref().is_some_and(|query| {
query.split('&').any(|part| {
let (raw_key, _) = part.split_once('=').unwrap_or((part, ""));
let decoded_key = urlencoding::decode(raw_key).unwrap_or_default();
let key = decoded_key.as_ref();
let value_part = part.split_once('=').map(|(_, v)| v);
let value_empty = value_part.is_none() || value_part == Some("");
(key == "span_type_uid" || key == "phase") && value_empty
})
}) {
return Err((
StatusCode::BAD_REQUEST,
"span_type_uid and phase must not be empty".to_string(),
));
}
if let Some(ref hex_str) = params.span_type_uid {
if hex_str.is_empty() {
return Err((
StatusCode::BAD_REQUEST,
"span_type_uid must not be empty".to_string(),
));
}
match hex::decode(hex_str) {
Ok(bytes) if bytes.len() == 16 => {} _ => {
return Err((
StatusCode::BAD_REQUEST,
format!(
"invalid span_type_uid: must be 32 hex chars (16 bytes), got {hex_str:?}"
),
));
}
}
}
if let (Some(min), Some(max)) = (params.min_span_ns, params.max_span_ns) {
if min < 0 || max < 0 {
return Err((
StatusCode::BAD_REQUEST,
format!(
"span duration bounds must be non-negative: min_span_ns={min}, max_span_ns={max}"
),
));
}
if min > max {
return Err((
StatusCode::BAD_REQUEST,
format!("inverted span duration bounds: min_span_ns={min} > max_span_ns={max}"),
));
}
} else if params.min_span_ns.is_some_and(|v| v < 0) || params.max_span_ns.is_some_and(|v| v < 0)
{
return Err((
StatusCode::BAD_REQUEST,
"span duration bounds must be non-negative".to_string(),
));
}
if (params.min_span_ns.is_some() || params.max_span_ns.is_some())
&& params.span_type_uid.is_none()
{
return Err((
StatusCode::BAD_REQUEST,
"min_span_ns/max_span_ns require span_type_uid".to_string(),
));
}
if let Some(ref phase) = params.phase
&& !phase.is_empty()
&& phase != "on_cpu"
&& phase != "blocking"
{
return Err((
StatusCode::BAD_REQUEST,
format!("invalid phase: must be 'on_cpu' or 'blocking', got {phase:?}"),
));
}
let Some(agg) = state
.agg_context_for(
params.bucket.as_deref(),
params.prefix.as_deref(),
params.aws_region.as_deref(),
creds,
)
.await?
else {
return Err((
StatusCode::NOT_FOUND,
"flamegraph requires demand-driven aggregation (start with --agg or supply a bucket)"
.to_string(),
));
};
let scope = scope_from_params(¶ms);
let opts = RefineOpts {
max_files: params.max_files,
};
let Some(resolved) = refine::resolve(&agg, &scope, opts).await else {
return Err((
StatusCode::NOT_FOUND,
"no source files match this scope".to_string(),
));
};
let op = OperationMetrics::flamegraph(
resolved.files_matched as u32,
resolved.files_folded_in(resolved.folded()) as u32,
None,
);
let stream = flamegraph_stream(agg, resolved, ¶ms, state.fold_limits.clone());
Ok((
Extension(op),
Sse::new(stream).keep_alive(KeepAlive::default()),
))
}
fn scope_from_params(params: &FlamegraphParams) -> Scope {
Scope {
start_ns: params.start_ns,
end_ns: params.end_ns,
service: params.service.clone(),
hosts: params.host.clone(),
}
}
fn sample_filter(params: &FlamegraphParams) -> SampleFilter {
let mut facets = HashMap::new();
for def in FACETS {
let value = match def.name {
"source" => {
let effective_source = match params.phase.as_deref() {
Some("on_cpu") => Some("cpu".to_string()),
Some("blocking") => Some("sched".to_string()),
_ => params.source.clone(),
};
let raw = effective_source.unwrap_or_else(|| def.default_filter.to_string());
if raw == "all" { String::new() } else { raw }
}
"thread_class" => params
.thread_class
.clone()
.unwrap_or_else(|| def.default_filter.to_string()),
"spawn_location" => params
.spawn_location
.clone()
.unwrap_or_else(|| def.default_filter.to_string()),
"host" => {
String::new()
}
_ => def.default_filter.to_string(),
};
facets.insert(def.name, value);
}
let span_type_uid = match params.span_type_uid.as_deref() {
Some(hex_str) if !hex_str.is_empty() => {
match hex::decode(hex_str) {
Ok(bytes) if bytes.len() == 16 => {
let mut uid = [0u8; 16];
uid.copy_from_slice(&bytes);
Some(uid)
}
_ => {
Some([0u8; 16])
}
}
}
_ => None,
};
SampleFilter {
start_ns: params.start_ns,
end_ns: params.end_ns,
min_poll_ns: params.min_poll_ns,
max_poll_ns: params.max_poll_ns,
facets,
span_type_uid,
min_span_ns: params.min_span_ns,
max_span_ns: params.max_span_ns,
}
}
struct StreamCtx {
filter: SampleFilter,
service: Option<String>,
hosts: Vec<String>,
from: Option<String>,
to: Option<String>,
start_ns: Option<i64>,
end_ns: Option<i64>,
min_poll_ns: Option<i64>,
max_poll_ns: Option<i64>,
}
struct FlamegraphSink {
ctx: StreamCtx,
accum: FlamegraphAccum,
}
impl fold_stream::FoldSink for FlamegraphSink {
async fn seed_batch(
&mut self,
agg: &AggContext,
full_keys: &[String],
) -> Vec<fold_stream::PartOutcome> {
let seed = aggregate::fetch_folded_sample_parts(
&*agg.output,
&agg.output_bucket,
&agg.output_prefix,
full_keys,
)
.await;
let mut outcomes = Vec::with_capacity(seed.len());
for (leaf, result) in seed {
let outcome = match result {
Ok((samples, dict)) => match self.accum.merge(samples, dict) {
Ok(()) => fold_stream::PartOutcome::Folded { leaf },
Err(e) => {
fold_stream::rate_limited_warn("flamegraph: seed merge failed", &e);
fold_stream::PartOutcome::Failed {
key: leaf,
error: format!("merge: {e}"),
}
}
},
Err(msg) => fold_stream::PartOutcome::Failed {
key: leaf,
error: msg,
},
};
outcomes.push(outcome);
}
outcomes
}
async fn fold_one(&mut self, agg: &AggContext, f: &Folded) -> fold_stream::PartOutcome {
match aggregate::fetch_sample_parts(
&*agg.output,
&agg.output_bucket,
&agg.output_prefix,
&f.full_key,
)
.await
{
Some((samples, dict)) => match self.accum.merge(samples, dict) {
Ok(()) => fold_stream::PartOutcome::Folded {
leaf: aggregate::part_leaf_of(&f.full_key),
},
Err(e) => {
fold_stream::rate_limited_warn("flamegraph: merge failed", &e);
fold_stream::PartOutcome::Failed {
key: f.raw_key.clone(),
error: format!("merge: {e}"),
}
}
},
None => {
fold_stream::PartOutcome::Failed {
key: f.raw_key.clone(),
error: "sample parts GET failed (not found)".to_string(),
}
}
}
}
fn snapshot_event(
&self,
resolved: &Resolved,
files_folded: usize,
folded_set_id: &str,
target_folded_set_id: Option<&str>,
hosts_folded: usize,
errors: &FoldErrors,
) -> Event {
let snap = self.accum.snapshot();
let coverage = fold_stream::coverage_from(
resolved,
files_folded,
folded_set_id,
target_folded_set_id,
hosts_folded,
errors,
snap.total_samples,
);
let resp = build_response(&self.ctx, &snap, coverage);
Event::default().json_data(&resp).unwrap_or_else(|e| {
fold_stream::rate_limited_warn(
"flamegraph: event serialize failed",
&anyhow::anyhow!(e),
);
Event::default().comment("serialize error")
})
}
}
fn flamegraph_stream(
agg: AggContext,
resolved: Resolved,
params: &FlamegraphParams,
limits: aggregate::FoldLimits,
) -> impl Stream<Item = Result<Event, Infallible>> + use<> {
let ctx = StreamCtx {
filter: sample_filter(params),
service: params.service.clone(),
hosts: params.host.clone(),
from: params.from.clone(),
to: params.to.clone(),
start_ns: params.start_ns,
end_ns: params.end_ns,
min_poll_ns: params.min_poll_ns,
max_poll_ns: params.max_poll_ns,
};
let accum = FlamegraphAccum::new(ctx.filter.clone());
fold_stream::drive(agg, resolved, limits, FlamegraphSink { ctx, accum })
}
fn build_response(ctx: &StreamCtx, snap: &AggSnapshot, coverage: Coverage) -> FlamegraphResponse {
let tree = build_flamegraph_tree(&snap.stack_counts, snap.stacks_dict);
let filters: HashMap<String, String> = ctx
.filter
.facets
.iter()
.map(|(k, v)| (k.to_string(), v.clone()))
.collect();
FlamegraphResponse {
tree,
total_samples: snap.total_samples,
coverage: Some(coverage),
metadata: FlamegraphMetadata {
service: ctx.service.clone(),
hosts: snap.hosts,
time_range: match (&ctx.from, &ctx.to) {
(Some(f), Some(t)) => Some(format!("{f}–{t}")),
_ => None,
},
min_timestamp_ns: snap.min_ts,
max_timestamp_ns: snap.max_ts,
facets: snap.facets.clone(),
poll_duration_histogram: snap.poll_duration_histogram.clone(),
scope: ScopeEcho {
service: ctx.service.clone(),
hosts: ctx.hosts.clone(),
start_ns: ctx.start_ns,
end_ns: ctx.end_ns,
min_poll_ns: ctx.min_poll_ns,
max_poll_ns: ctx.max_poll_ns,
filters,
},
},
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_flamegraph_tree() {
let mut stacks_dict = HashMap::new();
let stack_a = vec![0u8; 16];
stacks_dict.insert(
stack_a.clone(),
vec!["bar".to_string(), "foo".to_string(), "main".to_string()],
);
let mut stack_b = vec![0u8; 16];
stack_b[0] = 1;
stacks_dict.insert(
stack_b.clone(),
vec!["baz".to_string(), "foo".to_string(), "main".to_string()],
);
let stack_counts = vec![(stack_a, 10), (stack_b, 5)];
let tree = build_flamegraph_tree(&stack_counts, &stacks_dict);
assert_eq!(tree.name, "(all)");
assert_eq!(tree.count, 15);
assert_eq!(tree.children.len(), 1); let main_node = &tree.children[0];
assert_eq!(main_node.name, "main");
assert_eq!(main_node.count, 15);
let foo_node = &main_node.children[0];
assert_eq!(foo_node.name, "foo");
assert_eq!(foo_node.count, 15);
assert_eq!(foo_node.children.len(), 2); }
}