use std::future::Future;
use tracing::Instrument;
use crate::Result;
pub const METRIC_REQUESTS: &str = "canton_client_requests_total";
pub const METRIC_ERRORS: &str = "canton_client_errors_total";
pub const TRANSPORT_GRPC: &str = "grpc";
pub const TRANSPORT_JSON: &str = "json";
pub async fn instrument<T, F>(method: &'static str, transport: &'static str, fut: F) -> Result<T>
where
F: Future<Output = Result<T>>,
{
metrics::counter!(METRIC_REQUESTS, "method" => method, "transport" => transport).increment(1);
let span = tracing::info_span!("canton.rpc", method = method, transport = transport);
async move {
let result = fut.await;
let trace_id = current_trace_id().unwrap_or_default();
match &result {
Ok(_) => tracing::debug!(method, transport, trace_id, "rpc completed"),
Err(error) => {
let retriable = error.is_retriable();
metrics::counter!(
METRIC_ERRORS,
"method" => method,
"transport" => transport,
"retriable" => retriable.to_string(),
)
.increment(1);
tracing::warn!(
method,
transport,
retriable,
trace_id,
error = %error,
"rpc failed",
);
}
}
result
}
.instrument(span)
.await
}
pub fn instrument_stream<T, S>(
method: &'static str,
transport: &'static str,
stream: S,
) -> impl futures_core::Stream<Item = Result<T>> + Send
where
S: futures_core::Stream<Item = Result<T>> + Send,
T: Send,
{
use tokio_stream::StreamExt as _;
let span = tracing::info_span!("canton.stream", method = method, transport = transport);
async_stream::stream! {
tokio::pin!(stream);
let mut items = 0u64;
loop {
let next = stream.next().instrument(span.clone()).await;
match next {
Some(Ok(item)) => {
items += 1;
yield Ok(item);
}
Some(Err(error)) => {
let retriable = error.is_retriable();
metrics::counter!(
METRIC_ERRORS,
"method" => method,
"transport" => transport,
"retriable" => retriable.to_string(),
)
.increment(1);
span.in_scope(|| {
tracing::warn!(
method,
transport,
retriable,
items,
trace_id = current_trace_id().unwrap_or_default(),
error = %error,
"stream failed",
);
});
yield Err(error);
}
None => {
span.in_scope(|| {
tracing::debug!(
method,
transport,
items,
trace_id = current_trace_id().unwrap_or_default(),
"stream ended",
);
});
return;
}
}
}
}
}
#[must_use]
pub fn current_trace_id() -> Option<String> {
#[cfg(feature = "otel")]
{
otel::current_trace_id()
}
#[cfg(not(feature = "otel"))]
{
None
}
}
#[cfg(feature = "otel")]
pub mod otel {
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, Ordering};
use opentelemetry::metrics::MeterProvider as _;
use opentelemetry::propagation::TextMapPropagator as _;
use opentelemetry::trace::TracerProvider as _;
use opentelemetry_otlp::WithExportConfig as _;
use opentelemetry_sdk::propagation::TraceContextPropagator;
pub fn otlp_tracer(
service_name: &'static str,
endpoint: impl Into<String>,
) -> Result<opentelemetry_sdk::trace::Tracer, opentelemetry::trace::TraceError> {
Ok(otlp_tracer_provider(service_name, endpoint)?.tracer(service_name))
}
pub fn otlp_tracer_provider(
service_name: &'static str,
endpoint: impl Into<String>,
) -> Result<opentelemetry_sdk::trace::TracerProvider, opentelemetry::trace::TraceError> {
let exporter = opentelemetry_otlp::SpanExporter::builder()
.with_tonic()
.with_endpoint(endpoint.into())
.build()?;
Ok(opentelemetry_sdk::trace::TracerProvider::builder()
.with_batch_exporter(exporter, opentelemetry_sdk::runtime::Tokio)
.with_resource(opentelemetry_sdk::Resource::new(vec![
opentelemetry::KeyValue::new("service.name", service_name),
]))
.build())
}
pub fn otlp_metrics(
service_name: &'static str,
endpoint: impl Into<String>,
) -> Result<opentelemetry_sdk::metrics::SdkMeterProvider, Box<dyn std::error::Error>> {
let exporter = opentelemetry_otlp::MetricExporter::builder()
.with_tonic()
.with_endpoint(endpoint.into())
.build()?;
let reader = opentelemetry_sdk::metrics::PeriodicReader::builder(
exporter,
opentelemetry_sdk::runtime::Tokio,
)
.build();
let provider = opentelemetry_sdk::metrics::SdkMeterProvider::builder()
.with_reader(reader)
.with_resource(opentelemetry_sdk::Resource::new(vec![
opentelemetry::KeyValue::new("service.name", service_name),
]))
.build();
metrics::set_global_recorder(OtelRecorder::new(provider.meter(service_name)))
.map_err(|e| -> Box<dyn std::error::Error> { Box::new(e) })?;
Ok(provider)
}
struct OtelRecorder {
meter: opentelemetry::metrics::Meter,
counters: Mutex<HashMap<String, opentelemetry::metrics::Counter<u64>>>,
gauges: Mutex<HashMap<String, opentelemetry::metrics::Gauge<f64>>>,
histograms: Mutex<HashMap<String, opentelemetry::metrics::Histogram<f64>>>,
}
impl OtelRecorder {
fn new(meter: opentelemetry::metrics::Meter) -> Self {
Self {
meter,
counters: Mutex::new(HashMap::new()),
gauges: Mutex::new(HashMap::new()),
histograms: Mutex::new(HashMap::new()),
}
}
}
fn attributes(key: &metrics::Key) -> Vec<opentelemetry::KeyValue> {
key.labels()
.map(|label| {
opentelemetry::KeyValue::new(label.key().to_string(), label.value().to_string())
})
.collect()
}
struct BridgedCounter {
counter: opentelemetry::metrics::Counter<u64>,
attributes: Vec<opentelemetry::KeyValue>,
last_absolute: AtomicU64,
}
impl metrics::CounterFn for BridgedCounter {
fn increment(&self, value: u64) {
self.counter.add(value, &self.attributes);
}
fn absolute(&self, value: u64) {
let previous = self.last_absolute.swap(value, Ordering::SeqCst);
self.counter
.add(value.saturating_sub(previous), &self.attributes);
}
}
struct BridgedGauge {
gauge: opentelemetry::metrics::Gauge<f64>,
attributes: Vec<opentelemetry::KeyValue>,
value: Mutex<f64>,
}
impl BridgedGauge {
fn apply(&self, change: impl FnOnce(f64) -> f64) {
let mut current = self
.value
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*current = change(*current);
self.gauge.record(*current, &self.attributes);
}
}
impl metrics::GaugeFn for BridgedGauge {
fn increment(&self, value: f64) {
self.apply(|current| current + value);
}
fn decrement(&self, value: f64) {
self.apply(|current| current - value);
}
fn set(&self, value: f64) {
self.apply(|_| value);
}
}
struct BridgedHistogram {
histogram: opentelemetry::metrics::Histogram<f64>,
attributes: Vec<opentelemetry::KeyValue>,
}
impl metrics::HistogramFn for BridgedHistogram {
fn record(&self, value: f64) {
self.histogram.record(value, &self.attributes);
}
}
impl metrics::Recorder for OtelRecorder {
fn describe_counter(
&self,
_key: metrics::KeyName,
_unit: Option<metrics::Unit>,
_description: metrics::SharedString,
) {
}
fn describe_gauge(
&self,
_key: metrics::KeyName,
_unit: Option<metrics::Unit>,
_description: metrics::SharedString,
) {
}
fn describe_histogram(
&self,
_key: metrics::KeyName,
_unit: Option<metrics::Unit>,
_description: metrics::SharedString,
) {
}
fn register_counter(
&self,
key: &metrics::Key,
_metadata: &metrics::Metadata<'_>,
) -> metrics::Counter {
let name = key.name().to_string();
let counter = {
let mut counters = self
.counters
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
counters
.entry(name.clone())
.or_insert_with(|| self.meter.u64_counter(name).build())
.clone()
};
metrics::Counter::from_arc(Arc::new(BridgedCounter {
counter,
attributes: attributes(key),
last_absolute: AtomicU64::new(0),
}))
}
fn register_gauge(
&self,
key: &metrics::Key,
_metadata: &metrics::Metadata<'_>,
) -> metrics::Gauge {
let name = key.name().to_string();
let gauge = {
let mut gauges = self
.gauges
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
gauges
.entry(name.clone())
.or_insert_with(|| self.meter.f64_gauge(name).build())
.clone()
};
metrics::Gauge::from_arc(Arc::new(BridgedGauge {
gauge,
attributes: attributes(key),
value: Mutex::new(0.0),
}))
}
fn register_histogram(
&self,
key: &metrics::Key,
_metadata: &metrics::Metadata<'_>,
) -> metrics::Histogram {
let name = key.name().to_string();
let histogram = {
let mut histograms = self
.histograms
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
histograms
.entry(name.clone())
.or_insert_with(|| self.meter.f64_histogram(name).build())
.clone()
};
metrics::Histogram::from_arc(Arc::new(BridgedHistogram {
histogram,
attributes: attributes(key),
}))
}
}
fn trace_context_carrier() -> std::collections::HashMap<String, String> {
use opentelemetry::trace::TraceContextExt as _;
use tracing_opentelemetry::OpenTelemetrySpanExt as _;
let context = tracing::Span::current().context();
let mut carrier = std::collections::HashMap::new();
if context.span().span_context().is_valid() {
TraceContextPropagator::new().inject_context(&context, &mut carrier);
}
carrier
}
pub(super) fn current_trace_id() -> Option<String> {
use opentelemetry::trace::TraceContextExt as _;
use tracing_opentelemetry::OpenTelemetrySpanExt as _;
let context = tracing::Span::current().context();
let span_context = context.span().span_context().clone();
span_context.is_valid().then(|| {
format!(
"{:032x}",
u128::from_be_bytes(span_context.trace_id().to_bytes())
)
})
}
pub fn inject_trace_context(headers: &mut http::HeaderMap) {
for (key, value) in trace_context_carrier() {
if let (Ok(name), Ok(val)) = (
http::header::HeaderName::try_from(key),
http::HeaderValue::from_str(&value),
) {
headers.insert(name, val);
}
}
}
pub fn inject_trace_context_metadata(metadata: &mut tonic::metadata::MetadataMap) {
for (key, value) in trace_context_carrier() {
if let (Ok(name), Ok(val)) = (
tonic::metadata::MetadataKey::from_bytes(key.as_bytes()),
tonic::metadata::MetadataValue::try_from(value),
) {
metadata.insert(name, val);
}
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use crate::Error;
use std::sync::{Arc, Mutex};
use tokio_stream::StreamExt as _;
use tracing::subscriber::set_default;
use tracing_subscriber::Layer;
use tracing_subscriber::layer::{Context, SubscriberExt};
use tracing_subscriber::registry::LookupSpan;
#[derive(Clone, Default)]
struct SpanCapture(Arc<Mutex<Vec<String>>>);
impl<S> Layer<S> for SpanCapture
where
S: tracing::Subscriber + for<'a> LookupSpan<'a>,
{
fn on_new_span(
&self,
attrs: &tracing::span::Attributes<'_>,
_id: &tracing::span::Id,
_ctx: Context<'_, S>,
) {
self.0
.lock()
.unwrap()
.push(attrs.metadata().name().to_string());
}
}
#[tokio::test]
async fn instrument_emits_span_and_metrics() {
let recorder = metrics_util::debugging::DebuggingRecorder::new();
let snapshotter = recorder.snapshotter();
recorder.install().expect("install metrics recorder");
let captured = SpanCapture::default();
let subscriber = tracing_subscriber::registry().with(captured.clone());
let _guard = set_default(subscriber);
let ok: Result<u8> = instrument("version", TRANSPORT_GRPC, async { Ok(1) }).await;
assert_eq!(ok.unwrap(), 1);
let err: Result<u8> = instrument("ledger_end", TRANSPORT_GRPC, async {
Err(Error::InvalidRequest("boom".into()))
})
.await;
assert!(err.is_err());
{
let spans = captured.0.lock().unwrap();
assert!(
spans.iter().filter(|n| *n == "canton.rpc").count() >= 2,
"expected canton.rpc spans, saw {spans:?}"
);
}
let snapshot = snapshotter.snapshot().into_vec();
let counter_total = |name: &str| -> u64 {
snapshot
.iter()
.filter(|(key, _, _, _)| key.key().name() == name)
.filter_map(|(_, _, _, value)| match value {
metrics_util::debugging::DebugValue::Counter(c) => Some(*c),
_ => None,
})
.sum()
};
assert_eq!(counter_total(METRIC_REQUESTS), 2, "two requests counted");
assert_eq!(counter_total(METRIC_ERRORS), 1, "one error counted");
let source = tokio_stream::iter(vec![
Ok(1u8),
Err(Error::Connection("the participant went away".into())),
]);
let stream = instrument_stream("updates", TRANSPORT_GRPC, source);
tokio::pin!(stream);
let mut outcomes = Vec::new();
while let Some(item) = stream.next().await {
outcomes.push(item.is_ok());
}
assert_eq!(outcomes, vec![true, false], "both items reach the caller");
let snapshot = snapshotter.snapshot().into_vec();
let counter_total = |name: &str| -> u64 {
snapshot
.iter()
.filter(|(key, _, _, _)| key.key().name() == name)
.filter_map(|(_, _, _, value)| match value {
metrics_util::debugging::DebugValue::Counter(c) => Some(*c),
_ => None,
})
.sum()
};
assert_eq!(
counter_total(METRIC_ERRORS),
2,
"the stream's mid-life failure is counted too"
);
let spans = captured.0.lock().unwrap();
assert!(
spans.iter().any(|name| name == "canton.stream"),
"expected a canton.stream span, saw {spans:?}"
);
}
#[cfg(feature = "otel")]
#[test]
fn inject_trace_context_is_a_noop_without_a_context() {
let mut headers = http::HeaderMap::new();
super::otel::inject_trace_context(&mut headers);
assert!(
headers.is_empty(),
"no trace context should be injected outside a span, saw {headers:?}"
);
}
#[cfg(feature = "otel")]
#[test]
fn the_trace_id_in_a_structured_event_is_the_active_span_s() {
use opentelemetry::trace::TracerProvider as _;
assert_eq!(
super::current_trace_id(),
None,
"no tracer installed means no trace id to name"
);
let provider = opentelemetry_sdk::trace::TracerProvider::builder().build();
let otel_layer = tracing_opentelemetry::layer().with_tracer(provider.tracer("test"));
let subscriber = tracing_subscriber::registry().with(otel_layer);
let _guard = set_default(subscriber);
let span = tracing::info_span!("test.rpc");
let _entered = span.enter();
let trace_id = super::current_trace_id().expect("a span is active");
assert_eq!(
trace_id.len(),
32,
"a W3C trace id is 32 hex digits: {trace_id}"
);
assert!(
trace_id.chars().all(|c| c.is_ascii_hexdigit()),
"not hex: {trace_id}"
);
assert_ne!(
trace_id, "00000000000000000000000000000000",
"the all-zero id means no trace, and must not be reported as one"
);
let mut headers = http::HeaderMap::new();
super::otel::inject_trace_context(&mut headers);
let traceparent = headers["traceparent"].to_str().expect("ascii").to_string();
assert!(
traceparent.contains(&trace_id),
"traceparent {traceparent} should carry trace id {trace_id}"
);
}
#[cfg(feature = "otel")]
#[test]
fn trace_context_is_injected_under_a_tracer() {
use opentelemetry::trace::TracerProvider as _;
let provider = opentelemetry_sdk::trace::TracerProvider::builder().build();
let otel_layer = tracing_opentelemetry::layer().with_tracer(provider.tracer("test"));
let subscriber = tracing_subscriber::registry().with(otel_layer);
let _guard = set_default(subscriber);
let span = tracing::info_span!("test.rpc");
let _entered = span.enter();
let mut headers = http::HeaderMap::new();
super::otel::inject_trace_context(&mut headers);
assert!(
headers.contains_key("traceparent"),
"expected a W3C traceparent header, saw {headers:?}"
);
let mut metadata = tonic::metadata::MetadataMap::new();
super::otel::inject_trace_context_metadata(&mut metadata);
assert!(
metadata.get("traceparent").is_some(),
"expected traceparent in gRPC metadata"
);
}
}