Skip to main content

deepstrike_core/context/
utility.rs

1//! Deterministic value-aware selection over indivisible context units.
2
3// DEL-1 migration window (0.2.67 → removed 0.2.68): this module still reads/writes the
4use 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
45/// Select complete units to retain under `target_tokens`.
46///
47/// Mandatory dependencies are retained even when they alone exceed the target;
48/// callers can then escalate pressure honestly instead of silently deleting the
49/// evidence required to continue the task.
50pub 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
69/// Measurement-aware planner entry point used by context partitions. The parallel slice is
70/// host-owned evidence aligned with `messages`; an empty or short slice falls back to the
71/// deprecated dual-write field during the migration window.
72pub 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        // Under the old per-character CJK vocabulary, any Chinese unit sharing two
410        // common characters (中/文/回…) with an active directive was marked mandatory,
411        // so compression could never archive unrelated Chinese history.
412        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}