lean_ctx/proxy/
effort_routing.rs1use serde_json::Value;
39use std::sync::atomic::{AtomicU64, Ordering};
40
41use crate::core::config::Effort;
42
43#[derive(Debug, Clone, Copy, PartialEq, Eq)]
45pub enum TurnClass {
46 Routine,
48 Full,
50}
51
52#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
54pub struct RoutingStats {
55 pub routine_count: u64,
56 pub full_count: u64,
57}
58
59static ROUTINE_COUNT: AtomicU64 = AtomicU64::new(0);
60static FULL_COUNT: AtomicU64 = AtomicU64::new(0);
61
62pub fn classify_turn(messages: &Value) -> TurnClass {
66 let Some(arr) = messages.as_array() else {
67 return TurnClass::Full;
68 };
69
70 if arr.is_empty() {
71 return TurnClass::Full;
72 }
73
74 let last = &arr[arr.len() - 1];
75 let role = last.get("role").and_then(Value::as_str).unwrap_or("");
76
77 if role == "tool" {
78 classify_tool_result(last, arr)
79 } else {
80 TurnClass::Full
81 }
82}
83
84pub fn classify_turn_responses(input: &Value) -> TurnClass {
86 let Some(arr) = input.as_array() else {
87 return TurnClass::Full;
88 };
89 if arr.is_empty() {
90 return TurnClass::Full;
91 }
92
93 let last = &arr[arr.len() - 1];
95 let item_type = last.get("type").and_then(Value::as_str).unwrap_or("");
96
97 if item_type == "function_call_output" {
98 let output = last.get("output").and_then(Value::as_str).unwrap_or("");
99 if is_routine_tool_output(output) {
100 return TurnClass::Routine;
101 }
102 }
103
104 TurnClass::Full
105}
106
107pub fn classify_turn_anthropic(messages: &Value) -> TurnClass {
109 let Some(arr) = messages.as_array() else {
110 return TurnClass::Full;
111 };
112 if arr.is_empty() {
113 return TurnClass::Full;
114 }
115
116 let last = &arr[arr.len() - 1];
117 let role = last.get("role").and_then(Value::as_str).unwrap_or("");
118
119 if role != "user" {
120 return TurnClass::Full;
121 }
122
123 let content = last.get("content");
125 if let Some(Value::Array(blocks)) = content {
126 let all_tool_results = !blocks.is_empty()
127 && blocks
128 .iter()
129 .all(|b| b.get("type").and_then(Value::as_str) == Some("tool_result"));
130
131 if all_tool_results {
132 let has_errors = blocks.iter().any(|b| {
134 b.get("is_error") == Some(&Value::Bool(true))
135 || b.get("content")
136 .and_then(|c| c.as_str().or_else(|| extract_text_from_content(c)))
137 .is_some_and(contains_error_indicators)
138 });
139
140 if has_errors {
141 return TurnClass::Full;
142 }
143
144 if blocks.len() > 3 {
146 return TurnClass::Full;
147 }
148
149 let all_routine = blocks.iter().all(|b| {
151 let text = b
152 .get("content")
153 .and_then(|c| c.as_str().or_else(|| extract_text_from_content(c)))
154 .unwrap_or("");
155 is_routine_tool_output(text)
156 });
157
158 if all_routine {
159 return TurnClass::Routine;
160 }
161 }
162 }
163
164 TurnClass::Full
165}
166
167pub fn effort_for_turn(class: TurnClass, base: Effort) -> Effort {
170 match class {
171 TurnClass::Routine => {
172 ROUTINE_COUNT.fetch_add(1, Ordering::Relaxed);
173 Effort::Minimal
175 }
176 TurnClass::Full => {
177 FULL_COUNT.fetch_add(1, Ordering::Relaxed);
178 base
179 }
180 }
181}
182
183pub fn stats() -> RoutingStats {
185 RoutingStats {
186 routine_count: ROUTINE_COUNT.load(Ordering::Relaxed),
187 full_count: FULL_COUNT.load(Ordering::Relaxed),
188 }
189}
190
191fn classify_tool_result(msg: &Value, _all_messages: &[Value]) -> TurnClass {
196 let content = msg.get("content").and_then(Value::as_str).unwrap_or("");
197
198 if contains_error_indicators(content) {
199 return TurnClass::Full;
200 }
201
202 if is_routine_tool_output(content) {
203 return TurnClass::Routine;
204 }
205
206 TurnClass::Full
207}
208
209fn is_routine_tool_output(content: &str) -> bool {
211 if content.is_empty() || content.len() < 10 {
212 return false;
213 }
214
215 if contains_error_indicators(content) {
217 return false;
218 }
219
220 if content.len() > 8000 {
222 return false;
223 }
224
225 let routine_signals = [
227 "deps ", "[unchanged", "[lean-ctx]", "lines:", "exit_code: 0",
234 "Command completed",
235 "0 errors",
236 "All tests passed",
237 "no changes",
238 "nothing to commit",
239 "Already up to date",
240 "Build succeeded",
241 "matches in",
243 "0 matches",
244 ];
245
246 routine_signals.iter().any(|sig| content.contains(sig))
247}
248
249fn contains_error_indicators(content: &str) -> bool {
251 let lower = content.to_ascii_lowercase();
252 let indicators = [
253 "error",
254 "failed",
255 "failure",
256 "fatal",
257 "panic",
258 "exception",
259 "traceback",
260 "stack trace",
261 "segfault",
262 "abort",
263 "denied",
264 "permission",
265 "not found",
266 "timed out",
267 "exit_code: 1",
268 "exit code 1",
269 "compilation error",
270 "syntax error",
271 "type error",
272 ];
273
274 indicators.iter().any(|ind| lower.contains(ind))
275}
276
277fn extract_text_from_content(content: &Value) -> Option<&str> {
279 if let Some(arr) = content.as_array() {
280 for block in arr {
281 if block.get("type").and_then(Value::as_str) == Some("text") {
282 return block.get("text").and_then(Value::as_str);
283 }
284 }
285 }
286 None
287}
288
289#[cfg(test)]
290mod tests {
291 use super::*;
292 use serde_json::json;
293
294 #[test]
295 fn user_message_is_always_full() {
296 let messages = json!([
297 {"role": "user", "content": "What does this function do?"}
298 ]);
299 assert_eq!(classify_turn(&messages), TurnClass::Full);
300 }
301
302 #[test]
303 fn successful_file_read_is_routine() {
304 let messages = json!([
305 {"role": "assistant", "content": "Let me read that file."},
306 {"role": "tool", "content": "main.rs 50L\n deps serde\n[lean-ctx] full source: ..."}
307 ]);
308 assert_eq!(classify_turn(&messages), TurnClass::Routine);
309 }
310
311 #[test]
312 fn error_tool_result_is_full() {
313 let messages = json!([
314 {"role": "tool", "content": "error[E0308]: mismatched types\n --> src/main.rs:5:12"}
315 ]);
316 assert_eq!(classify_turn(&messages), TurnClass::Full);
317 }
318
319 #[test]
320 fn successful_shell_is_routine() {
321 let messages = json!([
322 {"role": "tool", "content": "Command completed in 150ms\nexit_code: 0\nAll tests passed"}
323 ]);
324 assert_eq!(classify_turn(&messages), TurnClass::Routine);
325 }
326
327 #[test]
328 fn anthropic_tool_result_routine() {
329 let messages = json!([
330 {"role": "user", "content": [
331 {"type": "tool_result", "tool_use_id": "abc", "content": "[unchanged 5L]\n[lean-ctx] cached"}
332 ]}
333 ]);
334 assert_eq!(classify_turn_anthropic(&messages), TurnClass::Routine);
335 }
336
337 #[test]
338 fn anthropic_tool_result_with_error() {
339 let messages = json!([
340 {"role": "user", "content": [
341 {"type": "tool_result", "tool_use_id": "abc", "is_error": true, "content": "Tool failed"}
342 ]}
343 ]);
344 assert_eq!(classify_turn_anthropic(&messages), TurnClass::Full);
345 }
346
347 #[test]
348 fn effort_mapping() {
349 assert_eq!(
350 effort_for_turn(TurnClass::Routine, Effort::High),
351 Effort::Minimal
352 );
353 assert_eq!(effort_for_turn(TurnClass::Full, Effort::High), Effort::High);
354 assert_eq!(
355 effort_for_turn(TurnClass::Full, Effort::Medium),
356 Effort::Medium
357 );
358 }
359
360 #[test]
361 fn empty_messages_is_full() {
362 assert_eq!(classify_turn(&json!([])), TurnClass::Full);
363 assert_eq!(classify_turn(&json!(null)), TurnClass::Full);
364 }
365
366 #[test]
367 fn deterministic_classification() {
368 let messages = json!([
369 {"role": "tool", "content": "Build succeeded\nexit_code: 0\nCommand completed in 2s"}
370 ]);
371 let c1 = classify_turn(&messages);
372 let c2 = classify_turn(&messages);
373 assert_eq!(c1, c2, "classification must be deterministic");
374 }
375}