1use std::cmp::Ordering;
5use std::collections::BTreeSet;
6use std::ops::Range;
7
8use super::measurement::TokenMeasurement;
9use super::token_engine::ContextTokenEngine;
10use super::units::unit_boundaries;
11use crate::lexical::{overlap_count, terms};
12use crate::types::message::{Content, ContentPart, CoreMessage};
13
14pub struct UtilitySelectionContext<'a> {
15 pub goal: &'a str,
16 pub criteria: &'a [String],
17 pub preserved_refs: &'a [String],
18 pub active_directives: &'a [String],
19}
20
21#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct UtilityUnitScore {
23 pub range: Range<usize>,
24 pub tokens: u32,
25 pub mandatory: bool,
26 pub goal_overlap: u32,
27 pub has_unresolved: bool,
28 pub referenced_later: bool,
29 pub is_error_or_decision: bool,
30 pub recency: u32,
31 pub token_cost: u32,
32 pub prefix_invalidation_cost: u32,
33 pub utility: i64,
34}
35
36#[derive(Debug, Clone, Default, PartialEq, Eq)]
37pub struct UtilityArchivePlan {
38 pub archived_ranges: Vec<Range<usize>>,
39 pub retained_ranges: Vec<Range<usize>>,
40 pub archived_tokens: u32,
41 pub retained_tokens: u32,
42 pub scores: Vec<UtilityUnitScore>,
43}
44
45pub fn plan_utility_archive(
51 messages: &[CoreMessage],
52 total_tokens: u32,
53 target_tokens: u32,
54 preserve_recent_units: usize,
55 engine: &ContextTokenEngine,
56 context: &UtilitySelectionContext<'_>,
57) -> UtilityArchivePlan {
58 plan_utility_archive_with_measurements(
59 messages,
60 &[],
61 total_tokens,
62 target_tokens,
63 preserve_recent_units,
64 engine,
65 context,
66 )
67}
68
69pub fn plan_utility_archive_with_measurements(
73 messages: &[CoreMessage],
74 measurements: &[TokenMeasurement],
75 total_tokens: u32,
76 target_tokens: u32,
77 preserve_recent_units: usize,
78 engine: &ContextTokenEngine,
79 context: &UtilitySelectionContext<'_>,
80) -> UtilityArchivePlan {
81 let ranges = unit_boundaries(messages);
82 if ranges.is_empty() {
83 return UtilityArchivePlan::default();
84 }
85 let unit_texts = ranges
86 .iter()
87 .map(|range| unit_text(&messages[range.clone()]))
88 .collect::<Vec<_>>();
89 let goal_terms = terms(
90 std::iter::once(context.goal)
91 .chain(context.criteria.iter().map(String::as_str))
92 .collect::<Vec<_>>()
93 .join(" ")
94 .as_str(),
95 );
96 let recent_start = ranges.len().saturating_sub(preserve_recent_units);
97 let denominator = total_tokens.max(1);
98 let unit_count = ranges.len().max(1) as u32;
99 let mut scores = Vec::with_capacity(ranges.len());
100
101 for (index, range) in ranges.iter().enumerate() {
102 let slice = &messages[range.clone()];
103 let text = &unit_texts[index];
104 let folded_text = text.to_lowercase();
105 let tokens = range
106 .clone()
107 .map(|message_index| {
108 measurements
109 .get(message_index)
110 .map(|measurement| measurement.tokens)
111 .unwrap_or_else(|| {
112 let message = &messages[message_index];
113 engine.count_message(message)
114 })
115 })
116 .sum::<u32>();
117 let goal_overlap = if goal_terms.is_empty() {
118 0
119 } else {
120 overlap_count(&terms(text), &goal_terms)
121 };
122 let has_unresolved = has_unresolved(slice, &folded_text);
123 let referenced_later = unit_referenced_later(slice, text, &unit_texts[index + 1..]);
124 let is_error_or_decision = is_error_or_decision(slice, &folded_text);
125 let dependency = context
126 .preserved_refs
127 .iter()
128 .any(|reference| contains_folded(text, reference))
129 || context
130 .active_directives
131 .iter()
132 .any(|directive| directive_dependency(text, directive));
133 let mandatory = index >= recent_start || has_unresolved || dependency;
134 let recency = ((index as u64 + 1) * 1_000 / u64::from(unit_count)) as u32;
135 let token_cost = (u64::from(tokens) * 1_000 / u64::from(denominator)) as u32;
136 let prefix_invalidation_cost =
137 ((ranges.len() - index) as u64 * 1_000 / u64::from(unit_count)) as u32;
138 let utility = i64::from(goal_overlap) * 4_000
139 + if has_unresolved { 20_000 } else { 0 }
140 + if referenced_later { 5_000 } else { 0 }
141 + if is_error_or_decision { 6_000 } else { 0 }
142 + i64::from(recency) * 2
143 - i64::from(token_cost) * 2
144 - i64::from(prefix_invalidation_cost);
145 scores.push(UtilityUnitScore {
146 range: range.clone(),
147 tokens,
148 mandatory,
149 goal_overlap,
150 has_unresolved,
151 referenced_later,
152 is_error_or_decision,
153 recency,
154 token_cost,
155 prefix_invalidation_cost,
156 utility,
157 });
158 }
159
160 if total_tokens <= target_tokens {
161 return UtilityArchivePlan {
162 archived_ranges: Vec::new(),
163 retained_ranges: ranges,
164 archived_tokens: 0,
165 retained_tokens: scores.iter().map(|score| score.tokens).sum(),
166 scores,
167 };
168 }
169
170 let mut retained = scores
171 .iter()
172 .enumerate()
173 .filter_map(|(index, score)| score.mandatory.then_some(index))
174 .collect::<BTreeSet<_>>();
175 let mut retained_tokens = retained
176 .iter()
177 .map(|index| scores[*index].tokens)
178 .sum::<u32>();
179 let mut optional = scores
180 .iter()
181 .enumerate()
182 .filter_map(|(index, score)| (!score.mandatory).then_some(index))
183 .collect::<Vec<_>>();
184 optional.sort_by(|left, right| compare_density(&scores[*right], &scores[*left]));
185 for index in optional {
186 let tokens = scores[index].tokens;
187 if retained_tokens.saturating_add(tokens) <= target_tokens {
188 retained.insert(index);
189 retained_tokens = retained_tokens.saturating_add(tokens);
190 }
191 }
192
193 let retained_ranges = ranges
194 .iter()
195 .enumerate()
196 .filter_map(|(index, range)| retained.contains(&index).then_some(range.clone()))
197 .collect::<Vec<_>>();
198 let archived_ranges = ranges
199 .iter()
200 .enumerate()
201 .filter_map(|(index, range)| (!retained.contains(&index)).then_some(range.clone()))
202 .collect::<Vec<_>>();
203 let archived_tokens = scores
204 .iter()
205 .enumerate()
206 .filter_map(|(index, score)| (!retained.contains(&index)).then_some(score.tokens))
207 .sum();
208 UtilityArchivePlan {
209 archived_ranges,
210 retained_ranges,
211 archived_tokens,
212 retained_tokens,
213 scores,
214 }
215}
216
217fn compare_density(left: &UtilityUnitScore, right: &UtilityUnitScore) -> Ordering {
218 let left_density = i128::from(left.utility) * i128::from(right.tokens.max(1));
219 let right_density = i128::from(right.utility) * i128::from(left.tokens.max(1));
220 left_density
221 .cmp(&right_density)
222 .then_with(|| left.utility.cmp(&right.utility))
223 .then_with(|| left.range.start.cmp(&right.range.start))
224}
225
226fn unit_text(messages: &[CoreMessage]) -> String {
227 let mut text = String::new();
228 let mut first_part = true;
229 for message in messages {
230 match &message.content {
231 Content::Text(content) => append_unit_part(&mut text, &mut first_part, content),
232 Content::Parts(content_parts) => {
233 for part in content_parts {
234 match part {
235 ContentPart::Text { text: content } => {
236 append_unit_part(&mut text, &mut first_part, content)
237 }
238 ContentPart::ToolResult {
239 call_id, output, ..
240 } => {
241 append_unit_part(&mut text, &mut first_part, call_id.as_str());
242 text.push(' ');
243 text.push_str(output);
244 }
245 ContentPart::Image { source, .. } => append_unit_part(
246 &mut text,
247 &mut first_part,
248 match source {
249 crate::types::durable_content::DurableSource::Url { url } => url,
250 _ => "[image]",
251 },
252 ),
253 ContentPart::Audio { .. } => {
254 append_unit_part(&mut text, &mut first_part, "audio")
255 }
256 }
257 }
258 }
259 }
260 for call in &message.tool_calls {
261 append_unit_part(&mut text, &mut first_part, call.id.as_str());
262 text.push(' ');
263 text.push_str(call.name.as_str());
264 text.push(' ');
265 text.push_str(&call.arguments.to_string());
266 }
267 }
268 text
269}
270
271fn append_unit_part(text: &mut String, first_part: &mut bool, part: &str) {
272 if !*first_part {
273 text.push('\n');
274 }
275 *first_part = false;
276 text.push_str(part);
277}
278
279fn contains_folded(text: &str, pattern: &str) -> bool {
280 !pattern.trim().is_empty() && text.to_lowercase().contains(&pattern.to_lowercase())
281}
282
283fn directive_dependency(text: &str, directive: &str) -> bool {
284 if contains_folded(text, directive) {
285 return true;
286 }
287 let directive_terms = terms(directive);
288 if directive_terms.is_empty() {
289 return false;
290 }
291 let threshold = directive_terms.len().min(2);
292 terms(text).intersection(&directive_terms).count() >= threshold
293}
294
295fn has_unresolved(messages: &[CoreMessage], folded_text: &str) -> bool {
296 let mut opened = BTreeSet::new();
297 let mut resolved = BTreeSet::new();
298 for message in messages {
299 for call in &message.tool_calls {
300 opened.insert(call.id.to_string());
301 }
302 if let Content::Parts(parts) = &message.content {
303 for part in parts {
304 if let ContentPart::ToolResult {
305 call_id, is_error, ..
306 } = part
307 {
308 if *is_error {
309 return true;
310 }
311 resolved.insert(call_id.to_string());
312 }
313 }
314 }
315 }
316 opened.iter().any(|call_id| !resolved.contains(call_id))
317 || marker_folded(
318 folded_text,
319 &[
320 "unresolved",
321 "open question",
322 "retry",
323 "blocked",
324 "待确认",
325 "未解决",
326 "重试",
327 "阻塞",
328 ],
329 )
330}
331
332fn is_error_or_decision(messages: &[CoreMessage], folded_text: &str) -> bool {
333 messages.iter().any(|message| {
334 matches!(&message.content, Content::Parts(parts) if parts.iter().any(|part| matches!(part, ContentPart::ToolResult { is_error: true, .. })))
335 }) || marker_folded(
336 folded_text,
337 &[
338 "error", "failed", "failure", "exception", "decision", "decided", "must", "should",
339 "错误", "失败", "异常", "决定", "选择", "必须", "应当",
340 ],
341 )
342}
343
344fn marker_folded(folded_text: &str, markers: &[&str]) -> bool {
345 markers.iter().any(|marker| folded_text.contains(marker))
346}
347
348fn unit_referenced_later(messages: &[CoreMessage], text: &str, later: &[String]) -> bool {
349 let mut references = messages
350 .iter()
351 .flat_map(|message| message.tool_calls.iter().map(|call| call.id.to_string()))
352 .collect::<BTreeSet<_>>();
353 references.extend(
354 text.split_whitespace()
355 .map(|token| token.trim_matches(|character: char| character.is_ascii_punctuation()))
356 .filter(|token| token.contains('/') || token.contains("://"))
357 .filter(|token| token.len() > 3)
358 .map(str::to_string),
359 );
360 references.iter().any(|reference| {
361 later
362 .iter()
363 .any(|later_text| contains_folded(later_text, reference))
364 })
365}
366
367#[cfg(test)]
368mod tests {
369 use super::*;
370 use crate::types::message::{ContentPart, ToolCall};
371
372 #[test]
373 fn unit_text_preserves_empty_part_separators() {
374 let messages = vec![CoreMessage::user(""), CoreMessage::user("next")];
375 assert_eq!(unit_text(&messages), "\nnext");
376 }
377
378 #[test]
379 fn unresolved_tool_unit_is_mandatory() {
380 let mut call = CoreMessage::assistant("working");
381 call.tool_calls.push(ToolCall {
382 id: "call-1".into(),
383 name: "read".into(),
384 arguments: serde_json::json!({"path": "/work/a"}),
385 });
386 let mut recent = CoreMessage::user("recent");
387 let messages = vec![call, recent];
388 let engine = ContextTokenEngine::char_approx();
389 let plan = plan_utility_archive(
390 &messages,
391 40,
392 20,
393 1,
394 &engine,
395 &UtilitySelectionContext {
396 goal: "",
397 criteria: &[],
398 preserved_refs: &[],
399 active_directives: &[],
400 },
401 );
402 assert!(plan.scores[0].mandatory);
403 assert!(plan.scores[0].has_unresolved);
404 assert_eq!(plan.retained_tokens, 2);
405 }
406
407 #[test]
408 fn chinese_directive_dependency_requires_bigram_overlap_not_shared_characters() {
409 let mut unrelated = CoreMessage::assistant("我们在文中回顾了天气");
413 let mut on_topic = CoreMessage::user("已按要求保持中文回答");
414 let mut recent = CoreMessage::user("recent");
415 let messages = vec![unrelated, on_topic, recent];
416 let plan = plan_utility_archive(
417 &messages,
418 70,
419 10,
420 1,
421 &ContextTokenEngine::char_approx(),
422 &UtilitySelectionContext {
423 goal: "",
424 criteria: &[],
425 preserved_refs: &[],
426 active_directives: &["必须用中文回答".into()],
427 },
428 );
429 assert!(
430 !plan.scores[0].mandatory,
431 "unrelated Chinese text must not bind to the directive"
432 );
433 assert!(
434 plan.scores[1].mandatory,
435 "text restating the directive must stay mandatory"
436 );
437 }
438
439 #[test]
440 fn preserved_ref_keeps_complete_tool_unit() {
441 let mut call = CoreMessage::assistant("read artifact");
442 call.tool_calls.push(ToolCall {
443 id: "call-keep".into(),
444 name: "read".into(),
445 arguments: serde_json::json!({}),
446 });
447 let mut result = CoreMessage::tool(vec![ContentPart::ToolResult {
448 call_id: "call-keep".into(),
449 output: "artifact".into(),
450 is_error: false,
451 durable_content: None,
452 }]);
453 let messages = vec![call, result];
454 let plan = plan_utility_archive(
455 &messages,
456 40,
457 0,
458 0,
459 &ContextTokenEngine::char_approx(),
460 &UtilitySelectionContext {
461 goal: "",
462 criteria: &[],
463 preserved_refs: &["call-keep".into()],
464 active_directives: &[],
465 },
466 );
467 assert!(plan.scores[0].mandatory);
468 assert_eq!(plan.archived_ranges, Vec::<Range<usize>>::new());
469 }
470
471 #[test]
472 fn measurement_aware_planner_ignores_stale_message_projection() {
473 let mut message = CoreMessage::user("short");
474 let engine = ContextTokenEngine::char_approx();
475 let measurements = vec![TokenMeasurement::for_message(&message, 2)];
476 let plan = plan_utility_archive_with_measurements(
477 &[message],
478 &measurements,
479 2,
480 1,
481 0,
482 &engine,
483 &UtilitySelectionContext {
484 goal: "",
485 criteria: &[],
486 preserved_refs: &[],
487 active_directives: &[],
488 },
489 );
490 assert_eq!(plan.scores[0].tokens, 2);
491 }
492}