use opentelemetry::trace::{FutureExt, SpanKind, Status, TraceContextExt, Tracer};
use opentelemetry::{Context, KeyValue, global};
use std::future::Future;
use std::sync::Arc;
pub const PORT_OPERATION: &str = "port.operation";
pub const PORT_PROVIDER_HINT: &str = "port.provider_hint";
pub const PORT_NAME: &str = "port.name";
const INSTRUMENTATION_SCOPE: &str = "otel-bootstrap/instrumented-port";
#[derive(Debug, Clone)]
pub struct Instrumented<P> {
inner: P,
port_name: &'static str,
provider_hint: Option<Arc<str>>,
}
impl<P> Instrumented<P> {
pub fn new(inner: P, port_name: &'static str, provider_hint: Option<&str>) -> Self {
Self {
inner,
port_name,
provider_hint: provider_hint.map(Arc::from),
}
}
pub fn inner(&self) -> &P {
&self.inner
}
pub async fn call<'a, F, Fut, T, E>(&'a self, operation: &str, op: F) -> Result<T, E>
where
F: FnOnce(&'a P) -> Fut,
Fut: Future<Output = Result<T, E>>,
E: std::fmt::Display,
{
let tracer = global::tracer(INSTRUMENTATION_SCOPE);
let span_name = format!("{}.{}", self.port_name, operation);
let mut attributes = vec![
KeyValue::new(PORT_NAME, self.port_name),
KeyValue::new(PORT_OPERATION, operation.to_owned()),
];
if let Some(hint) = &self.provider_hint {
attributes.push(KeyValue::new(PORT_PROVIDER_HINT, hint.to_string()));
}
let parent_cx = Context::current();
let span = if parent_cx.span().span_context().is_valid() {
tracer
.span_builder(span_name)
.with_kind(SpanKind::Client)
.with_attributes(attributes)
.start_with_context(&tracer, &parent_cx)
} else {
tracer
.span_builder(span_name)
.with_kind(SpanKind::Client)
.with_attributes(attributes)
.start_with_context(&tracer, &Context::new())
};
let cx = Context::current_with_span(span);
let fut = op(&self.inner);
let result = fut.with_context(cx.clone()).await;
if result.is_err() {
cx.span().set_status(Status::error("port call failed"));
}
result
}
}
pub type InstrumentedArc<P> = Instrumented<Arc<P>>;
#[cfg(test)]
mod tests {
use super::*;
use opentelemetry_sdk::error::OTelSdkResult;
use opentelemetry_sdk::trace::{SdkTracerProvider, SpanData, SpanExporter};
use std::sync::{LazyLock, Mutex};
#[derive(Debug, Clone, Default)]
struct CapturingExporter {
spans: Arc<Mutex<Vec<SpanData>>>,
}
impl SpanExporter for CapturingExporter {
async fn export(&self, batch: Vec<SpanData>) -> OTelSdkResult {
self.spans.lock().unwrap().extend(batch);
Ok(())
}
}
struct Adapter;
impl Adapter {
async fn wrap(&self) -> Result<&'static str, &'static str> {
Ok("wrapped")
}
async fn unwrap(&self) -> Result<&'static str, &'static str> {
Err("boom")
}
}
fn install_capturing_tracer() -> (SdkTracerProvider, Arc<Mutex<Vec<SpanData>>>) {
let exporter = CapturingExporter::default();
let spans = exporter.spans.clone();
let provider = SdkTracerProvider::builder()
.with_simple_exporter(exporter)
.build();
(provider, spans)
}
static GLOBAL_TRACER_LOCK: LazyLock<tokio::sync::Mutex<()>> =
LazyLock::new(|| tokio::sync::Mutex::new(()));
#[tokio::test]
async fn call_emits_client_span_with_operation_and_provider_hint() {
let _guard = GLOBAL_TRACER_LOCK.lock().await;
let (provider, spans) = install_capturing_tracer();
opentelemetry::global::set_tracer_provider(provider.clone());
let wrapped = Instrumented::new(Adapter, "KmsProvider", Some("aws-kms"));
let out = wrapped.call("wrap", |inner| inner.wrap()).await.unwrap();
assert_eq!(out, "wrapped");
provider.force_flush().unwrap();
let captured = spans.lock().unwrap();
assert_eq!(
captured.len(),
1,
"exactly one span must be emitted per call"
);
let span = &captured[0];
assert_eq!(span.span_kind, opentelemetry::trace::SpanKind::Client);
assert_eq!(span.name, "KmsProvider.wrap");
let has_attr = |key: &str, value: &str| {
span.attributes
.iter()
.any(|kv| kv.key.as_str() == key && kv.value.as_str() == value)
};
assert!(has_attr(PORT_NAME, "KmsProvider"));
assert!(has_attr(PORT_OPERATION, "wrap"));
assert!(has_attr(PORT_PROVIDER_HINT, "aws-kms"));
}
#[tokio::test]
async fn call_records_error_status_and_returns_error_unchanged() {
let _guard = GLOBAL_TRACER_LOCK.lock().await;
let (provider, spans) = install_capturing_tracer();
opentelemetry::global::set_tracer_provider(provider.clone());
let wrapped = Instrumented::new(Adapter, "KmsProvider", None);
let out = wrapped.call("unwrap", |inner| inner.unwrap()).await;
assert_eq!(out, Err("boom"));
provider.force_flush().unwrap();
let captured = spans.lock().unwrap();
assert_eq!(captured.len(), 1);
let span = &captured[0];
assert_eq!(span.name, "KmsProvider.unwrap");
assert!(matches!(
&span.status,
opentelemetry::trace::Status::Error { .. }
));
let has_no_hint = !span
.attributes
.iter()
.any(|kv| kv.key.as_str() == PORT_PROVIDER_HINT);
assert!(has_no_hint, "no provider hint attribute when none supplied");
}
#[tokio::test]
async fn call_mints_valid_trace_id_with_no_parent() {
let _guard = GLOBAL_TRACER_LOCK.lock().await;
let (provider, spans) = install_capturing_tracer();
opentelemetry::global::set_tracer_provider(provider.clone());
let wrapped = Instrumented::new(Adapter, "KmsProvider", None);
wrapped.call("unwrap", |inner| inner.unwrap()).await.ok();
provider.force_flush().unwrap();
let captured = spans.lock().unwrap();
assert_eq!(captured.len(), 1);
assert!(
captured[0].span_context.is_valid(),
"root span must carry a valid (non-zero) trace id"
);
}
#[tokio::test]
async fn call_mints_fresh_trace_id_when_current_parent_is_invalid() {
use opentelemetry::trace::SpanContext;
let _guard = GLOBAL_TRACER_LOCK.lock().await;
let (provider, spans) = install_capturing_tracer();
opentelemetry::global::set_tracer_provider(provider.clone());
let invalid_parent = Context::new().with_remote_span_context(SpanContext::empty_context());
assert!(!invalid_parent.span().span_context().is_valid());
let _attach = invalid_parent.attach();
let wrapped = Instrumented::new(Adapter, "KmsProvider", None);
wrapped.call("unwrap", |inner| inner.unwrap()).await.ok();
provider.force_flush().unwrap();
let captured = spans.lock().unwrap();
assert_eq!(captured.len(), 1);
assert!(
captured[0].span_context.is_valid(),
"span must mint a fresh root trace id instead of inheriting the invalid parent's 0-bit id"
);
}
#[tokio::test]
async fn instrumented_arc_alias_wraps_shared_port_handle() {
let _guard = GLOBAL_TRACER_LOCK.lock().await;
let shared: Arc<Adapter> = Arc::new(Adapter);
let wrapped: InstrumentedArc<Adapter> = Instrumented::new(shared, "KmsProvider", None);
let out = wrapped.call("wrap", |inner| inner.wrap()).await.unwrap();
assert_eq!(out, "wrapped");
}
}