use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use serde_json::Value;
use crate::error::{Error, Result};
use crate::memory::{ContextProvider, SessionContext};
use crate::session::AgentSession;
use crate::types::Message;
pub trait HistoryProvider: ContextProvider {}
#[derive(PartialEq)]
enum MessageIdentity<'a> {
Id(&'a str),
Contents(&'a crate::types::Role, &'a [crate::types::Content]),
}
fn message_identity(message: &Message) -> MessageIdentity<'_> {
match real_message_id(message) {
Some(id) => MessageIdentity::Id(id),
None => MessageIdentity::Contents(&message.role, &message.contents),
}
}
fn real_message_id(message: &Message) -> Option<&str> {
message.message_id.as_deref().filter(|id| !id.is_empty())
}
pub fn filter_new_messages<'a>(existing: &[Message], incoming: &'a [Message]) -> &'a [Message] {
filter_new_messages_from(existing, incoming, StoredHistory::Complete)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StoredHistory {
Complete,
Window,
}
pub fn filter_new_messages_from<'a>(
existing: &[Message],
incoming: &'a [Message],
shape: StoredHistory,
) -> &'a [Message] {
if existing.is_empty() || incoming.len() <= existing.len() || !could_be_a_replay(existing) {
return incoming;
}
let matches = |start: usize| {
incoming[start..start + existing.len()]
.iter()
.zip(existing)
.all(|(a, b)| message_identity(a) == message_identity(b))
};
let found = match shape {
StoredHistory::Complete => matches(0).then_some(0),
StoredHistory::Window => matches(0).then_some(0).or_else(|| {
(1..(incoming.len() - existing.len()))
.rev()
.find(|s| matches(*s))
}),
};
match found {
Some(start) => &incoming[start + existing.len()..],
None => incoming,
}
}
fn could_be_a_replay(existing: &[Message]) -> bool {
existing
.iter()
.any(|m| real_message_id(m).is_some() || m.role != crate::types::Role::user())
}
pub fn new_run_messages(
existing: &[Message],
request_messages: &[Message],
response_messages: &[Message],
) -> Vec<Message> {
new_run_messages_from(
existing,
request_messages,
response_messages,
StoredHistory::Complete,
)
}
pub fn new_run_messages_from(
existing: &[Message],
request_messages: &[Message],
response_messages: &[Message],
shape: StoredHistory,
) -> Vec<Message> {
filter_new_messages_from(existing, request_messages, shape)
.iter()
.chain(response_messages)
.cloned()
.collect()
}
pub fn inject_stored_history(ctx: &mut SessionContext, stored: Vec<Message>) {
inject_stored_history_from(ctx, stored, StoredHistory::Complete)
}
pub fn inject_stored_history_from(
ctx: &mut SessionContext,
stored: Vec<Message>,
shape: StoredHistory,
) {
if stored.is_empty() {
return;
}
if filter_new_messages_from(&stored, &ctx.input_messages, shape).len()
< ctx.input_messages.len()
{
return;
}
let existing = std::mem::take(&mut ctx.messages);
ctx.messages = stored.into_iter().chain(existing).collect();
}
pub fn ensure_history_provider(session: &mut AgentSession) {
if session.service_session_id().is_none()
&& !session
.context_providers
.iter()
.any(|p| p.is_history_provider())
{
session
.context_providers
.insert(0, Arc::new(InMemoryHistoryProvider::new()));
}
}
#[derive(Default, Clone)]
pub struct InMemoryHistoryProvider {
messages: Arc<Mutex<Vec<Message>>>,
}
impl InMemoryHistoryProvider {
pub fn new() -> Self {
Self::default()
}
pub fn with_messages(messages: Vec<Message>) -> Self {
Self {
messages: Arc::new(Mutex::new(messages)),
}
}
pub fn list_messages(&self) -> Vec<Message> {
self.messages.lock().unwrap().clone()
}
pub fn to_dict(&self) -> Value {
serde_json::json!({ "messages": self.list_messages() })
}
pub fn from_dict(state: &Value) -> Result<Self> {
let messages = match state.get("messages") {
Some(v) if !v.is_null() => serde_json::from_value(v.clone()).map_err(|e| {
Error::Serialization(format!("failed to restore history provider: {e}"))
})?,
_ => Vec::new(),
};
Ok(Self::with_messages(messages))
}
}
#[async_trait]
impl ContextProvider for InMemoryHistoryProvider {
async fn before_run(&self, ctx: &mut SessionContext) -> Result<()> {
let stored = self.messages.lock().unwrap().clone();
inject_stored_history(ctx, stored);
Ok(())
}
async fn after_run(
&self,
request_messages: &[Message],
response_messages: &[Message],
error: Option<&Error>,
) -> Result<()> {
if error.is_none() {
let mut guard = self.messages.lock().unwrap();
let new = new_run_messages(&guard, request_messages, response_messages);
guard.extend(new);
}
Ok(())
}
fn is_history_provider(&self) -> bool {
true
}
}
impl HistoryProvider for InMemoryHistoryProvider {}
#[derive(Clone)]
pub struct FileHistoryProvider {
path: PathBuf,
messages: Arc<Mutex<Vec<Message>>>,
write_lock: Arc<tokio::sync::Mutex<()>>,
}
impl FileHistoryProvider {
pub fn new(path: impl Into<PathBuf>) -> Result<Self> {
let path = path.into();
let messages = if path.exists() {
let data = std::fs::read_to_string(&path)
.map_err(|e| Error::other(format!("failed to read history file {path:?}: {e}")))?;
if data.trim().is_empty() {
Vec::new()
} else {
let value: Value = serde_json::from_str(&data).map_err(|e| {
Error::Serialization(format!("failed to parse history file {path:?}: {e}"))
})?;
match value.get("messages") {
Some(v) if !v.is_null() => serde_json::from_value(v.clone()).map_err(|e| {
Error::Serialization(format!("failed to parse history file {path:?}: {e}"))
})?,
_ => Vec::new(),
}
}
} else {
Vec::new()
};
Ok(Self {
path,
messages: Arc::new(Mutex::new(messages)),
write_lock: Arc::new(tokio::sync::Mutex::new(())),
})
}
pub fn path(&self) -> &std::path::Path {
&self.path
}
pub fn list_messages(&self) -> Vec<Message> {
self.messages.lock().unwrap().clone()
}
pub fn to_dict(&self) -> Value {
serde_json::json!({ "messages": self.list_messages() })
}
async fn persist(&self, messages: &[Message]) -> Result<()> {
let dict = serde_json::json!({ "messages": messages });
let json = serde_json::to_string_pretty(&dict)
.map_err(|e| Error::Serialization(format!("failed to serialize history: {e}")))?;
let file_name = self
.path
.file_name()
.and_then(|f| f.to_str())
.unwrap_or("history.json");
let tmp = self
.path
.with_file_name(format!("{file_name}.tmp.{}", uuid::Uuid::new_v4()));
if let Err(e) = tokio::fs::write(&tmp, &json).await {
let _ = tokio::fs::remove_file(&tmp).await;
return Err(Error::other(format!(
"failed to write history temp file {tmp:?}: {e}"
)));
}
tokio::fs::rename(&tmp, &self.path).await.map_err(|e| {
let tmp = tmp.clone();
tokio::spawn(async move {
let _ = tokio::fs::remove_file(&tmp).await;
});
Error::other(format!(
"failed to finalize history file {:?}: {e}",
self.path
))
})
}
}
#[async_trait]
impl ContextProvider for FileHistoryProvider {
async fn before_run(&self, ctx: &mut SessionContext) -> Result<()> {
let stored = self.messages.lock().unwrap().clone();
inject_stored_history(ctx, stored);
Ok(())
}
async fn after_run(
&self,
request_messages: &[Message],
response_messages: &[Message],
error: Option<&Error>,
) -> Result<()> {
if error.is_some() {
return Ok(());
}
let _write = self.write_lock.lock().await;
let snapshot = {
let guard = self.messages.lock().unwrap();
let mut next = guard.clone();
next.extend(new_run_messages(
&guard,
request_messages,
response_messages,
));
next
};
self.persist(&snapshot).await?;
*self.messages.lock().unwrap() = snapshot;
Ok(())
}
fn is_history_provider(&self) -> bool {
true
}
}
impl HistoryProvider for FileHistoryProvider {}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Message;
#[tokio::test]
async fn before_run_prepends_stored_messages_ahead_of_existing_context_messages() {
let provider = InMemoryHistoryProvider::with_messages(vec![
Message::user("q1"),
Message::assistant("a1"),
]);
let mut ctx = SessionContext::new(vec![Message::user("q2")]);
ctx.messages
.push(Message::system("injected by another provider"));
provider.before_run(&mut ctx).await.unwrap();
let texts: Vec<String> = ctx.messages.iter().map(|m| m.text()).collect();
assert_eq!(
texts,
vec![
"q1".to_string(),
"a1".to_string(),
"injected by another provider".to_string(),
]
);
}
#[tokio::test]
async fn after_run_appends_only_on_success() {
let provider = InMemoryHistoryProvider::new();
provider
.after_run(&[Message::user("hi")], &[Message::assistant("hello")], None)
.await
.unwrap();
assert_eq!(provider.list_messages().len(), 2);
provider
.after_run(
&[Message::user("again")],
&[],
Some(&Error::service("boom")),
)
.await
.unwrap();
assert_eq!(provider.list_messages().len(), 2);
}
#[test]
fn to_dict_from_dict_round_trips_messages() {
let provider = InMemoryHistoryProvider::with_messages(vec![
Message::user("q1"),
Message::assistant("a1"),
]);
let state = provider.to_dict();
let restored = InMemoryHistoryProvider::from_dict(&state).unwrap();
let msgs = restored.list_messages();
assert_eq!(msgs.len(), 2);
assert_eq!(msgs[0].text(), "q1");
assert_eq!(msgs[1].text(), "a1");
}
#[test]
fn from_dict_tolerates_a_missing_messages_key() {
let restored = InMemoryHistoryProvider::from_dict(&serde_json::json!({})).unwrap();
assert!(restored.list_messages().is_empty());
}
#[test]
fn ensure_history_provider_attaches_once_and_skips_service_managed() {
let mut local = AgentSession::new();
ensure_history_provider(&mut local);
assert_eq!(local.context_providers.len(), 1);
assert!(local.context_providers[0].is_history_provider());
ensure_history_provider(&mut local);
assert_eq!(local.context_providers.len(), 1);
let mut service = AgentSession::service("svc-1");
ensure_history_provider(&mut service);
assert!(service.context_providers.is_empty());
}
#[tokio::test]
async fn file_history_provider_persists_and_reloads() {
let dir = std::env::temp_dir().join(format!("afr-history-test-{}", uuid::Uuid::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("history.json");
let provider = FileHistoryProvider::new(&path).unwrap();
assert!(provider.list_messages().is_empty());
provider
.after_run(&[Message::user("hi")], &[Message::assistant("hello")], None)
.await
.unwrap();
assert_eq!(provider.list_messages().len(), 2);
let reloaded = FileHistoryProvider::new(&path).unwrap();
let msgs = reloaded.list_messages();
assert_eq!(msgs.len(), 2);
assert_eq!(msgs[0].text(), "hi");
assert_eq!(msgs[1].text(), "hello");
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn file_history_provider_concurrent_runs_do_not_lose_messages() {
let dir = std::env::temp_dir().join(format!("afr-history-conc-{}", uuid::Uuid::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("history.json");
let provider = FileHistoryProvider::new(&path).unwrap();
const N: usize = 50;
let mut handles = Vec::new();
for i in 0..N {
let p = provider.clone();
handles.push(tokio::spawn(async move {
p.after_run(
&[Message::user(format!("q{i}"))],
&[Message::assistant(format!("a{i}"))],
None,
)
.await
.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
assert_eq!(provider.list_messages().len(), N * 2);
let reloaded = FileHistoryProvider::new(&path).unwrap();
assert_eq!(reloaded.list_messages().len(), N * 2);
let leftover: Vec<_> = std::fs::read_dir(&dir)
.unwrap()
.filter_map(|e| e.ok())
.filter(|e| e.file_name().to_string_lossy().contains(".tmp."))
.collect();
assert!(leftover.is_empty(), "temp files leaked: {leftover:?}");
std::fs::remove_dir_all(&dir).ok();
}
fn texts(messages: &[Message]) -> Vec<String> {
messages.iter().map(Message::text).collect()
}
#[test]
fn filter_new_messages_returns_everything_when_nothing_is_stored() {
let incoming = vec![Message::user("hi"), Message::assistant("hello")];
assert_eq!(filter_new_messages(&[], &incoming).len(), 2);
}
#[test]
fn filter_new_messages_drops_a_replayed_prefix() {
let existing = vec![Message::user("hi"), Message::assistant("hello")];
let incoming = vec![
Message::user("hi"),
Message::assistant("hello"),
Message::user("more"),
Message::assistant("sure"),
];
assert_eq!(
texts(filter_new_messages(&existing, &incoming)),
vec!["more".to_string(), "sure".to_string()]
);
}
#[test]
fn filter_new_messages_aligns_a_trimmed_window() {
let existing = vec![Message::user("q2"), Message::assistant("a2")];
let incoming = vec![
Message::user("q1"),
Message::assistant("a1"),
Message::user("q2"),
Message::assistant("a2"),
Message::user("q3"),
];
assert_eq!(
texts(filter_new_messages_from(
&existing,
&incoming,
StoredHistory::Window
)),
vec!["q3".to_string()]
);
assert_eq!(filter_new_messages(&existing, &incoming).len(), 5);
}
#[test]
fn a_window_prefers_an_anchored_match_over_a_later_one() {
let existing = vec![Message::user("q"), Message::assistant("a")];
let incoming = vec![
Message::user("q"),
Message::assistant("a"),
Message::user("filler"),
Message::user("q"),
Message::assistant("a"),
Message::user("new"),
];
assert_eq!(
texts(filter_new_messages_from(
&existing,
&incoming,
StoredHistory::Window
)),
vec![
"filler".to_string(),
"q".to_string(),
"a".to_string(),
"new".to_string()
],
"an at-cap list that opens the replay is a complete history, not a window"
);
let earlier = vec![
Message::user("q0"),
Message::assistant("a0"),
Message::user("q"),
Message::assistant("a"),
Message::user("filler"),
Message::user("q"),
Message::assistant("a"),
Message::user("new"),
];
assert_eq!(
texts(filter_new_messages_from(
&existing,
&earlier,
StoredHistory::Window
)),
vec!["new".to_string()]
);
}
#[test]
fn an_empty_message_id_is_not_an_identity() {
let empty_id = |m: Message| Message {
message_id: Some(String::new()),
..m
};
let existing = vec![
empty_id(Message::user("q")),
empty_id(Message::assistant("a")),
];
let incoming = vec![
empty_id(Message::user("something else")),
empty_id(Message::assistant("unrelated")),
empty_id(Message::user("new")),
];
assert_eq!(
texts(filter_new_messages(&existing, &incoming)).len(),
3,
"empty ids must fall back to role and contents, not match everything"
);
let user_only = vec![empty_id(Message::user("yes"))];
let repeat = vec![empty_id(Message::user("yes")), Message::user("question")];
assert_eq!(texts(filter_new_messages(&user_only, &repeat)).len(), 2);
}
#[test]
fn id_less_user_only_history_is_never_read_as_a_replay() {
let existing = vec![Message::user("yes")];
let incoming = vec![Message::user("yes"), Message::user("question")];
assert_eq!(
texts(filter_new_messages(&existing, &incoming)),
vec!["yes".to_string(), "question".to_string()],
"saying 'yes' again is not a replay of having said it"
);
let with_reply = vec![Message::user("yes"), Message::assistant("go on")];
let replayed = vec![
Message::user("yes"),
Message::assistant("go on"),
Message::user("question"),
];
assert_eq!(
texts(filter_new_messages(&with_reply, &replayed)),
vec!["question".to_string()]
);
let with_id = vec![Message {
message_id: Some("m1".to_string()),
..Message::user("yes")
}];
let replayed_by_id = vec![
Message {
message_id: Some("m1".to_string()),
..Message::user("yes")
},
Message::user("question"),
];
assert_eq!(
texts(filter_new_messages(&with_id, &replayed_by_id)),
vec!["question".to_string()]
);
}
#[test]
fn a_complete_history_never_matches_past_the_start() {
let existing = vec![Message::user("yes")];
let incoming = vec![
Message::user("preface"),
Message::user("yes"),
Message::user("question"),
];
assert_eq!(
texts(filter_new_messages(&existing, &incoming)),
vec![
"preface".to_string(),
"yes".to_string(),
"question".to_string()
],
"nothing may be dropped: the stored 'yes' is not where this input starts"
);
}
#[test]
fn a_complete_store_keeps_the_turns_between_two_occurrences() {
let existing = vec![Message::user("q"), Message::assistant("a")];
let incoming = vec![
Message::user("q"),
Message::assistant("a"),
Message::user("filler"),
Message::user("q"),
Message::assistant("a"),
Message::user("new"),
];
assert_eq!(
texts(filter_new_messages_from(
&existing,
&incoming,
StoredHistory::Complete
)),
vec![
"filler".to_string(),
"q".to_string(),
"a".to_string(),
"new".to_string()
]
);
}
#[test]
fn filter_new_messages_keeps_a_repeated_turn_it_cannot_align() {
let existing = vec![Message::user("ping"), Message::assistant("pong")];
let incoming = vec![Message::user("ping"), Message::assistant("pong!")];
assert_eq!(
texts(filter_new_messages(&existing, &incoming)),
vec!["ping".to_string(), "pong!".to_string()]
);
}
#[test]
fn filter_new_messages_matches_on_message_id_when_present() {
let with_id = |m: Message, id: &str| Message {
message_id: Some(id.to_string()),
..m
};
let stored = with_id(Message::user("hi"), "m1");
let replayed = with_id(Message::user("edited after the fact"), "m1");
let incoming = vec![replayed, Message::assistant("hello")];
assert_eq!(
texts(filter_new_messages(
std::slice::from_ref(&stored),
&incoming
)),
vec!["hello".to_string()]
);
let other = vec![with_id(Message::user("hi"), "m2")];
assert_eq!(
filter_new_messages(std::slice::from_ref(&stored), &other).len(),
1
);
}
#[tokio::test]
async fn replaying_the_transcript_does_not_duplicate_stored_history() {
let provider = InMemoryHistoryProvider::new();
provider
.after_run(&[Message::user("q1")], &[Message::assistant("a1")], None)
.await
.unwrap();
provider
.after_run(
&[
Message::user("q1"),
Message::assistant("a1"),
Message::user("q2"),
],
&[Message::assistant("a2")],
None,
)
.await
.unwrap();
assert_eq!(
texts(&provider.list_messages()),
vec![
"q1".to_string(),
"a1".to_string(),
"q2".to_string(),
"a2".to_string()
]
);
}
#[tokio::test]
async fn a_response_repeating_stored_history_is_still_stored() {
let provider = InMemoryHistoryProvider::with_messages(vec![
Message::user("q"),
Message::assistant("a"),
]);
provider
.after_run(
&[Message::user("q")],
&[Message::assistant("a"), Message::assistant("b")],
None,
)
.await
.unwrap();
assert_eq!(
texts(&provider.list_messages()),
vec![
"q".to_string(),
"a".to_string(),
"q".to_string(),
"a".to_string(),
"b".to_string()
],
"the new turn must not be swallowed by its own response"
);
}
#[tokio::test]
async fn before_run_does_not_prepend_history_the_input_already_carries() {
let provider = InMemoryHistoryProvider::with_messages(vec![
Message::user("q1"),
Message::assistant("a1"),
]);
let mut replayed = SessionContext::new(vec![
Message::user("q1"),
Message::assistant("a1"),
Message::user("q2"),
]);
provider.before_run(&mut replayed).await.unwrap();
assert!(
replayed.messages.is_empty(),
"stored history must not be injected on top of a replay of itself: {:?}",
texts(&replayed.messages)
);
let mut incremental = SessionContext::new(vec![Message::user("q2")]);
provider.before_run(&mut incremental).await.unwrap();
assert_eq!(
texts(&incremental.messages),
vec!["q1".to_string(), "a1".to_string()]
);
}
#[tokio::test]
async fn append_only_runs_still_accumulate_every_turn() {
let provider = InMemoryHistoryProvider::new();
for _ in 0..3 {
provider
.after_run(
&[Message::user("ping")],
&[Message::assistant("pong")],
None,
)
.await
.unwrap();
}
assert_eq!(provider.list_messages().len(), 6);
}
#[tokio::test]
async fn file_history_provider_does_not_duplicate_a_replayed_transcript() {
let dir = std::env::temp_dir().join(format!("afr-history-dedup-{}", uuid::Uuid::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("history.json");
let provider = FileHistoryProvider::new(&path).unwrap();
provider
.after_run(&[Message::user("q1")], &[Message::assistant("a1")], None)
.await
.unwrap();
provider
.after_run(
&[
Message::user("q1"),
Message::assistant("a1"),
Message::user("q2"),
],
&[Message::assistant("a2")],
None,
)
.await
.unwrap();
assert_eq!(texts(&provider.list_messages()).len(), 4);
let reloaded = FileHistoryProvider::new(&path).unwrap();
assert_eq!(
texts(&reloaded.list_messages()),
vec![
"q1".to_string(),
"a1".to_string(),
"q2".to_string(),
"a2".to_string()
]
);
std::fs::remove_dir_all(&dir).ok();
}
}