use std::sync::Arc;
use tracing::field::{Field, Visit};
use tracing::subscriber::Interest;
use tracing::{Event, Metadata, Subscriber};
use tracing_subscriber::layer::{Context, Filter};
use tracing_subscriber::registry::LookupSpan;
use tracing_subscriber::Layer;
use crate::lib_on::sql::{init_sql_state, send_sql_event, SqlEvent};
const TOASTY_QUERY_TARGET: &str = "toasty::query";
pub(crate) struct HotpathToastyLayer;
impl<S> Layer<S> for HotpathToastyLayer
where
S: Subscriber + for<'a> LookupSpan<'a>,
{
fn on_event(&self, event: &Event<'_>, _ctx: Context<'_, S>) {
let mut visitor = QueryVisitor::default();
event.record(&mut visitor);
let Some(sql) = visitor.statement else {
return;
};
send_sql_event(SqlEvent::Executed {
sql,
duration_nanos: visitor.duration_ns.unwrap_or(0),
timestamp_ns: crate::lib_on::current_elapsed_ns(),
source: crate::lib_on::caller_stack::current_caller(),
});
}
}
#[derive(Default)]
struct QueryVisitor {
statement: Option<Arc<str>>,
duration_ns: Option<u64>,
}
impl Visit for QueryVisitor {
fn record_str(&mut self, field: &Field, value: &str) {
if field.name() == "db.statement" {
self.statement = Some(Arc::from(value.trim()));
}
}
fn record_f64(&mut self, field: &Field, value: f64) {
if field.name() == "duration_ms" {
self.duration_ns = Some((value * 1e6) as u64);
}
}
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
if field.name() == "db.statement" && self.statement.is_none() {
let rendered = format!("{value:?}");
self.statement = Some(Arc::from(rendered.trim()));
}
}
}
struct ToastyQueryFilter;
impl<S> Filter<S> for ToastyQueryFilter {
fn enabled(&self, meta: &Metadata<'_>, _ctx: &Context<'_, S>) -> bool {
meta.target() == TOASTY_QUERY_TARGET
}
fn callsite_enabled(&self, meta: &Metadata<'_>) -> Interest {
if meta.target() == TOASTY_QUERY_TARGET {
Interest::always()
} else {
Interest::never()
}
}
}
pub fn toasty_tracing_layer<S>() -> impl Layer<S>
where
S: Subscriber + for<'a> LookupSpan<'a>,
{
init_sql_state();
HotpathToastyLayer.with_filter(ToastyQueryFilter)
}