use std::sync::atomic::Ordering;
use crate::cm_config::AgentConfig;
use crate::cm_types::Message;
use super::transforms::{
compress_tool_message_contents, drop_orphan_tool_messages, estimate_non_system_chars,
take_chat_timeline_markers, trim_messages_by_char_budget, trim_messages_by_count,
};
use super::{MESSAGE_PIPELINE_COUNTERS, MessagePipelineConfig, MessagePipelineCounters};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagePipelineStage {
SessionSyncStart,
AfterCompressTool,
AfterTrimByCount,
AfterTrimByCharBudget,
AfterSecondCompressTool,
AfterDropOrphanTool,
AfterMergeAssistantsInPlace,
}
impl MessagePipelineStage {
pub(crate) fn as_str(self) -> &'static str {
match self {
Self::SessionSyncStart => "start",
Self::AfterCompressTool => "after_compress_tool",
Self::AfterTrimByCount => "after_trim_count",
Self::AfterTrimByCharBudget => "after_trim_char_budget",
Self::AfterSecondCompressTool => "after_compress_tool_2",
Self::AfterDropOrphanTool => "after_drop_orphan_tool",
Self::AfterMergeAssistantsInPlace => "after_merge_assistants",
}
}
}
#[derive(Clone, Debug)]
pub struct PipelineStepSnapshot {
pub stage: MessagePipelineStage,
pub message_count: usize,
pub non_system_chars_est: usize,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct MessagePipelineDelta {
pub n_before: usize,
pub n_after: usize,
pub trim_count_hit: bool,
pub trim_char_hit: bool,
pub tool_compress_hits: usize,
}
#[derive(Default, Debug)]
pub struct MessagePipelineReport {
pub steps: Vec<PipelineStepSnapshot>,
}
impl MessagePipelineReport {
fn record(&mut self, stage: MessagePipelineStage, messages: &[Message]) {
self.steps.push(PipelineStepSnapshot {
stage,
message_count: messages.len(),
non_system_chars_est: estimate_non_system_chars(messages),
});
}
pub fn format_for_log(&self) -> String {
let parts: Vec<String> = self
.steps
.iter()
.map(|s| {
format!(
"{}:n={}:chars≈{}",
s.stage.as_str(),
s.message_count,
s.non_system_chars_est
)
})
.collect();
parts.join(" | ")
}
}
fn log_session_sync_step(stage: MessagePipelineStage, messages: &[Message]) {
log::trace!(
target: "crabmate::message_pipeline",
"session_sync_step stage={} message_count={} non_system_chars_est={}",
stage.as_str(),
messages.len(),
estimate_non_system_chars(messages),
);
}
fn record_and_trace(
report: &mut Option<&mut MessagePipelineReport>,
stage: MessagePipelineStage,
messages: &[Message],
) {
if let Some(r) = report.as_mut() {
r.record(stage, messages);
}
log_session_sync_step(stage, messages);
}
pub fn apply_session_sync_pipeline_with_config(
messages: &mut Vec<Message>,
cfg: MessagePipelineConfig,
mut report: Option<&mut MessagePipelineReport>,
) -> MessagePipelineDelta {
let ctr: &MessagePipelineCounters = &MESSAGE_PIPELINE_COUNTERS;
let parked = take_chat_timeline_markers(messages);
let n_before = messages.len();
record_and_trace(
&mut report,
MessagePipelineStage::SessionSyncStart,
messages,
);
let mut tool_compress_hits = compress_tool_message_contents(messages, cfg.tool_message_max_chars);
if tool_compress_hits > 0 {
ctr.tool_compress_hits
.fetch_add(tool_compress_hits as u64, Ordering::Relaxed);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterCompressTool,
messages,
);
let trim_count_hit = trim_messages_by_count(messages, cfg.max_message_history);
if trim_count_hit {
ctr.trim_count_hits.fetch_add(1, Ordering::Relaxed);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterTrimByCount,
messages,
);
let mut trim_char_hit = false;
if cfg.context_char_budget > 0 {
trim_char_hit = trim_messages_by_char_budget(
messages,
cfg.context_char_budget,
cfg.context_min_messages_after_system,
);
if trim_char_hit {
ctr.trim_char_budget_hits.fetch_add(1, Ordering::Relaxed);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterTrimByCharBudget,
messages,
);
let c2 = compress_tool_message_contents(messages, cfg.tool_message_max_chars);
if c2 > 0 {
ctr.tool_compress_hits
.fetch_add(c2 as u64, Ordering::Relaxed);
tool_compress_hits = tool_compress_hits.saturating_add(c2);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterSecondCompressTool,
messages,
);
}
let dropped = drop_orphan_tool_messages(messages);
if dropped > 0 {
ctr.orphan_tool_drops
.fetch_add(dropped as u64, Ordering::Relaxed);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterDropOrphanTool,
messages,
);
let n_after = messages.len();
messages.extend(parked);
MessagePipelineDelta {
n_before,
n_after,
trim_count_hit,
trim_char_hit,
tool_compress_hits,
}
}
pub fn apply_session_sync_pipeline(
messages: &mut Vec<Message>,
cfg: &AgentConfig,
report: Option<&mut MessagePipelineReport>,
) -> MessagePipelineDelta {
apply_session_sync_pipeline_with_config(messages, MessagePipelineConfig::from(cfg), report)
}