use std::borrow::Cow;
use std::sync::atomic::{AtomicBool, Ordering};
static ANNOTATION_ENABLED: AtomicBool = AtomicBool::new(false);
pub fn set_annotation_enabled(enabled: bool) {
ANNOTATION_ENABLED.store(enabled, Ordering::Relaxed);
}
pub fn annotation_enabled() -> bool {
ANNOTATION_ENABLED.load(Ordering::Relaxed)
}
#[cfg(feature = "tracing-context")]
const SAMPLED: u8 = 0x01;
#[cfg(feature = "tracing-context")]
pub fn current_traceparent() -> Option<String> {
crate::context::TracingContext::current()
.filter(|ctx| ctx.trace_flags & SAMPLED != 0)
.map(|ctx| ctx.traceparent)
}
#[cfg(not(feature = "tracing-context"))]
pub fn current_traceparent() -> Option<String> {
None
}
pub fn annotate_sql(sql: &str) -> Cow<'_, str> {
if !annotation_enabled() {
return Cow::Borrowed(sql);
}
match current_traceparent() {
Some(traceparent) => Cow::Owned(format!("{sql} /*traceparent='{traceparent}'*/")),
None => Cow::Borrowed(sql),
}
}
#[cfg(all(test, feature = "tracing-context"))]
mod tests {
use opentelemetry::trace::TracerProvider as _;
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
use super::*;
static ANNOTATION_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn annotation_lock() -> std::sync::MutexGuard<'static, ()> {
ANNOTATION_LOCK.lock().unwrap_or_else(|e| e.into_inner())
}
struct AnnotationEnabled(std::sync::MutexGuard<'static, ()>);
impl AnnotationEnabled {
fn for_test() -> Self {
let guard = annotation_lock();
set_annotation_enabled(true);
Self(guard)
}
}
impl Drop for AnnotationEnabled {
fn drop(&mut self) {
set_annotation_enabled(false);
}
}
fn subscriber_with_sampler(
sampler: opentelemetry_sdk::trace::Sampler,
) -> tracing::subscriber::DefaultGuard {
let provider = opentelemetry_sdk::trace::SdkTracerProvider::builder()
.with_sampler(sampler)
.build();
let tracer = provider.tracer("sql-commenter-test");
tracing_subscriber::registry()
.with(tracing_opentelemetry::layer().with_tracer(tracer))
.set_default()
}
#[test]
fn no_span_context_is_noop() {
assert!(current_traceparent().is_none());
assert_eq!(annotate_sql("SELECT 1"), "SELECT 1");
}
#[test]
fn annotates_within_sampled_span() {
let _enabled = AnnotationEnabled::for_test();
let _subscriber = subscriber_with_sampler(opentelemetry_sdk::trace::Sampler::AlwaysOn);
let span = tracing::info_span!("test_span");
let _guard = span.enter();
let tp = current_traceparent().expect("should have traceparent in span");
let parts: Vec<&str> = tp.split('-').collect();
assert_eq!(parts.len(), 4);
assert_eq!(parts[0], "00");
assert_eq!(parts[1].len(), 32);
assert_eq!(parts[2].len(), 16);
assert_eq!(parts[3], "01");
assert_eq!(
annotate_sql("SELECT 1"),
format!("SELECT 1 /*traceparent='{tp}'*/")
);
}
#[test]
fn unsampled_span_is_not_annotated() {
let _enabled = AnnotationEnabled::for_test();
let _subscriber = subscriber_with_sampler(opentelemetry_sdk::trace::Sampler::AlwaysOff);
let span = tracing::info_span!("test_span");
let _guard = span.enter();
assert!(current_traceparent().is_none());
assert_eq!(annotate_sql("SELECT 1"), "SELECT 1");
}
#[test]
fn annotation_disabled_by_default_and_toggleable() {
let _lock = annotation_lock();
let _subscriber = subscriber_with_sampler(opentelemetry_sdk::trace::Sampler::AlwaysOn);
let span = tracing::info_span!("test_span");
let _guard = span.enter();
assert_eq!(annotate_sql("SELECT 1"), "SELECT 1");
set_annotation_enabled(true);
assert!(annotate_sql("SELECT 1").contains("/*traceparent='"));
set_annotation_enabled(false);
}
}