use std::time::Duration;
use opentelemetry::KeyValue;
use opentelemetry::trace::{TraceId, TracerProvider as _};
use opentelemetry_otlp::{ExporterBuildError, WithExportConfig};
use opentelemetry_sdk::Resource;
use opentelemetry_sdk::trace::{RandomIdGenerator, Sampler, SdkTracerProvider};
use serde::Serialize;
use thiserror::Error;
use tracing::dispatcher::SetGlobalDefaultError;
use tracing_opentelemetry::OpenTelemetryLayer;
use tracing_subscriber::prelude::*;
use tracing_subscriber::{EnvFilter, Registry};
#[derive(Error, Debug)]
pub enum Error {
#[error("ExporterBuildError: {0}")]
ExporterBuildError(#[source] ExporterBuildError),
#[error("SetGlobalDefaultError: {0}")]
SetGlobalDefaultError(#[source] SetGlobalDefaultError),
}
pub fn get_trace_id() -> TraceId {
use opentelemetry::trace::TraceContextExt as _; use tracing_opentelemetry::OpenTelemetrySpanExt as _;
tracing::Span::current()
.context()
.span()
.span_context()
.trace_id()
}
#[derive(clap::ValueEnum, Clone, Debug, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum LogFormat {
Json,
Text,
}
pub async fn init(
log_filter: &str,
log_format: LogFormat,
tracing_url: Option<&str>,
trace_ratio: f64,
) -> Result<(), Error> {
let logger = match log_format {
LogFormat::Json => tracing_subscriber::fmt::layer().json().compact().boxed(),
LogFormat::Text => tracing_subscriber::fmt::layer().compact().boxed(),
};
let filter = EnvFilter::new(log_filter).add_directive("kanidm_client=error".parse().unwrap());
let collector = Registry::default().with(logger).with(filter);
if let Some(url) = tracing_url {
let exporter = opentelemetry_otlp::SpanExporter::builder()
.with_http()
.with_endpoint(url)
.with_timeout(Duration::from_secs(3))
.build()
.map_err(Error::ExporterBuildError)?;
let provider = SdkTracerProvider::builder()
.with_sampler(Sampler::TraceIdRatioBased(trace_ratio))
.with_id_generator(RandomIdGenerator::default())
.with_max_events_per_span(64)
.with_max_attributes_per_span(16)
.with_max_events_per_span(16)
.with_resource(
Resource::builder()
.with_service_name("kaniop")
.with_attribute(KeyValue::new("key", "value"))
.build(),
)
.with_batch_exporter(exporter)
.build();
let tracer = provider.tracer("opentelemetry-otlp");
let telemetry = OpenTelemetryLayer::new(tracer);
tracing::subscriber::set_global_default(collector.with(telemetry))
.map_err(Error::SetGlobalDefaultError)
} else {
tracing::subscriber::set_global_default(collector).map_err(Error::SetGlobalDefaultError)
}
}
#[cfg(all(test, feature = "integration-test"))]
mod test {
#[tokio::test]
async fn integration_get_trace_id_returns_valid_traces() {
use super::*;
let opentelemetry_endpoint_url = std::env::var("OPENTELEMETRY_ENDPOINT_URL").ok();
super::init(
"info",
LogFormat::Text,
opentelemetry_endpoint_url.as_deref(),
0.1,
)
.await
.unwrap();
#[tracing::instrument(name = "test_span")] fn test_trace_id() -> TraceId {
get_trace_id()
}
assert_ne!(test_trace_id(), TraceId::INVALID, "valid trace");
}
}