返回 CodeWhale
token_estimate_cache.rs
根目录 / crates / tui / src / core / engine / token_estimate_cache.rs
1 //! Process-local memoization for [`crate::compaction::estimate_input_tokens_for_pressure`].
2 //!
3 //! This is the engine's one pressure number: the same un-inflated estimate the
4 //! auto-compaction gate, the compaction preflight, the TUI context meter and
5 //! the `/context` headline read, so compaction receipts and the context-budget
6 //! snapshot agree with them (0.10.1 item 9). The 1.5x-inflated
7 //! `estimate_input_tokens_conservative` is request-overflow protection only
8 //! and is deliberately not cached here.
9 //!
10 //! The token estimator walks the full [`codewhale_models::Message`] history and the
11 //! active system prompt, which is by far the most expensive per-turn CPU cost
12 //! in the engine hot path. The same input data is queried from at least five
13 //! sites per turn: capacity pre/post tool checkpoints, error escalation,
14 //! the seam manager, and the trimmed-message budget check, plus four more
15 //! from the TUI footer, `/status`, `/debug`, and the context inspector.
16 //!
17 //! Without memoization, a 200-message history with 5 KB of tool results costs
18 //! ~2 ms per call; that is 20 ms of pure waste on a single turn. The estimator
19 //! itself is a pure function of `(messages, system_prompt)`, so a
20 //! content-versioned cache is safe: the caller bumps `messages_revision`
21 //! on every mutation, and we also include a fast fingerprint of the system
22 //! prompt as part of the key.
23 //!
24 //! The cache is process-local only — cross-session persistence is intentionally
25 //! out of scope (see PR #2520 for the cross-session prompt-base disk cache).
26
27 use std::collections::hash_map::DefaultHasher;
28 use std::hash::{Hash, Hasher};
29
30 use crate::compaction::estimate_input_tokens_for_pressure;
31 use codewhale_models::{Message, SystemPrompt};
32
33 /// Default capacity for the rolling audit ring. Sized so a 64-entry window
34 /// covers a full capacity controller observation cycle without unbounded
35 /// growth on long-running sessions.
36 const AUDIT_RING_CAPACITY: usize = 64;
37
38 /// Process-local memoization for `estimate_input_tokens_for_pressure`.
39 ///
40 /// The cache is keyed on the `(messages_revision, system_fingerprint)`
41 /// pair, both of which the engine bumps on every content change. On a hit
42 /// the previously stored token estimate is returned without re-walking the
43 /// message list. On a miss, the estimator runs and the result is stored
44 /// alongside the audit ring entry.
45 #[derive(Debug, Default, Clone)]
46 pub struct TokenEstimateCache {
47 /// Monotonic counter bumped by the engine on every message mutation.
48 messages_revision: u64,
49 /// Stable 64-bit hash of the current system prompt text. Computed once
50 /// per `lookup_or_compute` call when the cache misses.
51 system_fingerprint: u64,
52 /// Cached token count, valid iff both keys match the current inputs.
53 cached_tokens: Option<usize>,
54 /// Audit ring of recent (revision, tokens) pairs. The most recent entry
55 /// is the tail; the oldest is dropped when capacity is exceeded. Used by
56 /// observability to surface cache effectiveness to `/status`.
57 audit_ring: Vec<(u64, usize)>,
58 /// Number of cache hits since the cache was last cleared. Saturates at
59 /// `u64::MAX` (effectively never in practice).
60 hits: u64,
61 /// Number of cache misses since the cache was last cleared.
62 misses: u64,
63 }
64
65 impl TokenEstimateCache {
66 /// Construct a fresh, empty cache. `messages_revision` defaults to 0; the
67 /// engine must call [`bump_messages_revision`](Self::bump_messages_revision)
68 /// whenever a mutation occurs so the next lookup correctly invalidates.
69 #[must_use]
70 pub fn new() -> Self {
71 Self::default()
72 }
73
74 /// Returns the cached token estimate, recomputing on miss.
75 ///
76 /// `messages_revision` is the engine's monotonic counter; bump it on
77 /// every add/remove/clear. `system_prompt` may be `None`. `messages` is
78 /// borrowed for the duration of the call so a miss can re-tokenize.
79 pub fn lookup_or_compute(
80 &mut self,
81 messages_revision: u64,
82 system_prompt: Option<&SystemPrompt>,
83 messages: &[Message],
84 ) -> usize {
85 let system_fingerprint = fingerprint_system_prompt(system_prompt);
86
87 if self.messages_revision == messages_revision
88 && self.system_fingerprint == system_fingerprint
89 && let Some(tokens) = self.cached_tokens
90 {
91 self.hits = self.hits.saturating_add(1);
92 return tokens;
93 }
94
95 let tokens = estimate_input_tokens_for_pressure(messages, system_prompt);
96 self.messages_revision = messages_revision;
97 self.system_fingerprint = system_fingerprint;
98 self.cached_tokens = Some(tokens);
99 self.misses = self.misses.saturating_add(1);
100 self.push_audit(messages_revision, tokens);
101 tokens
102 }
103
104 /// Record a messages-revision bump. The engine calls this whenever
105 /// `session.messages` is mutated. Calling it with a value smaller than
106 /// the current value is a no-op (the cache is monotonic).
107 #[allow(dead_code)] // exposed for future wiring of /clear and reset paths; tests exercise it
108 pub fn bump_messages_revision(&mut self, revision: u64) {
109 if revision > self.messages_revision {
110 self.messages_revision = revision;
111 self.cached_tokens = None;
112 }
113 }
114
115 /// Forget all cached state. Used by `/clear` and session reset paths.
116 #[allow(dead_code)] // exposed for future wiring of /clear and reset paths; tests exercise it
117 pub fn invalidate(&mut self) {
118 self.cached_tokens = None;
119 self.system_fingerprint = 0;
120 self.audit_ring.clear();
121 self.hits = 0;
122 self.misses = 0;
123 }
124
125 /// Returns `(hits, misses)` counters since the last `invalidate` call.
126 #[allow(dead_code)] // surfaced via /status in a follow-up; tests exercise it
127 #[must_use]
128 pub fn stats(&self) -> (u64, u64) {
129 (self.hits, self.misses)
130 }
131
132 /// Returns the most recent `(revision, tokens)` audit entries, newest
133 /// first. Bounded by [`AUDIT_RING_CAPACITY`].
134 #[allow(dead_code)] // surfaced via /status in a follow-up; tests exercise it
135 #[must_use]
136 pub fn recent_audit(&self) -> &[(u64, usize)] {
137 &self.audit_ring
138 }
139
140 fn push_audit(&mut self, revision: u64, tokens: usize) {
141 if self.audit_ring.len() >= AUDIT_RING_CAPACITY {
142 self.audit_ring.remove(0);
143 }
144 self.audit_ring.push((revision, tokens));
145 }
146 }
147
148 /// Stable 64-bit hash of the system prompt text. Walks the same shape the
149 /// estimator consumes: a `Text` variant or a list of `Blocks`. Returns 0
150 /// for `None` so the empty case is distinguishable but cheap to compare.
151 fn fingerprint_system_prompt(system: Option<&SystemPrompt>) -> u64 {
152 let Some(system) = system else {
153 return 0;
154 };
155 let mut hasher = DefaultHasher::new();
156 match system {
157 SystemPrompt::Text(text) => {
158 "text".hash(&mut hasher);
159 text.hash(&mut hasher);
160 }
161 SystemPrompt::Blocks(blocks) => {
162 "blocks".hash(&mut hasher);
163 blocks.len().hash(&mut hasher);
164 for block in blocks {
165 block.block_type.hash(&mut hasher);
166 block.text.hash(&mut hasher);
167 }
168 }
169 }
170 hasher.finish()
171 }
172
173 #[cfg(test)]
174 mod tests {
175 use super::*;
176 use codewhale_models::Role;
177 use codewhale_models::{ContentBlock, SystemBlock};
178
179 fn user_text(s: &str) -> Message {
180 Message {
181 role: Role::User,
182 content: vec![ContentBlock::Text {
183 text: s.to_string(),
184 cache_control: None,
185 }],
186 }
187 }
188
189 fn sys_text(s: &str) -> SystemPrompt {
190 SystemPrompt::Text(s.to_string())
191 }
192
193 #[test]
194 fn first_call_is_a_miss() {
195 let mut cache = TokenEstimateCache::new();
196 let messages = vec![user_text("hello world")];
197 let tokens = cache.lookup_or_compute(1, None, &messages);
198 let (hits, misses) = cache.stats();
199 assert!(tokens > 0);
200 assert_eq!(hits, 0);
201 assert_eq!(misses, 1);
202 }
203
204 #[test]
205 fn repeated_call_with_same_revision_is_a_hit() {
206 let mut cache = TokenEstimateCache::new();
207 let messages = vec![user_text("hello world")];
208 let _ = cache.lookup_or_compute(1, None, &messages);
209 let _ = cache.lookup_or_compute(1, None, &messages);
210 let (hits, misses) = cache.stats();
211 assert_eq!(hits, 1);
212 assert_eq!(misses, 1);
213 }
214
215 #[test]
216 fn revision_bump_invalidates() {
217 let mut cache = TokenEstimateCache::new();
218 let messages = vec![user_text("hi")];
219 let a = cache.lookup_or_compute(1, None, &messages);
220 let b = cache.lookup_or_compute(2, None, &messages);
221 let (hits, misses) = cache.stats();
222 // Both calls were misses (different revisions), neither hit the cache.
223 assert_eq!(a, b);
224 assert_eq!(hits, 0);
225 assert_eq!(misses, 2);
226 }
227
228 #[test]
229 fn system_prompt_change_invalidates() {
230 let mut cache = TokenEstimateCache::new();
231 let messages = vec![user_text("hi")];
232 let _ = cache.lookup_or_compute(1, Some(&sys_text("alpha")), &messages);
233 let _ = cache.lookup_or_compute(1, Some(&sys_text("beta")), &messages);
234 let (hits, misses) = cache.stats();
235 assert_eq!(hits, 0);
236 assert_eq!(misses, 2);
237 }
238
239 #[test]
240 fn bump_messages_revision_clears_cache() {
241 let mut cache = TokenEstimateCache::new();
242 let messages = vec![user_text("x")];
243 let _ = cache.lookup_or_compute(1, None, &messages);
244 cache.bump_messages_revision(2);
245 let _ = cache.lookup_or_compute(2, None, &messages);
246 let (hits, misses) = cache.stats();
247 assert_eq!(hits, 0);
248 assert_eq!(misses, 2);
249 }
250
251 #[test]
252 fn bump_to_smaller_revision_is_noop() {
253 let mut cache = TokenEstimateCache::new();
254 let messages = vec![user_text("x")];
255 let _ = cache.lookup_or_compute(5, None, &messages);
256 cache.bump_messages_revision(2);
257 // revision went down, cache should still be valid for revision 5
258 let _ = cache.lookup_or_compute(5, None, &messages);
259 let (hits, _) = cache.stats();
260 assert_eq!(hits, 1, "downward revision bumps must not invalidate");
261 }
262
263 #[test]
264 fn invalidate_resets_state() {
265 let mut cache = TokenEstimateCache::new();
266 let messages = vec![user_text("x")];
267 let _ = cache.lookup_or_compute(1, None, &messages);
268 let _ = cache.lookup_or_compute(1, None, &messages);
269 cache.invalidate();
270 let (hits, misses) = cache.stats();
271 assert_eq!(hits, 0);
272 assert_eq!(misses, 0);
273 }
274
275 #[test]
276 fn blocks_system_prompt_yields_distinct_fingerprint() {
277 let blocks_a = SystemPrompt::Blocks(vec![SystemBlock {
278 block_type: "text".to_string(),
279 text: "alpha".to_string(),
280 cache_control: None,
281 }]);
282 let blocks_b = SystemPrompt::Blocks(vec![SystemBlock {
283 block_type: "text".to_string(),
284 text: "beta".to_string(),
285 cache_control: None,
286 }]);
287 let mut cache = TokenEstimateCache::new();
288 let messages = vec![user_text("hi")];
289 let _ = cache.lookup_or_compute(1, Some(&blocks_a), &messages);
290 let _ = cache.lookup_or_compute(1, Some(&blocks_b), &messages);
291 let (hits, misses) = cache.stats();
292 assert_eq!(hits, 0);
293 assert_eq!(misses, 2);
294 }
295
296 #[test]
297 fn audit_ring_records_recent_pairs() {
298 let mut cache = TokenEstimateCache::new();
299 let messages = vec![user_text("hi")];
300 for rev in 1..=5 {
301 let _ = cache.lookup_or_compute(rev, None, &messages);
302 }
303 let ring = cache.recent_audit();
304 assert_eq!(ring.len(), 5);
305 assert_eq!(ring.last().copied(), Some((5, ring.last().unwrap().1)));
306 }
307
308 #[test]
309 fn audit_ring_bounded_by_capacity() {
310 let mut cache = TokenEstimateCache::new();
311 let messages = vec![user_text("hi")];
312 for rev in 1..=(AUDIT_RING_CAPACITY + 10) as u64 {
313 let _ = cache.lookup_or_compute(rev, None, &messages);
314 }
315 let ring = cache.recent_audit();
316 assert_eq!(ring.len(), AUDIT_RING_CAPACITY);
317 // newest entry should be the most recent revision we asked for
318 assert_eq!(ring.last().unwrap().0, (AUDIT_RING_CAPACITY + 10) as u64);
319 }
320 }
321
321 lines RUST