use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::fmt;
use std::sync::Arc;
use crate::runtime::message::RuntimeMessage;
use crate::runtime::typed_id::{EventId, SessionId};
#[derive(Clone)]
pub enum MessageFilter {
TimeRange {
from: Option<DateTime<Utc>>,
to: Option<DateTime<Utc>>,
},
EventTypes(Vec<String>),
ToolName(String),
Search(String),
ExcludeIds(Vec<EventId>),
IncludeIds(Vec<EventId>),
Custom(Arc<dyn Fn(&RuntimeMessage) -> bool + Send + Sync>),
}
impl fmt::Debug for MessageFilter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::TimeRange { from, to } => f
.debug_struct("TimeRange")
.field("from", from)
.field("to", to)
.finish(),
Self::EventTypes(types) => f.debug_tuple("EventTypes").field(types).finish(),
Self::ToolName(name) => f.debug_tuple("ToolName").field(name).finish(),
Self::Search(query) => f.debug_tuple("Search").field(query).finish(),
Self::ExcludeIds(ids) => f.debug_tuple("ExcludeIds").field(ids).finish(),
Self::IncludeIds(ids) => f.debug_tuple("IncludeIds").field(ids).finish(),
Self::Custom(_) => f.debug_tuple("Custom").field(&"<fn>").finish(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum InjectionPosition {
Start,
End,
BeforeIndex(usize),
AfterIndex(usize),
}
#[derive(Debug, Clone)]
pub struct InjectedMessage {
pub position: InjectionPosition,
pub message: RuntimeMessage,
}
impl InjectedMessage {
pub fn at_start(message: RuntimeMessage) -> Self {
Self {
position: InjectionPosition::Start,
message,
}
}
pub fn at_end(message: RuntimeMessage) -> Self {
Self {
position: InjectionPosition::End,
message,
}
}
pub fn before_index(index: usize, message: RuntimeMessage) -> Self {
Self {
position: InjectionPosition::BeforeIndex(index),
message,
}
}
pub fn after_index(index: usize, message: RuntimeMessage) -> Self {
Self {
position: InjectionPosition::AfterIndex(index),
message,
}
}
}
#[derive(Debug, Clone)]
pub struct FilterContext {
pub total_count: usize,
pub filtered_count: usize,
pub excluded_count: usize,
}
pub trait PrependTransform: Send + Sync {
fn transform(&self, ctx: &FilterContext) -> Option<RuntimeMessage>;
}
pub struct ExcludedNoticeTransform {
pub format: String,
}
impl ExcludedNoticeTransform {
pub fn new(format: impl Into<String>) -> Self {
Self {
format: format.into(),
}
}
pub fn infinity_context() -> Self {
Self::new(
"[{} earlier messages are not in this context. \
To answer questions about earlier parts of the conversation, \
search them with the `query_history` tool.]",
)
}
}
impl PrependTransform for ExcludedNoticeTransform {
fn transform(&self, ctx: &FilterContext) -> Option<RuntimeMessage> {
if ctx.excluded_count > 0 {
let text = self.format.replace("{}", &ctx.excluded_count.to_string());
Some(RuntimeMessage::system(&text))
} else {
None
}
}
}
#[derive(Clone)]
pub struct MessageQuery {
pub session_id: SessionId,
pub filters: Vec<MessageFilter>,
pub injections: Vec<InjectedMessage>,
pub limit: Option<i64>,
pub offset: Option<i64>,
pub keep_head: Option<usize>,
pub after_sequence: Option<i64>,
pub prepend_transform: Option<Arc<dyn PrependTransform>>,
}
impl Default for MessageQuery {
fn default() -> Self {
Self {
session_id: SessionId::from_seed(0),
filters: Vec::new(),
injections: Vec::new(),
limit: None,
offset: None,
keep_head: None,
after_sequence: None,
prepend_transform: None,
}
}
}
impl std::fmt::Debug for MessageQuery {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MessageQuery")
.field("session_id", &self.session_id)
.field("filters", &self.filters)
.field("injections", &self.injections)
.field("limit", &self.limit)
.field("offset", &self.offset)
.field("keep_head", &self.keep_head)
.field("after_sequence", &self.after_sequence)
.field("prepend_transform", &self.prepend_transform.is_some())
.finish()
}
}
impl MessageQuery {
pub fn new(session_id: SessionId) -> Self {
Self {
session_id,
filters: Vec::new(),
injections: Vec::new(),
limit: None,
offset: None,
keep_head: None,
after_sequence: None,
prepend_transform: None,
}
}
pub fn with_prepend_transform(mut self, transform: Arc<dyn PrependTransform>) -> Self {
self.prepend_transform = Some(transform);
self
}
pub fn with_filter(mut self, filter: MessageFilter) -> Self {
self.filters.push(filter);
self
}
pub fn with_filters(mut self, filters: impl IntoIterator<Item = MessageFilter>) -> Self {
self.filters.extend(filters);
self
}
pub fn with_injection(mut self, injection: InjectedMessage) -> Self {
self.injections.push(injection);
self
}
pub fn with_injections(
mut self,
injections: impl IntoIterator<Item = InjectedMessage>,
) -> Self {
self.injections.extend(injections);
self
}
pub fn with_limit(mut self, limit: i64) -> Self {
self.limit = Some(limit);
self
}
pub fn with_offset(mut self, offset: i64) -> Self {
self.offset = Some(offset);
self
}
pub fn with_keep_head(mut self, keep_head: usize) -> Self {
self.keep_head = Some(keep_head);
self
}
pub fn after_sequence(mut self, sequence: i64) -> Self {
self.after_sequence = Some(sequence);
self
}
pub fn has_db_filters(&self) -> bool {
self.filters
.iter()
.any(|f| !matches!(f, MessageFilter::Custom(_)))
}
pub fn has_custom_filters(&self) -> bool {
self.filters
.iter()
.any(|f| matches!(f, MessageFilter::Custom(_)))
}
pub fn has_injections(&self) -> bool {
!self.injections.is_empty()
}
pub fn apply_injections(&self, messages: &mut Vec<RuntimeMessage>) {
let mut start_injections = Vec::new();
let mut end_injections = Vec::new();
let mut index_injections: Vec<_> = Vec::new();
for inj in &self.injections {
match &inj.position {
InjectionPosition::Start => start_injections.push(inj.message.clone()),
InjectionPosition::End => end_injections.push(inj.message.clone()),
InjectionPosition::BeforeIndex(idx) => {
index_injections.push((*idx, true, inj.message.clone()))
}
InjectionPosition::AfterIndex(idx) => {
index_injections.push((*idx, false, inj.message.clone()))
}
}
}
for msg in start_injections.into_iter().rev() {
messages.insert(0, msg);
}
index_injections.sort_by_key(|entry| std::cmp::Reverse(entry.0));
for (idx, is_before, msg) in index_injections {
let insert_idx = if is_before {
idx.min(messages.len())
} else {
idx.saturating_add(1).min(messages.len())
};
messages.insert(insert_idx, msg);
}
messages.extend(end_injections);
}
pub fn apply_window_bounds(&self, messages: &mut Vec<RuntimeMessage>) {
if let Some(offset) = self.offset {
let offset = offset.max(0) as usize;
if offset < messages.len() {
messages.drain(0..offset);
} else {
messages.clear();
}
}
if let Some(limit) = self.limit {
let limit = limit.max(0) as usize;
let keep_head = self.keep_head.unwrap_or(0).min(messages.len());
if messages.len() > keep_head + limit {
let drain_end = messages.len() - limit;
messages.drain(keep_head..drain_end);
}
}
}
pub fn prepend_excluded_notice(
&self,
messages: &mut Vec<RuntimeMessage>,
count_before_limit: usize,
) {
if let Some(ref transform) = self.prepend_transform {
let ctx = FilterContext {
total_count: count_before_limit,
filtered_count: messages.len(),
excluded_count: count_before_limit.saturating_sub(messages.len()),
};
if let Some(prepend_msg) = transform.transform(&ctx) {
messages.insert(0, prepend_msg);
}
}
}
pub fn apply_windowing(&self, messages: &mut Vec<RuntimeMessage>) {
let count_before_limit = messages.len();
self.apply_window_bounds(messages);
self.prepend_excluded_notice(messages, count_before_limit);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AnchoredWindow {
pub head_len: usize,
pub recent_start: usize,
}
impl AnchoredWindow {
pub fn hidden(&self) -> usize {
self.recent_start.saturating_sub(self.head_len)
}
}
pub fn anchored_window(
costs: &[usize],
keep_head: usize,
min_tail: usize,
max_tail: Option<usize>,
budget: usize,
) -> AnchoredWindow {
let len = costs.len();
let keep_head = keep_head.min(len);
let mut min_tail = min_tail.min(len);
if let Some(max_tail) = max_tail {
min_tail = min_tail.min(max_tail.max(1));
}
let mut recent_start = len - min_tail;
if recent_start <= keep_head {
return AnchoredWindow {
head_len: keep_head,
recent_start: keep_head,
};
}
let head_cost: usize = costs[..keep_head].iter().sum();
let mut window_cost: usize = head_cost + costs[recent_start..].iter().sum::<usize>();
let mut tail_count = min_tail;
while recent_start > keep_head {
if let Some(max_tail) = max_tail
&& tail_count >= max_tail
{
break;
}
let next = recent_start - 1;
let next_cost = costs[next];
if window_cost + next_cost > budget {
break;
}
recent_start = next;
window_cost += next_cost;
tail_count += 1;
}
AnchoredWindow {
head_len: keep_head,
recent_start,
}
}
pub trait MessageFilterProvider: Send + Sync {
fn apply_filters(&self, query: &mut MessageQuery, config: &serde_json::Value);
fn post_load(&self, messages: &mut Vec<RuntimeMessage>, config: &serde_json::Value) {
let _ = (messages, config);
}
fn priority(&self) -> i32 {
0
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime::message::RuntimeMessageRole;
fn messages() -> Vec<RuntimeMessage> {
["goal", "first", "second", "third", "latest"]
.into_iter()
.map(RuntimeMessage::user)
.collect()
}
fn assert_messages(actual: &[RuntimeMessage], expected: &[RuntimeMessage]) {
assert_eq!(
serde_json::to_value(actual).unwrap(),
serde_json::to_value(expected).unwrap()
);
}
#[test]
fn query_builder_preserves_filter_and_window_settings() {
let session = SessionId::new();
let query = MessageQuery::new(session)
.with_filter(MessageFilter::Search("needle".into()))
.with_filters([MessageFilter::EventTypes(vec!["input.message".into()])])
.with_limit(50)
.with_offset(10)
.with_keep_head(2)
.after_sequence(7);
assert_eq!(query.session_id, session);
assert!(
matches!(&query.filters[..], [MessageFilter::Search(s), MessageFilter::EventTypes(k)]
if s == "needle" && k == &["input.message"])
);
assert_eq!(
(
query.limit,
query.offset,
query.keep_head,
query.after_sequence
),
(Some(50), Some(10), Some(2), Some(7))
);
}
#[test]
fn filter_classification_handles_empty_db_custom_and_mixed_queries() {
let session = SessionId::new();
let empty = MessageQuery::new(session);
assert!(!empty.has_db_filters());
assert!(!empty.has_custom_filters());
let custom = MessageFilter::Custom(Arc::new(|m| m.role == RuntimeMessageRole::User));
let query = empty.clone().with_filter(custom.clone());
assert!(!query.has_db_filters());
assert!(query.has_custom_filters());
for filter in [
MessageFilter::Search("hello".into()),
MessageFilter::EventTypes(vec!["input.message".into()]),
MessageFilter::ToolName("lookup".into()),
MessageFilter::ExcludeIds(vec![EventId::new()]),
MessageFilter::IncludeIds(vec![EventId::new()]),
MessageFilter::TimeRange {
from: Some(Utc::now()),
to: None,
},
] {
let db = empty.clone().with_filter(filter);
assert!(db.has_db_filters());
assert!(!db.has_custom_filters());
let mixed = db.with_filter(custom.clone());
assert!(mixed.has_db_filters());
assert!(mixed.has_custom_filters());
}
}
#[test]
fn injections_preserve_messages_at_each_position_and_clamp_large_indices() {
for original in [vec![], messages()] {
let len = original.len();
for (position, index) in [
(InjectionPosition::Start, 0),
(InjectionPosition::End, len),
(InjectionPosition::BeforeIndex(1), 1.min(len)),
(InjectionPosition::AfterIndex(1), 2.min(len)),
(InjectionPosition::BeforeIndex(10), len),
(InjectionPosition::AfterIndex(10), len),
(InjectionPosition::BeforeIndex(usize::MAX), len),
(InjectionPosition::AfterIndex(usize::MAX), len),
] {
let inserted = RuntimeMessage::system("injected");
let query = MessageQuery::new(SessionId::new()).with_injection(InjectedMessage {
position,
message: inserted.clone(),
});
assert!(query.has_injections());
let mut actual = original.clone();
query.apply_injections(&mut actual);
let mut expected = original.clone();
expected.insert(index, inserted);
assert_messages(&actual, &expected);
}
}
}
#[test]
fn start_and_end_injections_preserve_declared_order() {
let injected: Vec<_> = ["start one", "start two", "end one", "end two"]
.into_iter()
.map(RuntimeMessage::system)
.collect();
let query = MessageQuery::new(SessionId::new())
.with_injection(InjectedMessage::at_start(injected[0].clone()))
.with_injections([
InjectedMessage::at_start(injected[1].clone()),
InjectedMessage::at_end(injected[2].clone()),
InjectedMessage::at_end(injected[3].clone()),
]);
for original in [vec![], messages()] {
let expected = [
injected[..2].to_vec(),
original.clone(),
injected[2..].to_vec(),
]
.concat();
let mut actual = original;
query.apply_injections(&mut actual);
assert_messages(&actual, &expected);
}
}
#[test]
fn indexed_injections_do_not_shift_other_original_targets() {
let original = messages();
let before = RuntimeMessage::system("before first");
let after = RuntimeMessage::system("after third");
for injections in [
vec![
InjectedMessage::before_index(1, before.clone()),
InjectedMessage::after_index(3, after.clone()),
],
vec![
InjectedMessage::after_index(3, after.clone()),
InjectedMessage::before_index(1, before.clone()),
],
] {
let query = MessageQuery::new(SessionId::new()).with_injections(injections);
let mut actual = original.clone();
query.apply_injections(&mut actual);
assert_messages(
&actual,
&[
original[0].clone(),
before.clone(),
original[1].clone(),
original[2].clone(),
original[3].clone(),
after.clone(),
original[4].clone(),
],
);
}
}
#[test]
fn window_bounds_preserve_exact_head_and_tail_after_offset() {
let original = messages();
type WindowCase = (Option<i64>, Option<i64>, Option<usize>, &'static [usize]);
let cases: &[WindowCase] = &[
(None, None, None, &[0, 1, 2, 3, 4]),
(None, Some(2), None, &[3, 4]),
(None, Some(2), Some(0), &[3, 4]),
(None, Some(2), Some(1), &[0, 3, 4]),
(None, Some(3), Some(3), &[0, 1, 2, 3, 4]),
(None, Some(1), Some(10), &[0, 1, 2, 3, 4]),
(Some(1), Some(1), Some(1), &[1, 4]),
(Some(-1), Some(2), None, &[3, 4]),
(Some(5), Some(2), Some(1), &[]),
(Some(9), None, None, &[]),
(None, Some(0), None, &[]),
(None, Some(-1), None, &[]),
(None, Some(0), Some(1), &[0]),
(None, None, Some(1), &[0, 1, 2, 3, 4]),
];
for &(offset, limit, keep_head, indices) in cases {
let query = MessageQuery {
offset,
limit,
keep_head,
..MessageQuery::new(SessionId::new())
};
let mut actual = original.clone();
query.apply_window_bounds(&mut actual);
let expected: Vec<_> = indices.iter().map(|&i| original[i].clone()).collect();
assert_messages(&actual, &expected);
}
let mut unchanged = original.clone();
let default = MessageQuery::default();
assert!(!default.has_injections());
default.apply_windowing(&mut unchanged);
default.apply_injections(&mut unchanged);
assert_messages(&unchanged, &original);
}
#[test]
fn windowing_notice_counts_offset_and_limit_without_losing_retained_messages() {
let original = messages();
let query = MessageQuery::new(SessionId::new())
.with_offset(1)
.with_limit(2)
.with_prepend_transform(Arc::new(ExcludedNoticeTransform::new("{} hidden")));
let mut actual = original.clone();
query.apply_windowing(&mut actual);
assert_eq!(actual[0].role, RuntimeMessageRole::System);
assert_eq!(actual[0].text(), Some("3 hidden"));
assert_messages(&actual[1..], &original[3..]);
let no_exclusion = MessageQuery::new(SessionId::new())
.with_prepend_transform(Arc::new(ExcludedNoticeTransform::new("{} hidden")));
let mut unchanged = original.clone();
no_exclusion.apply_windowing(&mut unchanged);
assert_messages(&unchanged, &original);
}
#[test]
fn anchored_window_respects_exact_budget_and_contiguous_recent_block() {
let costs = [5, 11, 7, 3, 2];
for (budget, recent_start) in [(0, 4), (9, 4), (10, 3), (16, 3), (17, 2), (28, 1), (100, 1)]
{
let window = anchored_window(&costs, 1, 1, None, budget);
assert_eq!(
window,
AnchoredWindow {
head_len: 1,
recent_start
}
);
assert_eq!(window.hidden(), recent_start - 1);
}
}
#[test]
fn anchored_window_caps_recent_tail_and_preserves_over_budget_anchors() {
let costs = [100, 11, 7, 3, 100];
for (minimum, maximum, budget, recent_start) in [
(1, None, 0, 4),
(2, None, 1, 3),
(5, Some(2), 1000, 3),
(1, Some(2), 1000, 3),
(0, Some(0), 0, 5),
] {
assert_eq!(
anchored_window(&costs, 1, minimum, maximum, budget),
AnchoredWindow {
head_len: 1,
recent_start
}
);
}
}
#[test]
fn anchored_window_handles_empty_and_overlapping_anchors() {
for (costs, head, tail, expected_head) in [
(&[][..], 1, 2, 0),
(&[10, 10, 10][..], 1, 10, 1),
(&[10, 10, 10][..], 9, 1, 3),
(&[10, 10, 10][..], 1, 2, 1),
] {
let window = anchored_window(costs, head, tail, None, 1);
assert_eq!(
window,
AnchoredWindow {
head_len: expected_head,
recent_start: expected_head
}
);
assert_eq!(window.hidden(), 0);
}
}
}