use opentelemetry::KeyValue;
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
pub mod attributes {
pub const TRACE_SESSION_ID: &str = "langfuse.session.id";
pub const TRACE_USER_ID: &str = "langfuse.user.id";
pub const TRACE_TAGS: &str = "langfuse.trace.tags";
pub const TRACE_METADATA: &str = "langfuse.trace.metadata";
pub const TRACE_NAME: &str = "langfuse.trace.name";
}
#[derive(Clone, Default)]
pub struct LangfuseContext {
attributes: Arc<RwLock<HashMap<String, String>>>,
}
impl LangfuseContext {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn set_session_id(&self, session_id: impl Into<String>) -> &Self {
self.set_attribute(attributes::TRACE_SESSION_ID, session_id)
}
pub fn set_user_id(&self, user_id: impl Into<String>) -> &Self {
self.set_attribute(attributes::TRACE_USER_ID, user_id)
}
pub fn add_tags(&self, tags: impl IntoIterator<Item = String>) -> &Self {
let tags: Vec<String> = tags.into_iter().collect();
self.set_attribute(
attributes::TRACE_TAGS,
serde_json::to_string(&tags).expect("serializing strings cannot fail"),
)
}
pub fn add_tag(&self, tag: impl Into<String>) -> &Self {
let mut attributes = self.attributes.write().unwrap();
let mut tags = attributes
.get(attributes::TRACE_TAGS)
.and_then(|value| serde_json::from_str::<Vec<String>>(value).ok())
.unwrap_or_default();
tags.push(tag.into());
attributes.insert(
attributes::TRACE_TAGS.to_string(),
serde_json::to_string(&tags).expect("serializing strings cannot fail"),
);
drop(attributes);
self
}
pub fn set_metadata(&self, metadata: impl Into<serde_json::Value>) -> &Self {
self.set_attribute(attributes::TRACE_METADATA, metadata.into().to_string())
}
pub fn set_attribute(&self, key: impl Into<String>, value: impl Into<String>) -> &Self {
self.attributes
.write()
.unwrap()
.insert(key.into(), value.into());
self
}
pub fn set_trace_name(&self, name: impl Into<String>) -> &Self {
self.set_attribute(attributes::TRACE_NAME, name)
}
pub fn clear(&self) {
self.attributes.write().unwrap().clear();
}
#[must_use]
pub fn get_attributes(&self) -> Vec<KeyValue> {
self.attributes
.read()
.unwrap()
.iter()
.map(|(key, value)| KeyValue::new(key.clone(), value.clone()))
.collect()
}
#[must_use]
pub fn has_attribute(&self, key: &str) -> bool {
self.attributes.read().unwrap().contains_key(key)
}
#[must_use]
pub fn get_attribute(&self, key: &str) -> Option<String> {
self.attributes.read().unwrap().get(key).cloned()
}
}
#[cfg(test)]
mod tests {
use super::{attributes, LangfuseContext};
#[test]
fn attributes_survive_clone_and_clear() {
let context = LangfuseContext::new();
context.set_session_id("session-1").add_tag("first");
let cloned = context.clone();
cloned.add_tag("second");
assert_eq!(
context.get_attribute(attributes::TRACE_SESSION_ID),
Some("session-1".to_string())
);
assert_eq!(
context.get_attribute(attributes::TRACE_TAGS),
Some(r#"["first","second"]"#.to_string())
);
assert_eq!(context.get_attributes().len(), 2);
cloned.clear();
assert_eq!(context.get_attributes().len(), 0);
}
}