1use crate::AgentTraceContext;
2use crate::error::{TelemetryInitError, TelemetryShutdownError};
3use crate::gen_ai_metrics::GenAiMetrics;
4use crate::genai_constants;
5use crate::otel_observer::OtelInstrumentation;
6use crate::span_guard::{ErrorKind, SpanGuard};
7use crate::trace_context::{extract_trace_context, inject_trace_context};
8use aether_core::events::{
9 AgentObserver, DynObserverFactory, McpRequestInstrumentation, ObserverFactory, TraceContext,
10};
11use opentelemetry::metrics::MeterProvider as _;
12use opentelemetry::trace::{SpanBuilder, SpanKind, TraceId, TraceState, TracerProvider as _};
13use opentelemetry::{Context, KeyValue};
14use opentelemetry_otlp::{MetricExporter, SpanExporter, WithExportConfig, WithHttpConfig};
15use opentelemetry_sdk::Resource;
16use opentelemetry_sdk::metrics::SdkMeterProvider;
17use opentelemetry_sdk::metrics::periodic_reader_with_async_runtime::PeriodicReader;
18use opentelemetry_sdk::runtime::Tokio;
19use opentelemetry_sdk::trace::span_processor_with_async_runtime::BatchSpanProcessor;
20use opentelemetry_sdk::trace::{IdGenerator, RandomIdGenerator, Sampler, SdkTracerProvider};
21use std::collections::HashMap;
22use std::str::FromStr;
23
24pub struct TelemetryRuntime {
25 tracer_provider: SdkTracerProvider,
26 meter_provider: SdkMeterProvider,
27 instrumentation: OtelInstrumentation,
28}
29
30struct OtelObserverFactory {
31 instrumentation: OtelInstrumentation,
32}
33
34struct OtelMcpRequestInstrumentation {
35 span: SpanGuard,
36}
37
38pub struct TelemetryConfig {
42 pub endpoint: Option<String>,
44 pub traces_endpoint: Option<String>,
47 pub metrics_endpoint: Option<String>,
50 pub headers: HashMap<String, String>,
51 pub service_name: String,
52 pub service_version: String,
53 pub sample_ratio: f64,
54 pub capture_content: bool,
55 pub trace_context: Option<AgentTraceContext>,
56 pub traces_enabled: bool,
57 pub metrics_enabled: bool,
58}
59
60impl TelemetryRuntime {
61 pub fn new(config: &TelemetryConfig) -> Result<Self, TelemetryInitError> {
62 if !(0.0..=1.0).contains(&config.sample_ratio) {
63 return Err(TelemetryInitError::InvalidSampleRatio(config.sample_ratio));
64 }
65
66 let trace_root = resolve_trace_root(config.trace_context.as_ref())?;
67 let http_client = build_http_client(&config.headers)?;
68 let resource = Resource::builder()
69 .with_service_name(config.service_name.clone())
70 .with_attribute(KeyValue::new(genai_constants::SERVICE_VERSION, config.service_version.clone()))
71 .build();
72 let scope = genai_constants::genai_instrumentation_scope(config.service_version.clone());
73 let tracer_provider =
74 build_tracer_provider(config, resource.clone(), http_client.clone(), trace_root.trace_id)?;
75 let meter_provider = build_meter_provider(config, resource, http_client)?;
76 let instrumentation = OtelInstrumentation {
77 tracer: tracer_provider.tracer_with_scope(scope.clone()),
78 metrics: GenAiMetrics::new(&meter_provider.meter_with_scope(scope)),
79 capture_content: config.capture_content,
80 root_parent: trace_root.parent,
81 agent_name: None,
82 };
83
84 Ok(Self { tracer_provider, meter_provider, instrumentation })
85 }
86
87 pub fn observer_factory(&self) -> DynObserverFactory {
88 std::sync::Arc::new(OtelObserverFactory { instrumentation: self.instrumentation.clone() })
89 }
90
91 pub fn shutdown_or_log(&self) {
94 if let Err(error) = self.shutdown() {
95 tracing::error!("Failed to shutdown telemetry: {error}");
96 }
97 }
98
99 pub fn shutdown(&self) -> Result<(), TelemetryShutdownError> {
102 let traces = self.tracer_provider.shutdown().map_err(TelemetryShutdownError::Trace);
103 let metrics = self.meter_provider.shutdown().map_err(TelemetryShutdownError::Metric);
104 traces.and(metrics)
105 }
106}
107
108impl ObserverFactory for OtelObserverFactory {
109 fn agent(&self, agent_name: Option<&str>, parent: Option<&TraceContext>) -> Box<dyn AgentObserver> {
110 let mut instrumentation = self.instrumentation.clone();
111 instrumentation.agent_name = agent_name.map(str::to_string);
112 if let Some(parent) = parent.and_then(extract_trace_context) {
113 instrumentation.root_parent = Some(parent);
114 }
115 Box::new(crate::OtelObserver::new(instrumentation))
116 }
117
118 fn tool_call_request(&self, tool_name: &str, parent: Option<&TraceContext>) -> Box<dyn McpRequestInstrumentation> {
119 let remote_parent = parent.and_then(extract_trace_context);
120 let attributes = vec![
121 KeyValue::new(genai_constants::GEN_AI_OPERATION_NAME, "execute_tool"),
122 KeyValue::new(genai_constants::GEN_AI_TOOL_NAME, tool_name.to_string()),
123 KeyValue::new(genai_constants::MCP_METHOD_NAME, TOOLS_CALL_METHOD),
124 KeyValue::new(genai_constants::MCP_TOOL_NAME, tool_name.to_string()),
125 ];
126 let builder = SpanBuilder::from_name(format!("{TOOLS_CALL_METHOD} {tool_name}"))
127 .with_kind(SpanKind::Server)
128 .with_attributes(attributes);
129 let context = self
130 .instrumentation
131 .start_span(builder, remote_parent.as_ref().or(self.instrumentation.root_parent.as_ref()));
132 Box::new(OtelMcpRequestInstrumentation { span: SpanGuard::new(context, "MCP request cancelled") })
133 }
134}
135
136impl McpRequestInstrumentation for OtelMcpRequestInstrumentation {
137 fn trace_context(&self) -> Option<TraceContext> {
138 inject_trace_context(self.span.context())
139 }
140
141 fn finish(mut self: Box<Self>, error: Option<&str>) {
142 match error {
143 Some(error) => self.span.end_error(Some(ErrorKind::McpError), error),
144 None => self.span.end_ok(),
145 }
146 }
147}
148
149impl TelemetryConfig {
150 fn signal_endpoint(&self, signal: &str) -> Result<String, TelemetryInitError> {
151 let signal_specific_endpoint = match signal {
152 "traces" => self.traces_endpoint.as_deref(),
153 "metrics" => self.metrics_endpoint.as_deref(),
154 _ => None,
155 };
156 if let Some(endpoint) = signal_specific_endpoint.filter(|endpoint| !endpoint.is_empty()) {
157 return Ok(endpoint.to_string());
158 }
159
160 Ok(format!("{}/v1/{signal}", self.required_endpoint()?.trim_end_matches('/')))
161 }
162
163 fn required_endpoint(&self) -> Result<&str, TelemetryInitError> {
164 self.endpoint.as_deref().filter(|endpoint| !endpoint.is_empty()).ok_or(TelemetryInitError::MissingOtlpEndpoint)
165 }
166}
167
168const TOOLS_CALL_METHOD: &str = "tools/call";
169
170#[derive(Default)]
171struct TraceRoot {
172 parent: Option<Context>,
173 trace_id: Option<TraceId>,
174}
175
176#[derive(Debug)]
177struct TraceIdGenerator {
178 trace_id: TraceId,
179 random: RandomIdGenerator,
180}
181
182impl IdGenerator for TraceIdGenerator {
183 fn new_trace_id(&self) -> TraceId {
184 self.trace_id
185 }
186
187 fn new_span_id(&self) -> opentelemetry::trace::SpanId {
188 self.random.new_span_id()
189 }
190}
191
192fn resolve_trace_root(trace_context: Option<&AgentTraceContext>) -> Result<TraceRoot, TelemetryInitError> {
193 let Some(trace_context) = trace_context else {
194 return Ok(TraceRoot::default());
195 };
196
197 match trace_context {
198 AgentTraceContext::Parent(trace_context) => {
199 if let Some(tracestate) = &trace_context.tracestate {
200 TraceState::from_str(tracestate).map_err(|_| TelemetryInitError::InvalidTraceContext("tracestate"))?;
201 }
202
203 let parent =
204 extract_trace_context(trace_context).ok_or(TelemetryInitError::InvalidTraceContext("traceparent"))?;
205
206 Ok(TraceRoot { parent: Some(parent), trace_id: None })
207 }
208 AgentTraceContext::Root { trace_id } => {
209 let trace_id = parse_trace_id(trace_id)?;
210 Ok(TraceRoot { parent: None, trace_id: Some(trace_id) })
211 }
212 }
213}
214
215fn parse_trace_id(value: &str) -> Result<TraceId, TelemetryInitError> {
216 let valid_hex =
217 value.len() == 32 && value.bytes().all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte));
218 if !valid_hex {
219 return Err(TelemetryInitError::InvalidTraceContext("traceId"));
220 }
221
222 let trace_id = TraceId::from_hex(value).map_err(|_| TelemetryInitError::InvalidTraceContext("traceId"))?;
223 if trace_id == TraceId::INVALID {
224 return Err(TelemetryInitError::InvalidTraceContext("traceId"));
225 }
226
227 Ok(trace_id)
228}
229
230fn build_tracer_provider(
231 config: &TelemetryConfig,
232 resource: Resource,
233 http_client: reqwest::Client,
234 trace_id: Option<TraceId>,
235) -> Result<SdkTracerProvider, TelemetryInitError> {
236 let builder = SdkTracerProvider::builder().with_resource(resource);
237 let builder = match trace_id {
238 Some(trace_id) => {
239 builder.with_id_generator(TraceIdGenerator { trace_id, random: RandomIdGenerator::default() })
240 }
241 None => builder,
242 };
243 if !config.traces_enabled {
244 return Ok(builder.with_sampler(Sampler::AlwaysOff).build());
245 }
246
247 let endpoint = config.signal_endpoint("traces")?;
248 let exporter = SpanExporter::builder()
249 .with_http()
250 .with_endpoint(endpoint)
251 .with_http_client(http_client)
252 .build()
253 .map_err(TelemetryInitError::TraceExporter)?;
254
255 Ok(builder
256 .with_sampler(Sampler::ParentBased(Box::new(Sampler::TraceIdRatioBased(config.sample_ratio))))
257 .with_span_processor(BatchSpanProcessor::builder(exporter, Tokio).build())
258 .build())
259}
260
261fn build_meter_provider(
262 config: &TelemetryConfig,
263 resource: Resource,
264 http_client: reqwest::Client,
265) -> Result<SdkMeterProvider, TelemetryInitError> {
266 let builder = SdkMeterProvider::builder().with_resource(resource);
267 if !config.metrics_enabled {
268 return Ok(builder.build());
269 }
270
271 let endpoint = config.signal_endpoint("metrics")?;
272 let exporter = MetricExporter::builder()
273 .with_http()
274 .with_endpoint(endpoint)
275 .with_http_client(http_client)
276 .build()
277 .map_err(TelemetryInitError::MetricExporter)?;
278
279 Ok(builder.with_reader(PeriodicReader::builder(exporter, Tokio).build()).build())
280}
281
282fn build_http_client(headers: &HashMap<String, String>) -> Result<reqwest::Client, TelemetryInitError> {
283 let mut parsed = reqwest::header::HeaderMap::new();
284 for (key, value) in headers {
285 let name = reqwest::header::HeaderName::from_bytes(key.as_bytes())
286 .map_err(|_| TelemetryInitError::InvalidHeaderName(key.clone()))?;
287 let value = reqwest::header::HeaderValue::from_str(value)
288 .map_err(|_| TelemetryInitError::InvalidHeaderValue(key.clone()))?;
289 parsed.insert(name, value);
290 }
291 reqwest::Client::builder().default_headers(parsed).build().map_err(TelemetryInitError::HttpClient)
292}