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