1use std::path::PathBuf;
2
3use crate::codec::OutputEncoding;
4use crate::errors::TokenFoldError;
5use crate::input::{CompressionInput, InputFormat};
6use crate::token_estimator::TokenEstimator;
7
8const MIN_LOSSY_TTL_SECONDS: u64 = 86_400;
13
14#[derive(Debug, Clone, PartialEq)]
16pub struct CompressionPolicy {
17 pub target_tokens: Option<usize>,
18 pub reserve_output_tokens: usize,
19 pub preset: Preset,
20 pub task_scope: TaskScope,
21 pub encoding: OutputEncoding,
22 pub pruning: Option<PruningPolicy>,
23 pub preserve_latest_user_message: bool,
24 pub disabled: Vec<String>,
25 pub experimental: bool,
29 pub enable: Vec<String>,
33 pub store_originals: bool,
37 pub retrieval_namespace: String,
40 pub retrieval_ttl_seconds: Option<u64>,
44 pub retrieval_backend: String,
48 pub retrieval_store_path: Option<PathBuf>,
51 pub lossy: Option<LossyPath>,
58 pub lossy_ratio: f64,
68 pub lossy_preserve: Vec<String>,
71 pub(crate) preview: bool,
83}
84
85#[derive(Debug, Clone, Copy, PartialEq, Eq)]
86pub enum Preset {
87 Conservative,
88 Balanced,
89 Aggressive,
90}
91
92#[derive(Debug, Clone, Copy, PartialEq, Eq)]
97pub enum LossyPath {
98 Heuristic,
99}
100
101#[derive(Debug, Clone, PartialEq)]
102pub struct PruningPolicy {
103 pub keep_ratio: Option<f64>,
104 pub preserve_paths: Vec<String>,
105 pub retrieval_store: Option<PathBuf>,
106 pub retrieval_namespace: Option<String>,
107}
108
109#[derive(Debug, Clone, Copy, PartialEq, Eq)]
110pub enum TaskScope {
111 All,
112 General,
113 CodeReview,
114 ChangeSummary,
115 Debugging,
116 Generation,
117 ApiOverview,
118 RetrievalQa,
119 AgentHistory,
120}
121
122impl CompressionPolicy {
123 pub fn builder() -> CompressionPolicyBuilder {
124 CompressionPolicyBuilder::default()
125 }
126
127 pub fn validate(&self) -> Result<(), TokenFoldError> {
134 if let Some(pruning) = &self.pruning {
135 if pruning.keep_ratio.is_none() && self.target_tokens.is_none() {
136 return Err(TokenFoldError::ConfigError(
137 "pruning requires target_tokens or keep_ratio".to_string(),
138 ));
139 }
140 if pruning
141 .keep_ratio
142 .is_some_and(|ratio| !(0.0 < ratio && ratio <= 1.0))
143 {
144 return Err(TokenFoldError::ConfigError(
145 "keep_ratio must be greater than 0 and at most 1".to_string(),
146 ));
147 }
148 }
149 if self.disabled.iter().any(|id| id == "secret_redaction") {
150 return Err(TokenFoldError::ConfigError(
151 "secret_redaction cannot be disabled via CompressionPolicy.disabled".to_string(),
152 ));
153 }
154 if !(0.0..=1.0).contains(&self.lossy_ratio) {
155 return Err(TokenFoldError::ConfigError(format!(
156 "lossy_ratio must be between 0.0 and 1.0, got {}",
157 self.lossy_ratio
158 )));
159 }
160 if self.lossy.is_some() {
161 if self.retrieval_backend != "filesystem" {
165 return Err(TokenFoldError::ConfigError(format!(
166 "lossy pruning requires a durable retrieval backend (\"filesystem\"); \
167 {:?} would make dropped items unrecoverable",
168 self.retrieval_backend
169 )));
170 }
171 let effective_ttl = self
172 .retrieval_ttl_seconds
173 .unwrap_or(crate::retrieval_store::DEFAULT_TTL_SECONDS);
174 if effective_ttl < MIN_LOSSY_TTL_SECONDS {
175 return Err(TokenFoldError::ConfigError(format!(
176 "lossy pruning requires retrieval_ttl_seconds >= {MIN_LOSSY_TTL_SECONDS} \
177 (got {effective_ttl}); a near-immediate expiry has no real recoverability"
178 )));
179 }
180 }
181 Ok(())
182 }
183}
184
185#[derive(Debug, Clone, Default)]
186pub struct CompressionPolicyBuilder {
187 target_tokens: Option<usize>,
188 reserve_output_tokens: Option<usize>,
189 preset: Option<Preset>,
190 task_scope: Option<TaskScope>,
191 encoding: Option<OutputEncoding>,
192 pruning: Option<PruningPolicy>,
193 preserve_latest_user_message: Option<bool>,
194 disabled: Vec<String>,
195 experimental: bool,
196 enable: Vec<String>,
197 store_originals: bool,
198 retrieval_namespace: Option<String>,
199 retrieval_ttl_seconds: Option<u64>,
200 retrieval_backend: Option<String>,
201 retrieval_store_path: Option<PathBuf>,
202 lossy: Option<LossyPath>,
203 lossy_ratio: Option<f64>,
204 lossy_preserve: Vec<String>,
205 preview: bool,
206}
207
208impl CompressionPolicyBuilder {
209 pub fn target_tokens(mut self, target_tokens: usize) -> Self {
210 self.target_tokens = Some(target_tokens);
211 self
212 }
213
214 pub fn reserve_output_tokens(mut self, reserve_output_tokens: usize) -> Self {
215 self.reserve_output_tokens = Some(reserve_output_tokens);
216 self
217 }
218
219 pub fn preset(mut self, preset: Preset) -> Self {
220 self.preset = Some(preset);
221 self
222 }
223
224 pub fn task_scope(mut self, task_scope: TaskScope) -> Self {
225 self.task_scope = Some(task_scope);
226 self
227 }
228
229 pub fn encoding(mut self, encoding: OutputEncoding) -> Self {
230 self.encoding = Some(encoding);
231 self
232 }
233
234 pub fn pruning(mut self, pruning: PruningPolicy) -> Self {
235 self.lossy = Some(LossyPath::Heuristic);
236 self.lossy_ratio = Some(pruning.keep_ratio.unwrap_or(0.0));
237 self.lossy_preserve = pruning.preserve_paths.clone();
238 if pruning.retrieval_store.is_some() {
239 self.retrieval_store_path = pruning.retrieval_store.clone();
240 }
241 if pruning.retrieval_namespace.is_some() {
242 self.retrieval_namespace = pruning.retrieval_namespace.clone();
243 }
244 self.pruning = Some(pruning);
245 self
246 }
247
248 pub fn preserve_latest_user_message(mut self, preserve: bool) -> Self {
249 self.preserve_latest_user_message = Some(preserve);
250 self
251 }
252
253 pub fn disable(mut self, transform_id: impl Into<String>) -> Self {
254 self.disabled.push(transform_id.into());
255 self
256 }
257
258 pub fn experimental(mut self, experimental: bool) -> Self {
259 self.experimental = experimental;
260 self
261 }
262
263 pub fn enable(mut self, transform_id: impl Into<String>) -> Self {
264 self.enable.push(transform_id.into());
265 self
266 }
267
268 pub fn store_originals(mut self, store_originals: bool) -> Self {
269 self.store_originals = store_originals;
270 self
271 }
272
273 pub fn retrieval_namespace(mut self, namespace: impl Into<String>) -> Self {
274 self.retrieval_namespace = Some(namespace.into());
275 self
276 }
277
278 pub fn retrieval_ttl_seconds(mut self, ttl_seconds: Option<u64>) -> Self {
279 self.retrieval_ttl_seconds = ttl_seconds;
280 self
281 }
282
283 pub fn retrieval_backend(mut self, backend: impl Into<String>) -> Self {
284 self.retrieval_backend = Some(backend.into());
285 self
286 }
287
288 pub fn retrieval_store_path(mut self, store_path: Option<PathBuf>) -> Self {
289 self.retrieval_store_path = store_path;
290 self
291 }
292
293 pub fn lossy(mut self, lossy: LossyPath) -> Self {
294 self.lossy = Some(lossy);
295 self
296 }
297
298 pub fn lossy_ratio(mut self, ratio: f64) -> Self {
299 self.lossy_ratio = Some(ratio);
300 self
301 }
302
303 pub fn lossy_preserve(mut self, path: impl Into<String>) -> Self {
304 self.lossy_preserve.push(path.into());
305 self
306 }
307
308 pub fn preview(mut self, preview: bool) -> Self {
314 self.preview = preview;
315 self
316 }
317
318 pub fn build(self) -> Result<CompressionPolicy, TokenFoldError> {
319 let policy = CompressionPolicy {
320 target_tokens: self.target_tokens,
321 reserve_output_tokens: self.reserve_output_tokens.unwrap_or(0),
322 preset: self.preset.unwrap_or(Preset::Balanced),
323 task_scope: self.task_scope.unwrap_or(TaskScope::All),
324 encoding: self.encoding.unwrap_or_default(),
325 pruning: self.pruning,
326 preserve_latest_user_message: self.preserve_latest_user_message.unwrap_or(true),
327 disabled: self.disabled,
328 experimental: self.experimental,
329 enable: self.enable,
330 store_originals: self.store_originals,
331 retrieval_namespace: self
332 .retrieval_namespace
333 .unwrap_or_else(|| "default".to_string()),
334 retrieval_ttl_seconds: self.retrieval_ttl_seconds,
335 retrieval_backend: self
336 .retrieval_backend
337 .unwrap_or_else(|| "filesystem".to_string()),
338 retrieval_store_path: self.retrieval_store_path,
339 lossy: self.lossy,
340 lossy_ratio: self.lossy_ratio.unwrap_or(0.3),
341 lossy_preserve: self.lossy_preserve,
342 preview: self.preview,
343 };
344 policy.validate()?;
345 Ok(policy)
346 }
347}
348
349pub fn protected_floor(
351 input: &CompressionInput,
352 policy: &CompressionPolicy,
353 estimator: &dyn TokenEstimator,
354) -> usize {
355 estimator.count_bytes(&protected_segments(input, policy).concat())
356}
357
358pub fn protected_segments(input: &CompressionInput, policy: &CompressionPolicy) -> Vec<Vec<u8>> {
364 match input.format {
365 InputFormat::OpenAiJson => extract_openai_protected(&input.bytes, policy),
366 InputFormat::AnthropicJson => extract_anthropic_protected(&input.bytes, policy),
367 InputFormat::GitDiff => extract_diff_protected(&input.bytes),
368 InputFormat::PlainText
373 | InputFormat::CommandOutput
374 | InputFormat::Json
375 | InputFormat::Auto => Vec::new(),
376 }
377}
378
379fn extract_openai_protected(bytes: &[u8], policy: &CompressionPolicy) -> Vec<Vec<u8>> {
380 let Ok(value) = serde_json::from_slice::<serde_json::Value>(bytes) else {
381 return Vec::new();
382 };
383 let Some(messages) = value.get("messages").and_then(|m| m.as_array()) else {
384 return Vec::new();
385 };
386
387 let mut segments = Vec::new();
388 for message in messages {
389 if message.get("role").and_then(|r| r.as_str()) == Some("system")
390 && let Some(bytes) = message_content_bytes(message)
391 {
392 segments.push(bytes);
393 }
394 }
395 if policy.preserve_latest_user_message
396 && let Some(last_user) = messages
397 .iter()
398 .rev()
399 .find(|m| m.get("role").and_then(|r| r.as_str()) == Some("user"))
400 && let Some(bytes) = message_content_bytes(last_user)
401 {
402 segments.push(bytes);
403 }
404 segments
405}
406
407fn extract_anthropic_protected(bytes: &[u8], policy: &CompressionPolicy) -> Vec<Vec<u8>> {
408 let Ok(value) = serde_json::from_slice::<serde_json::Value>(bytes) else {
409 return Vec::new();
410 };
411
412 let mut segments = Vec::new();
413 match value.get("system") {
418 Some(serde_json::Value::String(text)) => segments.push(text.as_bytes().to_vec()),
419 Some(structured @ serde_json::Value::Array(_)) => {
420 if let Ok(bytes) = serde_json::to_vec(structured) {
421 segments.push(bytes);
422 }
423 }
424 _ => {}
425 }
426 if policy.preserve_latest_user_message
427 && let Some(last_user) =
428 value
429 .get("messages")
430 .and_then(|m| m.as_array())
431 .and_then(|messages| {
432 messages
433 .iter()
434 .rev()
435 .find(|m| m.get("role").and_then(|r| r.as_str()) == Some("user"))
436 })
437 && let Some(bytes) = message_content_bytes(last_user)
438 {
439 segments.push(bytes);
440 }
441 segments
442}
443
444fn message_content_bytes(message: &serde_json::Value) -> Option<Vec<u8>> {
445 match message.get("content") {
446 Some(serde_json::Value::String(text)) => Some(text.as_bytes().to_vec()),
447 Some(structured) => serde_json::to_vec(structured).ok(),
448 None => None,
449 }
450}
451
452fn extract_diff_protected(bytes: &[u8]) -> Vec<Vec<u8>> {
455 let text = String::from_utf8_lossy(bytes);
456 let mut segments = Vec::new();
457 for line in text.lines() {
458 if line.starts_with("diff --git")
459 || line.starts_with("--- ")
460 || line.starts_with("+++ ")
461 || line.starts_with("@@")
462 {
463 let mut segment = line.as_bytes().to_vec();
464 segment.push(b'\n');
465 segments.push(segment);
466 }
467 }
468 segments
469}
470
471#[cfg(test)]
472mod tests {
473 use super::*;
474 use crate::token_estimator::ByteHeuristicEstimator;
475
476 #[test]
477 fn default_mode_is_balanced() {
478 let policy = CompressionPolicy::builder().build().unwrap();
479 assert_eq!(policy.preset, Preset::Balanced);
480 }
481
482 #[test]
483 fn store_originals_defaults_to_false_with_a_default_namespace() {
484 let policy = CompressionPolicy::builder().build().unwrap();
485 assert!(!policy.store_originals);
486 assert_eq!(policy.retrieval_namespace, "default");
487 }
488
489 #[test]
490 fn store_originals_and_namespace_are_settable_via_the_builder() {
491 let policy = CompressionPolicy::builder()
492 .store_originals(true)
493 .retrieval_namespace("project-x")
494 .retrieval_ttl_seconds(Some(60))
495 .retrieval_backend("memory")
496 .retrieval_store_path(Some(std::path::PathBuf::from("/tmp/custom")))
497 .build()
498 .unwrap();
499 assert!(policy.store_originals);
500 assert_eq!(policy.retrieval_namespace, "project-x");
501 assert_eq!(policy.retrieval_ttl_seconds, Some(60));
502 assert_eq!(policy.retrieval_backend, "memory");
503 assert_eq!(
504 policy.retrieval_store_path,
505 Some(std::path::PathBuf::from("/tmp/custom"))
506 );
507 }
508
509 #[test]
510 fn retrieval_defaults_are_none_ttl_and_filesystem_backend() {
511 let policy = CompressionPolicy::builder().build().unwrap();
512 assert_eq!(policy.retrieval_ttl_seconds, None);
513 assert_eq!(policy.retrieval_backend, "filesystem");
514 assert_eq!(policy.retrieval_store_path, None);
515 }
516
517 #[test]
518 fn secret_redaction_cannot_be_disabled_through_policy() {
519 let err = CompressionPolicy::builder()
520 .disable("secret_redaction")
521 .build()
522 .unwrap_err();
523 assert!(matches!(err, TokenFoldError::ConfigError(_)));
524 }
525
526 #[test]
527 fn disabling_other_transforms_is_allowed() {
528 let policy = CompressionPolicy::builder()
529 .disable("json_minify")
530 .build()
531 .unwrap();
532 assert_eq!(policy.disabled, vec!["json_minify".to_string()]);
533 }
534
535 #[test]
536 fn floor_is_zero_for_plain_text() {
537 let input = CompressionInput::plain_text(b"just some plain text".to_vec());
538 let policy = CompressionPolicy::builder().build().unwrap();
539 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
540 assert_eq!(floor, 0);
541 }
542
543 #[test]
544 fn floor_covers_system_and_latest_user_message_for_openai_json() {
545 let payload = serde_json::json!({
546 "model": "gpt-4",
547 "messages": [
548 {"role": "system", "content": "You are a helpful assistant."},
549 {"role": "user", "content": "first question"},
550 {"role": "assistant", "content": "first answer"},
551 {"role": "user", "content": "second question"},
552 ]
553 });
554 let input = CompressionInput::openai_json(serde_json::to_vec(&payload).unwrap());
555 let policy = CompressionPolicy::builder().build().unwrap();
556 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
557
558 let expected_bytes = "You are a helpful assistant.".len() + "second question".len();
559 assert_eq!(
560 floor,
561 ByteHeuristicEstimator.count_bytes(&vec![0u8; expected_bytes])
562 );
563 assert!(floor < ByteHeuristicEstimator.count_bytes(input.bytes.as_slice()));
565 }
566
567 #[test]
568 fn floor_excludes_latest_user_message_when_policy_disables_preservation() {
569 let payload = serde_json::json!({
570 "messages": [
571 {"role": "system", "content": "system prompt"},
572 {"role": "user", "content": "question"},
573 ]
574 });
575 let input = CompressionInput::openai_json(serde_json::to_vec(&payload).unwrap());
576 let policy = CompressionPolicy::builder()
577 .preserve_latest_user_message(false)
578 .build()
579 .unwrap();
580 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
581 assert_eq!(floor, ByteHeuristicEstimator.count_bytes(b"system prompt"));
582 }
583
584 #[test]
585 fn floor_covers_system_and_latest_user_message_for_anthropic_json() {
586 let payload = serde_json::json!({
587 "system": "system prompt",
588 "messages": [
589 {"role": "user", "content": "first"},
590 {"role": "assistant", "content": "reply"},
591 {"role": "user", "content": "second"},
592 ]
593 });
594 let input = CompressionInput::anthropic_json(serde_json::to_vec(&payload).unwrap());
595 let policy = CompressionPolicy::builder().build().unwrap();
596 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
597 let expected_bytes = "system prompt".len() + "second".len();
598 assert_eq!(
599 floor,
600 ByteHeuristicEstimator.count_bytes(&vec![0u8; expected_bytes])
601 );
602 }
603
604 #[test]
605 fn floor_covers_structured_anthropic_system_content_not_just_a_plain_string() {
606 let payload = serde_json::json!({
611 "system": [{"type": "text", "text": "structured system prompt"}],
612 "messages": [
613 {"role": "user", "content": "first"},
614 ]
615 });
616 let input = CompressionInput::anthropic_json(serde_json::to_vec(&payload).unwrap());
617 let policy = CompressionPolicy::builder()
618 .preserve_latest_user_message(false)
619 .build()
620 .unwrap();
621 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
622 assert!(
623 floor > 0,
624 "a structured Anthropic `system` array must contribute to the protected floor"
625 );
626 let segments = protected_segments(&input, &policy);
627 let system_bytes = serde_json::to_vec(
628 &serde_json::json!([{"type": "text", "text": "structured system prompt"}]),
629 )
630 .unwrap();
631 assert!(
632 segments.contains(&system_bytes),
633 "the structured system content must be a protected segment, byte-for-byte"
634 );
635 }
636
637 #[test]
638 fn floor_keeps_diff_headers_and_hunk_markers_only() {
639 let diff =
640 b"diff --git a/f.rs b/f.rs\n--- a/f.rs\n+++ b/f.rs\n@@ -1,2 +1,2 @@\n-old\n+new\n";
641 let input = CompressionInput::git_diff(diff.to_vec());
642 let policy = CompressionPolicy::builder().build().unwrap();
643 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
644 assert!(floor > 0);
645 assert!(floor < ByteHeuristicEstimator.count_bytes(diff));
646 }
647
648 #[test]
649 fn lossy_defaults_to_disabled_with_a_default_ratio() {
650 let policy = CompressionPolicy::builder().build().unwrap();
651 assert_eq!(policy.lossy, None);
652 assert_eq!(policy.lossy_ratio, 0.3);
653 assert!(policy.lossy_preserve.is_empty());
654 }
655
656 #[test]
657 fn lossy_is_settable_via_the_builder() {
658 let policy = CompressionPolicy::builder()
659 .lossy(LossyPath::Heuristic)
660 .lossy_ratio(0.5)
661 .lossy_preserve("items")
662 .lossy_preserve("data.results")
663 .build()
664 .unwrap();
665 assert_eq!(policy.lossy, Some(LossyPath::Heuristic));
666 assert_eq!(policy.lossy_ratio, 0.5);
667 assert_eq!(policy.lossy_preserve, vec!["items", "data.results"]);
668 }
669
670 #[test]
671 fn lossy_refuses_memory_retrieval_backend() {
672 let err = CompressionPolicy::builder()
673 .lossy(LossyPath::Heuristic)
674 .retrieval_backend("memory")
675 .build()
676 .unwrap_err();
677 assert!(matches!(err, TokenFoldError::ConfigError(_)));
678 }
679
680 #[test]
681 fn lossy_refuses_a_ttl_below_the_floor() {
682 let err = CompressionPolicy::builder()
683 .lossy(LossyPath::Heuristic)
684 .retrieval_ttl_seconds(Some(60))
685 .build()
686 .unwrap_err();
687 assert!(matches!(err, TokenFoldError::ConfigError(_)));
688 }
689
690 #[test]
691 fn lossy_with_default_retrieval_settings_is_accepted() {
692 let policy = CompressionPolicy::builder()
695 .lossy(LossyPath::Heuristic)
696 .build()
697 .unwrap();
698 assert_eq!(policy.lossy, Some(LossyPath::Heuristic));
699 }
700
701 #[test]
702 fn lossy_ratio_outside_unit_interval_is_rejected() {
703 let err = CompressionPolicy::builder()
704 .lossy_ratio(1.5)
705 .build()
706 .unwrap_err();
707 assert!(matches!(err, TokenFoldError::ConfigError(_)));
708 let err = CompressionPolicy::builder()
709 .lossy_ratio(-0.1)
710 .build()
711 .unwrap_err();
712 assert!(matches!(err, TokenFoldError::ConfigError(_)));
713 }
714
715 #[test]
716 fn malformed_json_never_panics_and_yields_zero_floor() {
717 let input = CompressionInput::openai_json(b"{not json".to_vec());
718 let policy = CompressionPolicy::builder().build().unwrap();
719 let floor = protected_floor(&input, &policy, &ByteHeuristicEstimator);
720 assert_eq!(floor, 0);
721 }
722}