use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum MeasurementSource {
Native { provider: String },
LocalExact { tokenizer: String },
Postflight,
Heuristic,
HostProvided,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct TokenMeasurement {
pub fingerprint: String,
pub tokens: u32,
pub source: MeasurementSource,
pub confidence: MeasurementConfidence,
}
impl TokenMeasurement {
pub fn for_message(message: &crate::types::message::CoreMessage, tokens: u32) -> Self {
use sha2::{Digest as _, Sha256};
let digest = Sha256::digest(serde_json::to_vec(message).expect("message is serializable"));
let hex = digest
.iter()
.map(|b| format!("{b:02x}"))
.collect::<String>();
Self {
fingerprint: format!("sha256:{hex}"),
tokens,
source: MeasurementSource::HostProvided,
confidence: MeasurementConfidence::HighConfidence,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ToolMeasurement {
pub call_id: String,
pub tokens: u32,
}
impl ToolMeasurement {
pub fn new(call_id: impl Into<String>, tokens: u32) -> Self {
Self {
call_id: call_id.into(),
tokens,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MeasurementConfidence {
Exact,
HighConfidence,
LowConfidence,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct PromptMeasurement {
pub input_tokens: u32,
pub source: MeasurementSource,
pub confidence: MeasurementConfidence,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn constructs_all_three_measurement_source_variants() {
let native = MeasurementSource::Native {
provider: "anthropic".to_string(),
};
let local_exact = MeasurementSource::LocalExact {
tokenizer: "cl100k_base".to_string(),
};
let heuristic = MeasurementSource::Heuristic;
assert_ne!(native, local_exact);
assert_ne!(local_exact, heuristic);
}
#[test]
fn prompt_measurement_round_trips_through_json() {
let m = PromptMeasurement {
input_tokens: 1234,
source: MeasurementSource::Native {
provider: "openai".to_string(),
},
confidence: MeasurementConfidence::Exact,
};
let json = serde_json::to_string(&m).unwrap();
let back: PromptMeasurement = serde_json::from_str(&json).unwrap();
assert_eq!(m, back);
}
#[test]
fn postflight_source_round_trips_and_differs_from_preflight_kinds() {
let m = PromptMeasurement {
input_tokens: 1000,
source: MeasurementSource::Postflight,
confidence: MeasurementConfidence::Exact,
};
let json = serde_json::to_string(&m).unwrap();
let back: PromptMeasurement = serde_json::from_str(&json).unwrap();
assert_eq!(m, back);
assert_eq!(
serde_json::to_string(&m.source).unwrap(),
r#"{"kind":"postflight"}"#
);
}
#[test]
fn unknown_field_is_rejected() {
let raw = r#"{"input_tokens": 10, "source": {"kind": "heuristic"}, "confidence": "exact", "extra": true}"#;
let result: Result<PromptMeasurement, _> = serde_json::from_str(raw);
assert!(
result.is_err(),
"deny_unknown_fields must reject stray keys"
);
}
#[test]
fn tool_measurement_is_independent_from_tool_result_state() {
let measurement = ToolMeasurement::new("call-1", 42);
let json = serde_json::to_string(&measurement).unwrap();
let decoded: ToolMeasurement = serde_json::from_str(&json).unwrap();
assert_eq!(decoded, measurement);
}
}