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,
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(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>,
) {
let ctr: &MessagePipelineCounters = &MESSAGE_PIPELINE_COUNTERS;
record_and_trace(
&mut report,
MessagePipelineStage::SessionSyncStart,
messages,
);
let c1 = compress_tool_message_contents(messages, cfg.tool_message_max_chars);
if c1 > 0 {
ctr.tool_compress_hits
.fetch_add(c1 as u64, Ordering::Relaxed);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterCompressTool,
messages,
);
if trim_messages_by_count(messages, cfg.max_message_history) {
ctr.trim_count_hits.fetch_add(1, Ordering::Relaxed);
}
record_and_trace(
&mut report,
MessagePipelineStage::AfterTrimByCount,
messages,
);
if cfg.context_char_budget > 0 {
if trim_messages_by_char_budget(
messages,
cfg.context_char_budget,
cfg.context_min_messages_after_system,
) {
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);
}
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,
);
}
pub fn apply_session_sync_pipeline(
messages: &mut Vec<Message>,
cfg: &AgentConfig,
report: Option<&mut MessagePipelineReport>,
) {
apply_session_sync_pipeline_with_config(messages, MessagePipelineConfig::from(cfg), report);
}