use std::collections::HashMap;
use std::sync::LazyLock;
use std::sync::Mutex;
use serde::{Deserialize, Serialize};
use crate::core::context::get_current_task;
use crate::error::DexcostError;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct TraceLink {
pub provider: String,
pub trace_id: String,
}
static TRACE_LINKS: LazyLock<Mutex<HashMap<String, Vec<TraceLink>>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
fn lock_links() -> std::sync::MutexGuard<'static, HashMap<String, Vec<TraceLink>>> {
TRACE_LINKS.lock().unwrap_or_else(|e| {
eprintln!("[dexcost] trace-link mutex poisoned, recovering: {}", e);
e.into_inner()
})
}
pub fn link_trace(provider: &str, trace_id: &str) -> Result<(), DexcostError> {
let task = get_current_task().ok_or(DexcostError::NotInitialized)?;
let mut guard = lock_links();
let entry = guard.entry(task.task_id.clone()).or_default();
entry.push(TraceLink {
provider: provider.to_string(),
trace_id: trace_id.to_string(),
});
Ok(())
}
pub fn get_trace_links() -> Vec<TraceLink> {
let task = match get_current_task() {
Some(t) => t,
None => return Vec::new(),
};
let guard = lock_links();
guard.get(&task.task_id).cloned().unwrap_or_default()
}
pub fn get_trace_links_for_task(task_id: &str) -> Vec<TraceLink> {
lock_links().get(task_id).cloned().unwrap_or_default()
}
pub fn clear_trace_links_for_task(task_id: &str) {
lock_links().remove(task_id);
}
pub fn clear_all_trace_links() {
lock_links().clear();
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::context::with_task;
use crate::core::models::Task;
static TEST_LOCK: LazyLock<Mutex<()>> = LazyLock::new(|| Mutex::new(()));
fn test_lock() -> std::sync::MutexGuard<'static, ()> {
TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner())
}
#[allow(clippy::await_holding_lock)]
#[tokio::test]
async fn link_trace_inside_active_task() {
let _g = test_lock();
clear_all_trace_links();
let task = Task::new("trace_test");
let task_id = task.task_id.clone();
with_task(task, async {
link_trace("langfuse", "trace-abc-123").expect("link inside task");
link_trace("langsmith", "run-def-456").expect("second link");
let links = get_trace_links();
assert_eq!(links.len(), 2);
assert_eq!(links[0].provider, "langfuse");
assert_eq!(links[0].trace_id, "trace-abc-123");
assert_eq!(links[1].provider, "langsmith");
})
.await;
assert!(get_trace_links().is_empty());
let by_id = get_trace_links_for_task(&task_id);
assert_eq!(by_id.len(), 2);
clear_trace_links_for_task(&task_id);
assert!(get_trace_links_for_task(&task_id).is_empty());
}
#[tokio::test]
async fn link_trace_outside_active_task_errors() {
let _g = test_lock();
clear_all_trace_links();
let result = link_trace("langfuse", "trace-1");
assert!(matches!(result, Err(DexcostError::NotInitialized)));
assert!(get_trace_links().is_empty());
}
#[test]
fn trace_link_round_trips_json() {
let link = TraceLink {
provider: "otel".to_string(),
trace_id: "abcd".to_string(),
};
let s = serde_json::to_string(&link).unwrap();
let parsed: TraceLink = serde_json::from_str(&s).unwrap();
assert_eq!(parsed, link);
}
}