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