deepstrike_core/context/
measurement.rs1use serde::{Deserialize, Serialize};
9
10#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
14#[serde(tag = "kind", rename_all = "snake_case")]
15pub enum MeasurementSource {
16 Native { provider: String },
19 LocalExact { tokenizer: String },
23 Postflight,
26 Heuristic,
28 HostProvided,
30}
31
32#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
35#[serde(deny_unknown_fields)]
36pub struct TokenMeasurement {
37 pub fingerprint: String,
38 pub tokens: u32,
39 pub source: MeasurementSource,
40 pub confidence: MeasurementConfidence,
41}
42
43impl TokenMeasurement {
44 pub fn for_message(message: &crate::types::message::CoreMessage, tokens: u32) -> Self {
45 Self {
46 fingerprint: Self::message_fingerprint(message),
47 tokens,
48 source: MeasurementSource::HostProvided,
49 confidence: MeasurementConfidence::HighConfidence,
50 }
51 }
52
53 pub fn matches_message(&self, message: &crate::types::message::CoreMessage) -> bool {
54 self.fingerprint == Self::message_fingerprint(message)
55 }
56
57 fn message_fingerprint(message: &crate::types::message::CoreMessage) -> String {
58 use sha2::{Digest as _, Sha256};
59 let material =
60 super::execution::message_material(message, &crate::mm::handle::HandleTable::new());
61 let digest =
62 Sha256::digest(serde_json::to_vec(&material).expect("message is serializable"));
63 let hex = digest
64 .iter()
65 .map(|b| format!("{b:02x}"))
66 .collect::<String>();
67 format!("sha256:{hex}")
68 }
69}
70
71#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
74#[serde(deny_unknown_fields)]
75pub struct ToolMeasurement {
76 pub call_id: String,
77 pub tokens: u32,
78}
79
80impl ToolMeasurement {
81 pub fn new(call_id: impl Into<String>, tokens: u32) -> Self {
82 Self {
83 call_id: call_id.into(),
84 tokens,
85 }
86 }
87}
88
89#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
94#[serde(rename_all = "snake_case")]
95pub enum MeasurementConfidence {
96 Exact,
97 HighConfidence,
98 LowConfidence,
99}
100
101#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
103#[serde(deny_unknown_fields)]
104pub struct PromptMeasurement {
105 pub input_tokens: u32,
106 pub source: MeasurementSource,
107 pub confidence: MeasurementConfidence,
108}
109
110#[cfg(test)]
111mod tests {
112 use super::*;
113
114 #[test]
115 fn constructs_all_three_measurement_source_variants() {
116 let native = MeasurementSource::Native {
117 provider: "anthropic".to_string(),
118 };
119 let local_exact = MeasurementSource::LocalExact {
120 tokenizer: "cl100k_base".to_string(),
121 };
122 let heuristic = MeasurementSource::Heuristic;
123
124 assert_ne!(native, local_exact);
125 assert_ne!(local_exact, heuristic);
126 }
127
128 #[test]
129 fn prompt_measurement_round_trips_through_json() {
130 let m = PromptMeasurement {
131 input_tokens: 1234,
132 source: MeasurementSource::Native {
133 provider: "openai".to_string(),
134 },
135 confidence: MeasurementConfidence::Exact,
136 };
137 let json = serde_json::to_string(&m).unwrap();
138 let back: PromptMeasurement = serde_json::from_str(&json).unwrap();
139 assert_eq!(m, back);
140 }
141
142 #[test]
143 fn postflight_source_round_trips_and_differs_from_preflight_kinds() {
144 let m = PromptMeasurement {
146 input_tokens: 1000,
147 source: MeasurementSource::Postflight,
148 confidence: MeasurementConfidence::Exact,
149 };
150 let json = serde_json::to_string(&m).unwrap();
151 let back: PromptMeasurement = serde_json::from_str(&json).unwrap();
152 assert_eq!(m, back);
153 assert_eq!(
154 serde_json::to_string(&m.source).unwrap(),
155 r#"{"kind":"postflight"}"#
156 );
157 }
158
159 #[test]
160 fn unknown_field_is_rejected() {
161 let raw = r#"{"input_tokens": 10, "source": {"kind": "heuristic"}, "confidence": "exact", "extra": true}"#;
162 let result: Result<PromptMeasurement, _> = serde_json::from_str(raw);
163 assert!(
164 result.is_err(),
165 "deny_unknown_fields must reject stray keys"
166 );
167 }
168
169 #[test]
170 fn tool_measurement_is_independent_from_tool_result_state() {
171 let measurement = ToolMeasurement::new("call-1", 42);
172 let json = serde_json::to_string(&measurement).unwrap();
173 let decoded: ToolMeasurement = serde_json::from_str(&json).unwrap();
174 assert_eq!(decoded, measurement);
175 }
176}