#![allow(dead_code)]
use crate::agent_definition::AgentDefinition;
use crate::harness_definition::HarnessDefinition;
use crate::session::ExecutionSession;
use crate::typed_id::{AgentId, EventId, HarnessId, MessageId, SessionId};
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use uuid::Uuid;
use crate::error::Result;
use crate::message::Message;
use crate::message_filter::MessageQuery;
use crate::message_retriever::{InputMessage, MessageHistory, MessageRetriever};
use crate::{
execution_loading::AgentStore, execution_loading::HarnessStore, execution_loading::SessionStore,
};
use chrono::Utc;
#[derive(Debug, Default, Clone)]
pub(crate) struct TestMessageRetriever {
messages: Arc<RwLock<HashMap<SessionId, Vec<Message>>>>,
}
impl TestMessageRetriever {
pub(crate) fn new() -> Self {
Self {
messages: Arc::new(RwLock::new(HashMap::new())),
}
}
pub(crate) async fn sessions(&self) -> Vec<SessionId> {
self.messages.read().await.keys().copied().collect()
}
pub(crate) async fn clear(&self) {
self.messages.write().await.clear();
}
pub(crate) async fn clear_session(&self, session_id: SessionId) {
self.messages.write().await.remove(&session_id);
}
pub(crate) async fn seed(&self, session_id: SessionId, messages: Vec<Message>) {
self.messages.write().await.insert(session_id, messages);
}
pub(crate) async fn add(&self, session_id: SessionId, input: InputMessage) -> Result<Message> {
let message = Message {
id: MessageId::new(),
role: input.role,
content: input.content,
phase: None,
phase_source: None,
controls: input.controls,
metadata: input.metadata,
external_actor: None,
created_at: Utc::now(),
};
self.messages
.write()
.await
.entry(session_id)
.or_default()
.push(message.clone());
Ok(message)
}
pub(crate) async fn store(&self, session_id: SessionId, message: Message) -> Result<()> {
self.messages
.write()
.await
.entry(session_id)
.or_default()
.push(message);
Ok(())
}
}
#[async_trait]
impl MessageRetriever for TestMessageRetriever {
async fn get(&self, session_id: SessionId, message_id: MessageId) -> Result<Option<Message>> {
Ok(self
.messages
.read()
.await
.get(&session_id)
.and_then(|messages| messages.iter().find(|m| m.id == message_id).cloned()))
}
async fn load(&self, session_id: SessionId) -> Result<Vec<Message>> {
Ok(self
.messages
.read()
.await
.get(&session_id)
.cloned()
.unwrap_or_default())
}
async fn load_filtered(&self, query: MessageQuery) -> Result<Vec<Message>> {
use crate::message_filter::MessageFilter;
let mut messages = self.load(query.session_id).await?;
if let Some(after) = query.after_sequence {
messages = messages.into_iter().skip(after.max(0) as usize).collect();
}
for filter in &query.filters {
match filter {
MessageFilter::TimeRange { from, to } => {
messages.retain(|m| {
let after_from = from.is_none_or(|t| m.created_at >= t);
let before_to = to.is_none_or(|t| m.created_at <= t);
after_from && before_to
});
}
MessageFilter::Search(q) => {
let q_lower = q.to_lowercase();
messages.retain(|m| {
m.text()
.is_some_and(|t| t.to_lowercase().contains(&q_lower))
});
}
MessageFilter::Custom(predicate) => {
messages.retain(|m| predicate(m));
}
_ => {}
}
}
query.apply_windowing(&mut messages);
if query.has_injections() {
query.apply_injections(&mut messages);
}
Ok(messages)
}
async fn load_filtered_history(&self, query: MessageQuery) -> Result<MessageHistory> {
let source_sequence = self
.messages
.read()
.await
.get(&query.session_id)
.map(|messages| messages.len() as i64)
.unwrap_or(0);
Ok(MessageHistory {
messages: self.load_filtered(query).await?,
source_sequence: Some(source_sequence),
})
}
async fn count(&self, session_id: SessionId) -> Result<usize> {
Ok(self
.messages
.read()
.await
.get(&session_id)
.map(|m| m.len())
.unwrap_or(0))
}
}
#[derive(Debug, Default, Clone)]
pub(crate) struct TestAgentStore {
agents: Arc<RwLock<HashMap<AgentId, AgentDefinition>>>,
}
impl TestAgentStore {
pub(crate) fn new() -> Self {
Self {
agents: Arc::new(RwLock::new(HashMap::new())),
}
}
pub(crate) async fn add_agent(&self, agent: AgentDefinition) {
self.agents.write().await.insert(agent.id, agent);
}
pub(crate) async fn agent_ids(&self) -> Vec<AgentId> {
self.agents.read().await.keys().copied().collect()
}
pub(crate) async fn clear(&self) {
self.agents.write().await.clear();
}
}
#[async_trait]
impl AgentStore for TestAgentStore {
async fn get_agent(&self, agent_id: AgentId) -> Result<Option<AgentDefinition>> {
Ok(self.agents.read().await.get(&agent_id).cloned())
}
}
#[derive(Debug, Default, Clone)]
pub(crate) struct TestHarnessStore {
harnesses: Arc<RwLock<HashMap<HarnessId, HarnessDefinition>>>,
}
impl TestHarnessStore {
pub(crate) fn new() -> Self {
Self {
harnesses: Arc::new(RwLock::new(HashMap::new())),
}
}
pub(crate) async fn add_harness(&self, harness_id: HarnessId, harness: HarnessDefinition) {
self.harnesses.write().await.insert(harness_id, harness);
}
}
#[async_trait]
impl HarnessStore for TestHarnessStore {
async fn get_harness(&self, harness_id: HarnessId) -> Result<Option<HarnessDefinition>> {
Ok(self.harnesses.read().await.get(&harness_id).cloned())
}
}
#[derive(Debug, Default, Clone)]
pub(crate) struct TestSessionStore {
sessions: Arc<RwLock<HashMap<SessionId, ExecutionSession>>>,
}
impl TestSessionStore {
pub(crate) fn new() -> Self {
Self {
sessions: Arc::new(RwLock::new(HashMap::new())),
}
}
pub(crate) async fn add_session(&self, session: ExecutionSession) {
self.sessions.write().await.insert(session.id, session);
}
pub(crate) async fn session_ids(&self) -> Vec<SessionId> {
self.sessions.read().await.keys().copied().collect()
}
pub(crate) async fn clear(&self) {
self.sessions.write().await.clear();
}
}
#[async_trait]
impl SessionStore for TestSessionStore {
async fn get_session(&self, session_id: SessionId) -> Result<Option<ExecutionSession>> {
Ok(self.sessions.read().await.get(&session_id).cloned())
}
}
use crate::event_emitter::EventEmitter;
use crate::events::{Event, EventRequest};
#[derive(Debug, Default, Clone)]
pub(crate) struct TestEventEmitter {
events: Arc<RwLock<Vec<Event>>>,
sequence: Arc<RwLock<i32>>,
}
impl TestEventEmitter {
pub(crate) fn new() -> Self {
Self {
events: Arc::new(RwLock::new(Vec::new())),
sequence: Arc::new(RwLock::new(0)),
}
}
pub(crate) async fn events(&self) -> Vec<Event> {
self.events.read().await.clone()
}
pub(crate) async fn event_count(&self) -> usize {
self.events.read().await.len()
}
pub(crate) async fn clear(&self) {
self.events.write().await.clear();
*self.sequence.write().await = 0;
}
pub(crate) async fn events_by_type(&self, event_type: &str) -> Vec<Event> {
self.events
.read()
.await
.iter()
.filter(|e| e.event_type == event_type)
.cloned()
.collect()
}
pub(crate) async fn events_for_session(&self, session_id: Uuid) -> Vec<Event> {
self.events
.read()
.await
.iter()
.filter(|e| e.session_uuid() == session_id)
.cloned()
.collect()
}
}
#[async_trait]
impl EventEmitter for TestEventEmitter {
async fn emit(&self, request: EventRequest) -> Result<Event> {
let mut sequence = self.sequence.write().await;
*sequence += 1;
let seq = *sequence;
drop(sequence);
let event = request.into_event(EventId::new(), seq);
self.events.write().await.push(event.clone());
Ok(event)
}
}
#[cfg(test)]
mod tests {
use super::*;
use uuid::Uuid;
#[tokio::test]
async fn test_in_memory_message_retriever() {
let store = TestMessageRetriever::new();
let session_id: SessionId = Uuid::now_v7().into();
store
.store(session_id, Message::user("Hello"))
.await
.unwrap();
let messages = store.load(session_id).await.unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].text(), Some("Hello"));
}
#[tokio::test]
async fn test_in_memory_message_retriever_add_and_get() {
let store = TestMessageRetriever::new();
let session_id: SessionId = Uuid::now_v7().into();
let message = store
.add(session_id, InputMessage::user("Hello via add"))
.await
.unwrap();
let retrieved = store.get(session_id, message.id).await.unwrap();
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().text(), Some("Hello via add"));
let missing = store.get(session_id, MessageId::new()).await.unwrap();
assert!(missing.is_none());
}
#[tokio::test]
async fn test_message_retriever_add_returns_consistent_id() {
let store = TestMessageRetriever::new();
let session_id: SessionId = Uuid::now_v7().into();
let added = store
.add(session_id, InputMessage::user("Test consistency"))
.await
.unwrap();
let retrieved = store.get(session_id, added.id).await.unwrap();
assert!(
retrieved.is_some(),
"Message must be retrievable by the ID returned from add()"
);
let retrieved = retrieved.unwrap();
assert_eq!(
retrieved.id, added.id,
"Retrieved message ID must match the ID returned from add()"
);
let all_messages = store.load(session_id).await.unwrap();
let found = all_messages.iter().find(|m| m.id == added.id);
assert!(
found.is_some(),
"Message with returned ID must appear in load() results"
);
}
#[tokio::test]
async fn test_in_memory_event_emitter() {
use crate::events::{EventContext, EventRequest, InputMessageData};
let emitter = TestEventEmitter::new();
let session_id: SessionId = Uuid::now_v7().into();
let event_context = EventContext::empty();
let event1 = emitter
.emit(EventRequest::new(
session_id,
event_context.clone(),
InputMessageData::new(Message::user("test1")),
))
.await
.unwrap();
assert_eq!(event1.sequence, Some(1));
let event2 = emitter
.emit(EventRequest::new(
session_id,
event_context,
InputMessageData::new(Message::user("test2")),
))
.await
.unwrap();
assert_eq!(event2.sequence, Some(2));
let events = emitter.events().await;
assert_eq!(events.len(), 2);
assert_eq!(emitter.event_count().await, 2);
}
#[tokio::test]
async fn test_in_memory_event_emitter_filter_by_type() {
use crate::events::{
EventContext, EventRequest, INPUT_MESSAGE, InputMessageData, REASON_STARTED,
ReasonStartedData,
};
let emitter = TestEventEmitter::new();
let session_id: SessionId = Uuid::now_v7().into();
let event_context = EventContext::empty();
emitter
.emit(EventRequest::new(
session_id,
event_context.clone(),
InputMessageData::new(Message::user("test")),
))
.await
.unwrap();
emitter
.emit(EventRequest::new(
session_id,
event_context,
ReasonStartedData {
harness_id: HarnessId::from_seed(1),
agent_id: Some(AgentId::new()),
metadata: None,
},
))
.await
.unwrap();
let received_events = emitter.events_by_type(INPUT_MESSAGE).await;
assert_eq!(received_events.len(), 1);
let started_events = emitter.events_by_type(REASON_STARTED).await;
assert_eq!(started_events.len(), 1);
}
#[tokio::test]
async fn test_in_memory_event_emitter_filter_by_session() {
use crate::events::{EventContext, EventRequest, InputMessageData};
let emitter = TestEventEmitter::new();
let session1: SessionId = Uuid::now_v7().into();
let session2: SessionId = Uuid::now_v7().into();
let context = EventContext::empty();
emitter
.emit(EventRequest::new(
session1,
context.clone(),
InputMessageData::new(Message::user("session1")),
))
.await
.unwrap();
emitter
.emit(EventRequest::new(
session2,
context,
InputMessageData::new(Message::user("session2")),
))
.await
.unwrap();
let session1_events = emitter.events_for_session(session1.uuid()).await;
assert_eq!(session1_events.len(), 1);
let session2_events = emitter.events_for_session(session2.uuid()).await;
assert_eq!(session2_events.len(), 1);
}
#[tokio::test]
async fn test_in_memory_event_emitter_clear() {
use crate::events::{EventContext, EventRequest, InputMessageData};
let emitter = TestEventEmitter::new();
let session_id: SessionId = Uuid::now_v7().into();
let event_context = EventContext::empty();
emitter
.emit(EventRequest::new(
session_id,
event_context,
InputMessageData::new(Message::user("test")),
))
.await
.unwrap();
assert_eq!(emitter.event_count().await, 1);
emitter.clear().await;
assert_eq!(emitter.event_count().await, 0);
}
}