#[cfg(feature = "telemetry")]
use opentelemetry::{global, trace::TracerProvider as _, KeyValue};
#[cfg(feature = "telemetry")]
use opentelemetry_otlp::{SpanExporter, WithExportConfig};
#[cfg(feature = "telemetry")]
use opentelemetry_sdk::{
trace::{RandomIdGenerator, Sampler, TracerProvider},
Resource,
};
#[cfg(feature = "telemetry")]
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt, EnvFilter};
#[cfg(feature = "telemetry")]
static TRACER_PROVIDER: std::sync::OnceLock<TracerProvider> = std::sync::OnceLock::new();
#[cfg(feature = "telemetry")]
pub fn init_tracing() -> Result<(), Box<dyn std::error::Error>> {
let otlp_endpoint = std::env::var("OTEL_EXPORTER_OTLP_ENDPOINT")
.unwrap_or_else(|_| "http://localhost:4317".to_string());
let service_name =
std::env::var("OTEL_SERVICE_NAME").unwrap_or_else(|_| "shodh-memory".to_string());
let exporter = SpanExporter::builder()
.with_tonic()
.with_endpoint(&otlp_endpoint)
.build()?;
let resource = Resource::new(vec![
KeyValue::new("service.name", service_name.clone()),
KeyValue::new("service.version", env!("CARGO_PKG_VERSION")),
]);
let tracer_provider = TracerProvider::builder()
.with_batch_exporter(exporter, opentelemetry_sdk::runtime::Tokio)
.with_sampler(Sampler::ParentBased(Box::new(Sampler::AlwaysOn)))
.with_id_generator(RandomIdGenerator::default())
.with_resource(resource)
.build();
let tracer = tracer_provider.tracer("shodh-memory");
let _ = TRACER_PROVIDER.set(tracer_provider.clone());
global::set_tracer_provider(tracer_provider);
let telemetry_layer = tracing_opentelemetry::layer().with_tracer(tracer);
let env_filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info"));
tracing_subscriber::registry()
.with(env_filter)
.with(tracing_subscriber::fmt::layer())
.with(telemetry_layer)
.init();
tracing::info!(
service_name = %service_name,
otlp_endpoint = %otlp_endpoint,
"OpenTelemetry tracing initialized"
);
Ok(())
}
#[cfg(feature = "telemetry")]
pub fn shutdown_tracing() {
tracing::info!("Shutting down OpenTelemetry tracing");
if let Some(provider) = TRACER_PROVIDER.get() {
if let Err(e) = provider.shutdown() {
tracing::error!("Error shutting down tracer provider: {:?}", e);
}
}
}
#[cfg(feature = "telemetry")]
pub mod trace_propagation {
use axum::{extract::Request, middleware::Next, response::Response};
use opentelemetry::global;
use opentelemetry::propagation::Extractor;
use tracing::Span;
use tracing_opentelemetry::OpenTelemetrySpanExt;
struct HeaderExtractor<'a> {
headers: &'a axum::http::HeaderMap,
}
impl<'a> Extractor for HeaderExtractor<'a> {
fn get(&self, key: &str) -> Option<&str> {
self.headers.get(key)?.to_str().ok()
}
fn keys(&self) -> Vec<&str> {
self.headers.keys().map(|k| k.as_str()).collect()
}
}
pub async fn propagate_trace_context(req: Request, next: Next) -> Response {
let extractor = HeaderExtractor {
headers: req.headers(),
};
let parent_cx =
global::get_text_map_propagator(|propagator| propagator.extract(&extractor));
let current_span = Span::current();
current_span.set_parent(parent_cx);
next.run(req).await
}
}
#[cfg(all(test, feature = "telemetry"))]
mod tests {
use super::*;
#[test]
fn test_tracing_init_no_panic() {
let _ = init_tracing();
}
}