use std::fmt;
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::{
errors::GitError,
hash::ObjectHash,
internal::object::{
ObjectTrait,
types::{ActorRef, Header, ObjectType},
},
};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(deny_unknown_fields)]
pub struct RunUsage {
#[serde(flatten)]
header: Header,
run_id: Uuid,
input_tokens: u64,
output_tokens: u64,
total_tokens: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
cost_usd: Option<f64>,
}
impl RunUsage {
pub fn new(
created_by: ActorRef,
run_id: Uuid,
input_tokens: u64,
output_tokens: u64,
cost_usd: Option<f64>,
) -> Result<Self, String> {
Ok(Self {
header: Header::new(ObjectType::RunUsage, created_by)?,
run_id,
input_tokens,
output_tokens,
total_tokens: input_tokens + output_tokens,
cost_usd,
})
}
pub fn header(&self) -> &Header {
&self.header
}
pub fn run_id(&self) -> Uuid {
self.run_id
}
pub fn input_tokens(&self) -> u64 {
self.input_tokens
}
pub fn output_tokens(&self) -> u64 {
self.output_tokens
}
pub fn total_tokens(&self) -> u64 {
self.total_tokens
}
pub fn cost_usd(&self) -> Option<f64> {
self.cost_usd
}
pub fn is_consistent(&self) -> bool {
self.total_tokens == self.input_tokens + self.output_tokens
}
}
impl fmt::Display for RunUsage {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "RunUsage: {}", self.header.object_id())
}
}
impl ObjectTrait for RunUsage {
fn from_bytes(data: &[u8], _hash: ObjectHash) -> Result<Self, GitError>
where
Self: Sized,
{
serde_json::from_slice(data).map_err(|e| GitError::InvalidObjectInfo(e.to_string()))
}
fn get_type(&self) -> ObjectType {
ObjectType::RunUsage
}
fn get_size(&self) -> usize {
match serde_json::to_vec(self) {
Ok(v) => v.len(),
Err(e) => {
tracing::warn!("failed to compute RunUsage size: {}", e);
0
}
}
}
fn to_data(&self) -> Result<Vec<u8>, GitError> {
serde_json::to_vec(self).map_err(|e| GitError::InvalidObjectInfo(e.to_string()))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_run_usage_fields() {
let actor = ActorRef::agent("planner").expect("actor");
let usage = RunUsage::new(actor, Uuid::from_u128(0x1), 100, 40, Some(0.12)).expect("usage");
assert_eq!(usage.input_tokens(), 100);
assert_eq!(usage.output_tokens(), 40);
assert_eq!(usage.total_tokens(), 140);
assert!(usage.is_consistent());
assert_eq!(usage.cost_usd(), Some(0.12));
}
}