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;
match &result {
Ok(_) => tracing::debug!(method, transport, "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, error = %error, "rpc failed");
}
}
result
}
.instrument(span)
.await
}
#[cfg(feature = "otel")]
pub mod otel {
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> {
let exporter = opentelemetry_otlp::SpanExporter::builder()
.with_tonic()
.with_endpoint(endpoint.into())
.build()?;
let provider = opentelemetry_sdk::trace::TracerProvider::builder()
.with_batch_exporter(exporter, opentelemetry_sdk::runtime::Tokio)
.build();
Ok(provider.tracer(service_name))
}
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 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 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");
}
#[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 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"
);
}
}