Skip to main content

deepstrike_core/context/
utility.rs

1//! Deterministic value-aware selection over indivisible context units.
2
3use 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
44/// Select complete units to retain under `target_tokens`.
45///
46/// Mandatory dependencies are retained even when they alone exceed the target;
47/// callers can then escalate pressure honestly instead of silently deleting the
48/// evidence required to continue the task.
49pub 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
68/// Measurement-aware planner entry point used by context partitions. The parallel slice is
69/// host-owned evidence aligned with `messages`; an empty or short slice falls back to deterministic
70/// engine recomputation for the missing entries.
71pub 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        // Under the old per-character CJK vocabulary, any Chinese unit sharing two
409        // common characters (中/文/回…) with an active directive was marked mandatory,
410        // so compression could never archive unrelated Chinese history.
411        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}