use opentelemetry::global;
use opentelemetry::trace::TraceError;
use opentelemetry::{KeyValue, Value};
use opentelemetry_otlp::{Protocol, SpanExporter, WithExportConfig, WithHttpConfig};
use opentelemetry_sdk::propagation::TraceContextPropagator;
use opentelemetry_sdk::resource::Resource;
use opentelemetry_sdk::runtime;
use opentelemetry_sdk::trace::TracerProvider;
use std::collections::HashMap;
use std::env;
use std::time::Duration;
const DEFAULT_ENDPOINT: &str = "http://localhost:4318";
const DEFAULT_SERVICE_NAME: &str = "langchainrust";
const DEFAULT_TIMEOUT_SECS: u64 = 10;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct OtlpConfig {
pub endpoint: String,
pub service_name: String,
pub headers: HashMap<String, String>,
pub timeout: Duration,
}
impl Default for OtlpConfig {
fn default() -> Self {
Self {
endpoint: env::var("OTEL_EXPORTER_OTLP_ENDPOINT")
.ok()
.filter(|s| !s.is_empty())
.unwrap_or_else(|| DEFAULT_ENDPOINT.to_string()),
service_name: env::var("OTEL_SERVICE_NAME")
.ok()
.filter(|s| !s.is_empty())
.unwrap_or_else(|| DEFAULT_SERVICE_NAME.to_string()),
headers: env::var("OTEL_EXPORTER_OTLP_HEADERS")
.map(|raw| parse_headers(&raw))
.unwrap_or_default(),
timeout: env::var("OTEL_EXPORTER_OTLP_TIMEOUT")
.ok()
.and_then(|s| s.parse::<u64>().ok())
.map(Duration::from_secs)
.unwrap_or(Duration::from_secs(DEFAULT_TIMEOUT_SECS)),
}
}
}
pub struct OtlpGuard {
provider: TracerProvider,
}
impl Drop for OtlpGuard {
fn drop(&mut self) {
let _ = self.provider.shutdown();
}
}
pub fn install_otlp_pipeline() -> Result<OtlpGuard, TraceError> {
install_otlp_pipeline_with(OtlpConfig::default())
}
pub fn install_otlp_pipeline_with(config: OtlpConfig) -> Result<OtlpGuard, TraceError> {
let exporter: SpanExporter = SpanExporter::builder()
.with_http()
.with_endpoint(config.endpoint)
.with_protocol(Protocol::HttpJson)
.with_timeout(config.timeout)
.with_headers(config.headers)
.build()?;
let resource = Resource::new([KeyValue::new(
"service.name",
Value::from(config.service_name),
)]);
let provider = TracerProvider::builder()
.with_batch_exporter(exporter, runtime::Tokio)
.with_resource(resource)
.build();
global::set_text_map_propagator(TraceContextPropagator::new());
global::set_tracer_provider(provider.clone());
Ok(OtlpGuard { provider })
}
fn parse_headers(raw: &str) -> HashMap<String, String> {
raw.split(',')
.filter_map(|pair| {
let mut parts = pair.splitn(2, '=');
let key = parts.next()?.trim();
let value = parts.next()?;
if key.is_empty() {
return None;
}
Some((key.to_string(), percent_decode(value.trim())))
})
.collect()
}
fn percent_decode(input: &str) -> String {
let bytes = input.as_bytes();
let mut out: Vec<u8> = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' && i + 2 < bytes.len() {
let hi = (bytes[i + 1] as char).to_digit(16);
let lo = (bytes[i + 2] as char).to_digit(16);
if let (Some(hi), Some(lo)) = (hi, lo) {
out.push((hi * 16 + lo) as u8);
i += 3;
continue;
}
}
out.push(bytes[i]);
i += 1;
}
String::from_utf8(out).unwrap_or_else(|_| input.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_comma_separated_headers() {
let map = parse_headers("authorization=Bearer abc,x-custom=v");
assert_eq!(map.get("authorization").unwrap(), "Bearer abc");
assert_eq!(map.get("x-custom").unwrap(), "v");
}
#[test]
fn skips_malformed_pairs_and_blank_keys() {
let map = parse_headers("no-equals, ,=novalue,k=v");
assert_eq!(map.len(), 1);
assert_eq!(map.get("k").unwrap(), "v");
}
#[test]
fn decodes_percent_encoded_values() {
assert_eq!(percent_decode("a%2Cb"), "a,b");
assert_eq!(percent_decode("a%3Db"), "a=b");
assert_eq!(percent_decode("%41%42%43"), "ABC");
assert_eq!(percent_decode("plain"), "plain");
assert_eq!(percent_decode("100%"), "100%");
}
#[test]
fn default_config_falls_back_when_env_absent() {
env::remove_var("OTEL_EXPORTER_OTLP_ENDPOINT");
let config = OtlpConfig::default();
assert_eq!(config.endpoint, DEFAULT_ENDPOINT);
assert_eq!(config.timeout, Duration::from_secs(DEFAULT_TIMEOUT_SECS));
}
}