use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct TimeRecord {
#[serde(default)]
pub event_time: f64,
#[serde(default)]
pub message_time: Option<i64>,
}
impl TimeRecord {
pub fn now() -> Self {
Self {
event_time: Utc::now().timestamp_millis() as f64 / 1000.0,
message_time: None,
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct MessageContent {
#[serde(default)]
pub text: Option<String>,
#[serde(default)]
pub json_payload: Option<String>,
}
impl MessageContent {
pub fn from_text(text: impl Into<String>) -> Self {
Self {
text: Some(text.into()),
json_payload: None,
}
}
pub fn from_json(value: &Value) -> Self {
Self {
text: None,
json_payload: Some(value.to_string()),
}
}
pub fn as_text(&self) -> Option<&str> {
self.text.as_deref().or(self.json_payload.as_deref())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum EventType {
Cais,
Environment,
Runtime,
}
impl std::fmt::Display for EventType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EventType::Cais => write!(f, "cais"),
EventType::Environment => write!(f, "environment"),
EventType::Runtime => write!(f, "runtime"),
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct BaseEventFields {
#[serde(default)]
pub system_instance_id: String,
#[serde(default)]
pub time_record: TimeRecord,
#[serde(default)]
pub metadata: HashMap<String, Value>,
#[serde(default)]
pub event_metadata: Option<Vec<Value>>,
}
impl BaseEventFields {
pub fn new(system_instance_id: impl Into<String>) -> Self {
Self {
system_instance_id: system_instance_id.into(),
time_record: TimeRecord::now(),
metadata: HashMap::new(),
event_metadata: None,
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LMCAISEvent {
#[serde(flatten)]
pub base: BaseEventFields,
#[serde(default)]
pub model_name: String,
#[serde(default)]
pub provider: Option<String>,
#[serde(default)]
pub input_tokens: Option<i32>,
#[serde(default)]
pub output_tokens: Option<i32>,
#[serde(default)]
pub total_tokens: Option<i32>,
#[serde(default)]
pub cost_usd: Option<f64>,
#[serde(default)]
pub latency_ms: Option<i32>,
#[serde(default)]
pub span_id: Option<String>,
#[serde(default)]
pub trace_id: Option<String>,
#[serde(default)]
pub call_records: Option<Vec<LLMCallRecord>>,
#[serde(default)]
pub system_state_before: Option<Value>,
#[serde(default)]
pub system_state_after: Option<Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct EnvironmentEvent {
#[serde(flatten)]
pub base: BaseEventFields,
#[serde(default)]
pub reward: f64,
#[serde(default)]
pub terminated: bool,
#[serde(default)]
pub truncated: bool,
#[serde(default)]
pub system_state_before: Option<Value>,
#[serde(default)]
pub system_state_after: Option<Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct RuntimeEvent {
#[serde(flatten)]
pub base: BaseEventFields,
#[serde(default)]
pub actions: Vec<i32>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "event_type", rename_all = "snake_case")]
pub enum TracingEvent {
Cais(LMCAISEvent),
Environment(EnvironmentEvent),
Runtime(RuntimeEvent),
}
impl TracingEvent {
pub fn event_type(&self) -> EventType {
match self {
TracingEvent::Cais(_) => EventType::Cais,
TracingEvent::Environment(_) => EventType::Environment,
TracingEvent::Runtime(_) => EventType::Runtime,
}
}
pub fn base(&self) -> &BaseEventFields {
match self {
TracingEvent::Cais(e) => &e.base,
TracingEvent::Environment(e) => &e.base,
TracingEvent::Runtime(e) => &e.base,
}
}
pub fn time_record(&self) -> &TimeRecord {
&self.base().time_record
}
pub fn system_instance_id(&self) -> &str {
&self.base().system_instance_id
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LLMUsage {
#[serde(default)]
pub input_tokens: Option<i32>,
#[serde(default)]
pub output_tokens: Option<i32>,
#[serde(default)]
pub total_tokens: Option<i32>,
#[serde(default)]
pub reasoning_tokens: Option<i32>,
#[serde(default)]
pub reasoning_input_tokens: Option<i32>,
#[serde(default)]
pub reasoning_output_tokens: Option<i32>,
#[serde(default)]
pub cache_read_tokens: Option<i32>,
#[serde(default)]
pub cache_write_tokens: Option<i32>,
#[serde(default)]
pub billable_input_tokens: Option<i32>,
#[serde(default)]
pub billable_output_tokens: Option<i32>,
#[serde(default)]
pub cost_usd: Option<f64>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LLMRequestParams {
#[serde(default)]
pub temperature: Option<f64>,
#[serde(default)]
pub top_p: Option<f64>,
#[serde(default)]
pub max_tokens: Option<i32>,
#[serde(default)]
pub stop: Option<Vec<String>>,
#[serde(default)]
pub top_k: Option<i32>,
#[serde(default)]
pub presence_penalty: Option<f64>,
#[serde(default)]
pub frequency_penalty: Option<f64>,
#[serde(default)]
pub repetition_penalty: Option<f64>,
#[serde(default)]
pub seed: Option<i32>,
#[serde(default)]
pub n: Option<i32>,
#[serde(default)]
pub best_of: Option<i32>,
#[serde(default)]
pub response_format: Option<Value>,
#[serde(default)]
pub json_mode: Option<bool>,
#[serde(default)]
pub tool_config: Option<Value>,
#[serde(default)]
pub raw_params: HashMap<String, Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LLMContentPart {
#[serde(rename = "type", default)]
pub content_type: String,
#[serde(default)]
pub text: Option<String>,
#[serde(default)]
pub data: Option<Value>,
#[serde(default)]
pub mime_type: Option<String>,
#[serde(default)]
pub uri: Option<String>,
#[serde(default)]
pub base64_data: Option<String>,
#[serde(default)]
pub size_bytes: Option<i64>,
#[serde(default)]
pub sha256: Option<String>,
#[serde(default)]
pub width: Option<i32>,
#[serde(default)]
pub height: Option<i32>,
#[serde(default)]
pub duration_ms: Option<i32>,
#[serde(default)]
pub sample_rate: Option<i32>,
#[serde(default)]
pub channels: Option<i32>,
#[serde(default)]
pub language: Option<String>,
}
impl LLMContentPart {
pub fn text(text: impl Into<String>) -> Self {
Self {
content_type: "text".to_string(),
text: Some(text.into()),
data: None,
mime_type: None,
uri: None,
base64_data: None,
size_bytes: None,
sha256: None,
width: None,
height: None,
duration_ms: None,
sample_rate: None,
channels: None,
language: None,
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LLMMessage {
#[serde(default)]
pub role: String,
#[serde(default)]
pub parts: Vec<LLMContentPart>,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub tool_call_id: Option<String>,
#[serde(default)]
pub metadata: HashMap<String, Value>,
}
impl LLMMessage {
pub fn new(role: impl Into<String>, text: impl Into<String>) -> Self {
Self {
role: role.into(),
parts: vec![LLMContentPart::text(text)],
name: None,
tool_call_id: None,
metadata: HashMap::new(),
}
}
pub fn text(&self) -> Option<&str> {
self.parts.iter().find_map(|p| p.text.as_deref())
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ToolCallSpec {
#[serde(default)]
pub name: String,
#[serde(default)]
pub arguments_json: String,
#[serde(default)]
pub arguments: Option<Value>,
#[serde(default)]
pub call_id: Option<String>,
#[serde(default)]
pub index: Option<i32>,
#[serde(default)]
pub parent_call_id: Option<String>,
#[serde(default)]
pub metadata: HashMap<String, Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ToolCallResult {
#[serde(default)]
pub call_id: Option<String>,
#[serde(default)]
pub output_text: Option<String>,
#[serde(default)]
pub exit_code: Option<i32>,
#[serde(default)]
pub status: Option<String>,
#[serde(default)]
pub error_message: Option<String>,
#[serde(default)]
pub started_at: Option<DateTime<Utc>>,
#[serde(default)]
pub completed_at: Option<DateTime<Utc>>,
#[serde(default)]
pub duration_ms: Option<i32>,
#[serde(default)]
pub metadata: HashMap<String, Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LLMChunk {
#[serde(default)]
pub sequence_index: i32,
#[serde(default = "Utc::now")]
pub received_at: DateTime<Utc>,
#[serde(default)]
pub event_type: Option<String>,
#[serde(default)]
pub choice_index: Option<i32>,
#[serde(default)]
pub raw_json: Option<String>,
#[serde(default)]
pub delta_text: Option<String>,
#[serde(default)]
pub delta: Option<Value>,
#[serde(default)]
pub metadata: HashMap<String, Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct LLMCallRecord {
#[serde(default)]
pub call_id: String,
#[serde(default)]
pub api_type: String,
#[serde(default)]
pub provider: Option<String>,
#[serde(default)]
pub model_name: String,
#[serde(default)]
pub schema_version: Option<String>,
#[serde(default)]
pub started_at: Option<DateTime<Utc>>,
#[serde(default)]
pub completed_at: Option<DateTime<Utc>>,
#[serde(default)]
pub latency_ms: Option<i32>,
#[serde(default)]
pub request_params: LLMRequestParams,
#[serde(default)]
pub input_messages: Vec<LLMMessage>,
#[serde(default)]
pub input_text: Option<String>,
#[serde(default)]
pub tool_choice: Option<String>,
#[serde(default)]
pub output_messages: Vec<LLMMessage>,
#[serde(default)]
pub outputs: Vec<LLMMessage>,
#[serde(default)]
pub output_text: Option<String>,
#[serde(default)]
pub output_tool_calls: Vec<ToolCallSpec>,
#[serde(default)]
pub tool_results: Vec<ToolCallResult>,
#[serde(default)]
pub usage: Option<LLMUsage>,
#[serde(default)]
pub finish_reason: Option<String>,
#[serde(default)]
pub choice_index: Option<i32>,
#[serde(default)]
pub chunks: Option<Vec<LLMChunk>>,
#[serde(default)]
pub request_raw_json: Option<String>,
#[serde(default)]
pub response_raw_json: Option<String>,
#[serde(default)]
pub metadata: HashMap<String, Value>,
#[serde(default)]
pub provider_request_id: Option<String>,
#[serde(default)]
pub request_server_timing: Option<Value>,
#[serde(default)]
pub outcome: Option<String>,
#[serde(default)]
pub error: Option<Value>,
#[serde(default)]
pub token_traces: Option<Vec<Value>>,
#[serde(default)]
pub safety: Option<Value>,
#[serde(default)]
pub refusal: Option<Value>,
#[serde(default)]
pub redactions: Option<Vec<Value>>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct MarkovBlanketMessage {
#[serde(default)]
pub content: MessageContent,
#[serde(default)]
pub message_type: String,
#[serde(default)]
pub time_record: TimeRecord,
#[serde(default)]
pub metadata: HashMap<String, Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionTimeStep {
#[serde(default)]
pub step_id: String,
#[serde(default)]
pub step_index: i32,
#[serde(default = "Utc::now")]
pub timestamp: DateTime<Utc>,
#[serde(default)]
pub turn_number: Option<i32>,
#[serde(default)]
pub events: Vec<TracingEvent>,
#[serde(default)]
pub markov_blanket_messages: Vec<MarkovBlanketMessage>,
#[serde(default)]
pub step_metadata: HashMap<String, Value>,
#[serde(default)]
pub completed_at: Option<DateTime<Utc>>,
}
impl SessionTimeStep {
pub fn new(step_id: impl Into<String>, step_index: i32) -> Self {
Self {
step_id: step_id.into(),
step_index,
timestamp: Utc::now(),
turn_number: None,
events: Vec::new(),
markov_blanket_messages: Vec::new(),
step_metadata: HashMap::new(),
completed_at: None,
}
}
pub fn complete(&mut self) {
self.completed_at = Some(Utc::now());
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionTrace {
#[serde(default)]
pub session_id: String,
#[serde(default = "Utc::now")]
pub created_at: DateTime<Utc>,
#[serde(default)]
pub session_time_steps: Vec<SessionTimeStep>,
#[serde(default)]
pub event_history: Vec<TracingEvent>,
#[serde(default)]
pub markov_blanket_message_history: Vec<MarkovBlanketMessage>,
#[serde(default)]
pub metadata: HashMap<String, Value>,
}
impl SessionTrace {
pub fn new(session_id: impl Into<String>) -> Self {
Self {
session_id: session_id.into(),
created_at: Utc::now(),
session_time_steps: Vec::new(),
event_history: Vec::new(),
markov_blanket_message_history: Vec::new(),
metadata: HashMap::new(),
}
}
pub fn num_timesteps(&self) -> usize {
self.session_time_steps.len()
}
pub fn num_events(&self) -> usize {
self.event_history.len()
}
pub fn num_messages(&self) -> usize {
self.markov_blanket_message_history.len()
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct OutcomeReward {
#[serde(default = "default_objective_key")]
pub objective_key: String,
pub total_reward: f64,
#[serde(default)]
pub achievements_count: i32,
#[serde(default)]
pub total_steps: i32,
#[serde(default)]
pub reward_metadata: HashMap<String, Value>,
#[serde(default)]
pub annotation: Option<Value>,
}
fn default_objective_key() -> String {
"reward".to_string()
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct EventReward {
#[serde(default = "default_objective_key")]
pub objective_key: String,
pub reward_value: f64,
#[serde(default)]
pub reward_type: Option<String>,
#[serde(default)]
pub key: Option<String>,
#[serde(default)]
pub annotation: Option<Value>,
#[serde(default)]
pub source: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_time_record() {
let tr = TimeRecord::now();
assert!(tr.event_time > 0.0);
assert!(tr.message_time.is_none());
}
#[test]
fn test_message_content() {
let mc = MessageContent::from_text("hello");
assert_eq!(mc.as_text(), Some("hello"));
let mc = MessageContent::from_json(&serde_json::json!({"key": "value"}));
assert!(mc.json_payload.is_some());
}
#[test]
fn test_event_serialization() {
let event = TracingEvent::Cais(LMCAISEvent {
base: BaseEventFields::new("test-system"),
model_name: "gpt-4".to_string(),
provider: Some("openai".to_string()),
input_tokens: Some(100),
output_tokens: Some(50),
..Default::default()
});
let json = serde_json::to_string(&event).unwrap();
assert!(json.contains("cais"));
assert!(json.contains("gpt-4"));
let parsed: TracingEvent = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.event_type(), EventType::Cais);
}
#[test]
fn test_session_trace() {
let mut trace = SessionTrace::new("test-session");
assert_eq!(trace.num_timesteps(), 0);
let step = SessionTimeStep::new("step-1", 0);
trace.session_time_steps.push(step);
assert_eq!(trace.num_timesteps(), 1);
}
#[test]
fn test_llm_message() {
let msg = LLMMessage::new("user", "Hello, world!");
assert_eq!(msg.role, "user");
assert_eq!(msg.text(), Some("Hello, world!"));
}
}