use async_trait::async_trait;
use opentelemetry::global::{self, BoxedSpan, BoxedTracer};
use opentelemetry::trace::{Span, TraceContextExt, Tracer};
use opentelemetry::Context;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::base::CallbackHandler;
use crate::run_tree::RunTree;
use lc_schema::Message;
pub struct OtelHandler {
tracer: BoxedTracer,
spans: Arc<Mutex<HashMap<String, BoxedSpan>>>,
}
impl OtelHandler {
pub fn new(tracer: BoxedTracer) -> Self {
Self {
tracer,
spans: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn from_global(name: &str) -> Self {
Self::new(global::tracer(name.to_string()))
}
async fn start_span(&self, name: &str, run: &RunTree) {
let parent_cx = {
let spans = self.spans.lock().await;
run.parent_run_id
.and_then(|parent_id| {
spans
.get(&parent_id.to_string())
.map(|parent_span| parent_span.span_context().clone())
})
.map(|parent_ctx| Context::new().with_remote_span_context(parent_ctx))
};
let span = match parent_cx {
Some(cx) => self.tracer.start_with_context(name.to_string(), &cx),
None => self.tracer.start(name.to_string()),
};
self.spans.lock().await.insert(run.id.to_string(), span);
}
async fn end_span(&self, run: &RunTree) {
if let Some(mut span) = self.spans.lock().await.remove(&run.id.to_string()) {
span.end();
}
}
async fn add_event(&self, name: &str, run: &RunTree) {
let mut spans = self.spans.lock().await;
if let Some(span) = spans.get_mut(&run.id.to_string()) {
span.add_event(name.to_string(), Vec::new());
}
}
pub async fn active_span_count(&self) -> usize {
self.spans.lock().await.len()
}
}
#[async_trait]
impl CallbackHandler for OtelHandler {
async fn on_run_start(&self, run: &RunTree) {
self.start_span("run", run).await;
}
async fn on_run_end(&self, run: &RunTree) {
self.end_span(run).await;
}
async fn on_run_error(&self, run: &RunTree, error: &str) {
self.add_event(&format!("error: {}", error), run).await;
self.end_span(run).await;
}
async fn on_llm_start(&self, run: &RunTree, _messages: &[Message]) {
self.start_span("llm", run).await;
let mut spans = self.spans.lock().await;
if let Some(span) = spans.get_mut(&run.id.to_string()) {
if let Some(system) = run.metadata.get("model_provider") {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.system",
system.as_str().unwrap_or("unknown").to_string(),
));
}
if let Some(model) = run.metadata.get("model") {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.request.model",
model.as_str().unwrap_or("unknown").to_string(),
));
}
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.operation.name",
"chat".to_string(),
));
}
}
async fn on_llm_end(&self, run: &RunTree, _response: &str) {
let mut spans = self.spans.lock().await;
if let Some(span) = spans.get_mut(&run.id.to_string()) {
if let Some(reason) = run.metadata.get("finish_reason") {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.response.finish_reason",
reason.as_str().unwrap_or("stop").to_string(),
));
}
if let Some(model) = run.metadata.get("response_model") {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.response.model",
model.as_str().unwrap_or("").to_string(),
));
}
if let Some(tokens) = run.metadata.get("token_usage") {
if let Some(obj) = tokens.as_object() {
if let Some(p) = obj.get("prompt_tokens").and_then(|v| v.as_u64()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.client.token.usage.prompt_tokens",
p as i64,
));
}
if let Some(c) = obj.get("completion_tokens").and_then(|v| v.as_u64()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.client.token.usage.completion_tokens",
c as i64,
));
}
if let Some(p) = obj.get("cache_read_input_tokens").and_then(|v| v.as_u64()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.usage.cache_read.input_tokens",
p as i64,
));
}
if let Some(p) = obj
.get("cache_creation_input_tokens")
.and_then(|v| v.as_u64())
{
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.usage.cache_creation.input_tokens",
p as i64,
));
}
if let Some(p) = obj.get("reasoning_output_tokens").and_then(|v| v.as_u64()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.usage.reasoning.output_tokens",
p as i64,
));
}
}
}
if let Some(max) = run.metadata.get("max_tokens") {
if let Some(v) = max.as_u64() {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.request.max_tokens",
v as i64,
));
}
}
if let Some(temp) = run.metadata.get("temperature") {
if let Some(v) = temp.as_f64() {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.request.temperature",
v,
));
}
}
}
drop(spans);
self.end_span(run).await;
}
async fn on_llm_new_token(&self, run: &RunTree, _token: &str) {
self.add_event("token", run).await;
}
async fn on_llm_error(&self, run: &RunTree, error: &str) {
self.add_event(&format!("llm error: {}", error), run).await;
self.end_span(run).await;
}
async fn on_chain_start(&self, run: &RunTree, _inputs: &serde_json::Value) {
self.start_span("chain", run).await;
let mut spans = self.spans.lock().await;
if let Some(span) = spans.get_mut(&run.id.to_string()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.operation.name",
"chain".to_string(),
));
}
}
async fn on_chain_end(&self, run: &RunTree, _outputs: &serde_json::Value) {
self.end_span(run).await;
}
async fn on_chain_error(&self, run: &RunTree, error: &str) {
self.add_event(&format!("chain error: {}", error), run)
.await;
self.end_span(run).await;
}
async fn on_tool_start(&self, run: &RunTree, tool_name: &str, _input: &str) {
self.start_span("tool", run).await;
let mut spans = self.spans.lock().await;
if let Some(span) = spans.get_mut(&run.id.to_string()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.tool.name",
tool_name.to_string(),
));
}
}
async fn on_tool_end(&self, run: &RunTree, _output: &str) {
self.end_span(run).await;
}
async fn on_tool_error(&self, run: &RunTree, error: &str) {
self.add_event(&format!("tool error: {}", error), run).await;
self.end_span(run).await;
}
async fn on_retriever_start(&self, run: &RunTree, _query: &str) {
self.start_span("retriever", run).await;
let mut spans = self.spans.lock().await;
if let Some(span) = spans.get_mut(&run.id.to_string()) {
span.set_attribute(opentelemetry::KeyValue::new(
"gen_ai.operation.name",
"retrieve".to_string(),
));
}
}
async fn on_retriever_end(&self, run: &RunTree, _documents: &[serde_json::Value]) {
self.end_span(run).await;
}
async fn on_retriever_error(&self, run: &RunTree, error: &str) {
self.add_event(&format!("retriever error: {}", error), run)
.await;
self.end_span(run).await;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_from_global_does_not_panic() {
let _h = OtelHandler::from_global("test");
}
#[tokio::test]
async fn test_span_start_end_balance() {
let h = OtelHandler::from_global("test");
assert_eq!(h.active_span_count().await, 0);
let run1 = RunTree::new("run1", crate::RunType::Llm, serde_json::json!({}));
h.start_span("a", &run1).await;
let run2 = RunTree::new("run2", crate::RunType::Tool, serde_json::json!({}));
h.start_span("b", &run2).await;
assert_eq!(h.active_span_count().await, 2);
h.end_span(&run2).await;
assert_eq!(h.active_span_count().await, 1);
h.end_span(&run1).await;
assert_eq!(h.active_span_count().await, 0);
}
#[tokio::test]
async fn test_end_span_when_empty_is_noop() {
let h = OtelHandler::from_global("test");
let run = RunTree::new("nonexistent", crate::RunType::Llm, serde_json::json!({}));
h.end_span(&run).await;
assert_eq!(h.active_span_count().await, 0);
}
#[tokio::test]
async fn test_add_event_to_active_span() {
let h = OtelHandler::from_global("test");
let run = RunTree::new("run1", crate::RunType::Llm, serde_json::json!({}));
h.start_span("a", &run).await;
h.add_event("something", &run).await;
assert_eq!(h.active_span_count().await, 1);
h.end_span(&run).await;
}
}