pub mod magi_adapter;
pub mod magi_wiring;
pub mod messages;
pub mod provider;
use crate::agent::messages::{Content, Message, Role};
use crate::agent::provider::{Provider, ResponseChunk};
use crate::memory::clock::Clock;
use crate::memory::config::MemoryConfig;
use crate::memory::context::assemble_selective;
use crate::memory::decay::purge_expired_archives;
use crate::memory::embedding::EmbeddingProvider;
use crate::memory::error::MemoryError;
use crate::memory::judge::LlmDistillJudge;
use crate::memory::profile::{distill, render_profile};
use crate::memory::retrieval::reembed_pending;
use crate::memory::salience::assign_salience;
use crate::memory::store::{Memory, SqliteVectorStore, VectorStore};
use crate::memory::MemoryKind;
use crate::system::database::MemoryStore;
use crate::tools::Tool;
use anyhow::Result;
use futures::StreamExt;
use sha2::{Digest, Sha256};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::oneshot;
use tokio::time::{timeout, Duration};
use tokio_util::sync::CancellationToken;
static TURN_SEQ: AtomicU64 = AtomicU64::new(0);
const APPROVAL_TIMEOUT_SECS: u64 = 300;
pub const DEFAULT_MAX_TOOL_CALLS: usize = 15;
pub const MAX_TOOL_CALLS_ERROR: &str = "Maximum tool call limit reached";
const CONSULT_TOOL_NAME: &str = "consult";
const FORCED_CONSULT_TOOL_USE_ID: &str = "forced-consult";
const CONSULT_ALREADY_FORCED_MESSAGE: &str =
"consult already ran once for this forced query; no further invocations";
pub trait RunObserver: Send + Sync {
fn authorize(&self, tool_name: &str) -> bool;
fn on_tool_call(
&self,
id: &str,
name: &str,
input: &serde_json::Value,
result: &str,
ok: bool,
ms: u64,
);
fn on_final_turn(&self, text_block_count: usize);
fn on_usage(&self, _input_tokens: u64, _output_tokens: u64) {}
}
pub struct AgentRunConfig {
pub max_tool_calls: usize,
pub disable_repetitive_guard: bool,
pub observer: Option<Arc<dyn RunObserver>>,
pub cancel: CancellationToken,
pub system: Option<String>,
pub force_consult: bool,
}
impl Default for AgentRunConfig {
fn default() -> Self {
Self {
max_tool_calls: DEFAULT_MAX_TOOL_CALLS,
disable_repetitive_guard: false,
observer: None,
cancel: CancellationToken::new(),
system: None,
force_consult: false,
}
}
}
pub struct ApprovalRequest {
pub tool_name: String,
#[allow(dead_code)]
pub input: serde_json::Value,
pub tx: oneshot::Sender<bool>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum StreamPiece {
Content(String),
Reasoning(String),
Notice(String),
}
struct MemorySubsystem {
store: Arc<SqliteVectorStore>,
embedder: Arc<dyn EmbeddingProvider>,
clock: Arc<dyn Clock>,
cfg: MemoryConfig,
scope: String,
turns_since_open: usize,
}
pub struct Agent {
provider: Arc<dyn Provider>,
tools: Vec<Box<dyn Tool>>,
history: Vec<Message>,
memory: Option<Arc<dyn MemoryStore>>,
session_id: Option<String>,
approval_tx: Option<tokio::sync::mpsc::Sender<ApprovalRequest>>,
memory_subsystem: Option<MemorySubsystem>,
}
impl Agent {
pub fn new(provider: Arc<dyn Provider>) -> Self {
Self {
provider,
tools: Vec::new(),
history: Vec::new(),
memory: None,
session_id: None,
approval_tx: None,
memory_subsystem: None,
}
}
pub fn set_memory(&mut self, memory: Arc<dyn MemoryStore>, session_id: String) {
self.memory = Some(memory);
self.session_id = Some(session_id);
}
pub fn set_provider(&mut self, provider: Arc<dyn Provider>) {
self.provider = provider;
}
pub fn provider_is_static(&self) -> bool {
self.provider.is_static()
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn history(&self) -> &[Message] {
&self.history
}
pub async fn load_history(&mut self) -> Result<()> {
if let (Some(memory), Some(sid)) = (&self.memory, &self.session_id) {
let messages = memory.get_messages(sid).await?;
self.history = messages;
}
Ok(())
}
pub fn set_approval_channel(&mut self, tx: tokio::sync::mpsc::Sender<ApprovalRequest>) {
self.approval_tx = Some(tx);
}
pub fn register_tool(&mut self, tool: Box<dyn Tool>) {
self.tools.push(tool);
}
pub fn register_or_replace_tool(&mut self, tool: Box<dyn Tool>) {
if let Some(slot) = self.tools.iter_mut().find(|t| t.name() == tool.name()) {
*slot = tool;
} else {
self.tools.push(tool);
}
}
pub fn normalize_input(val: &serde_json::Value, depth: usize) -> Result<String> {
const MAX_DEPTH: usize = 10;
if depth > MAX_DEPTH {
return Err(anyhow::anyhow!("JSON nesting limit exceeded"));
}
match val {
serde_json::Value::Object(map) => {
let mut sorted_entries: Vec<_> = map.iter().collect();
sorted_entries.sort_by_key(|(k, _)| *k);
let mut parts = Vec::new();
for (k, v) in sorted_entries {
parts.push(format!("{}:{}", k, Self::normalize_input(v, depth + 1)?));
}
Ok(format!("{{{}}}", parts.join(",")))
}
serde_json::Value::Array(arr) => {
let mut normalized_elements = Vec::new();
for v in arr {
normalized_elements.push(Self::normalize_input(v, depth + 1)?);
}
Ok(format!("[{}]", normalized_elements.join(",")))
}
serde_json::Value::String(s) => Ok(s.trim().to_string()),
_ => Ok(val.to_string()),
}
}
pub fn sanitize_text(text: &str) -> String {
let mut result = String::with_capacity(text.len());
let mut chars = text.chars().peekable();
while let Some(c) = chars.next() {
if c == '\x1B' {
if let Some('[') = chars.peek() {
chars.next();
for next in chars.by_ref() {
if next.is_ascii_alphabetic() {
break;
}
}
continue;
}
}
if c.is_control() && c != '\n' && c != '\r' && c != '\t' {
continue;
}
result.push(c);
}
result
}
pub fn clear_history(&mut self) {
self.history.clear();
}
pub fn set_memory_subsystem(
&mut self,
store: Arc<SqliteVectorStore>,
embedder: Arc<dyn EmbeddingProvider>,
clock: Arc<dyn Clock>,
cfg: MemoryConfig,
) {
self.memory_subsystem = Some(MemorySubsystem {
store,
embedder,
clock,
cfg,
scope: "root".into(),
turns_since_open: 0,
});
}
pub async fn on_session_open(&self) -> Result<()> {
if let Some(s) = &self.memory_subsystem {
let _ = reembed_pending(&*s.store, &*s.embedder, &s.cfg, &s.scope).await;
let _ = purge_expired_archives(&*s.store, &*s.clock, &s.cfg).await;
}
Ok(())
}
pub async fn on_session_close(&self) -> anyhow::Result<()> {
let Some(sub) = self.memory_subsystem.as_ref() else {
return Ok(());
};
if !sub.cfg.distill_on_session_close || !sub.cfg.distill_enabled {
return Ok(());
}
let judge = LlmDistillJudge::new(self.provider.clone());
let _ = distill(
&*sub.store,
&judge,
&*sub.embedder,
&*sub.clock,
&sub.cfg,
&sub.scope,
)
.await
.map_err(|e| eprintln!("on_session_close distill: {e}"));
Ok(())
}
pub async fn query_streaming(
&mut self,
text: &str,
chunk_tx: tokio::sync::mpsc::Sender<StreamPiece>,
config: AgentRunConfig,
) -> Result<String> {
let user_msg = Message::user(text);
self.history.push(user_msg.clone());
if let (Some(memory), Some(sid)) = (&self.memory, &self.session_id) {
memory.add_message(sid, &user_msg).await?;
}
let selective_snapshot = self.memory_subsystem.as_ref().and_then(|s| {
(s.cfg.mode == "selective").then(|| {
(
s.store.clone(),
s.embedder.clone(),
s.clock.clone(),
s.cfg.clone(),
s.scope.clone(),
)
})
});
if let Some((store, embedder, clock, cfg, scope)) = selective_snapshot {
let session_id_str = self.session_id.clone().unwrap_or_default();
let live_profile = render_profile(&*store, &cfg, &scope)
.await
.unwrap_or_default();
let (working_messages, assembly_notices) = match assemble_selective(
&*store,
&*embedder,
&*clock,
&cfg,
"",
&live_profile,
&user_msg,
&scope,
)
.await
{
Ok(assembled) => (assembled.messages, assembled.notices),
Err(MemoryError::BudgetUnsatisfiable) => {
return Err(anyhow::anyhow!("context assembly failed: budget unsatisfiable (system+profile exceed budget — check config)"));
}
Err(e) => {
let summary = summarize_assembly_error(&e);
let _ = chunk_tx.send(StreamPiece::Notice(summary)).await;
(self.history.clone(), vec![])
}
};
for notice in assembly_notices {
let _ = chunk_tx.send(StreamPiece::Notice(notice)).await;
}
write_turn_to_memory(
&store,
&*embedder,
&*clock,
&cfg,
&scope,
&session_id_str,
text,
Role::User,
Some(chunk_tx.clone()),
)
.await;
let (full_text, final_text) = self
.run_tool_loop(working_messages, &chunk_tx, &config, text)
.await?;
write_turn_to_memory(
&store,
&*embedder,
&*clock,
&cfg,
&scope,
&session_id_str,
&full_text,
Role::Assistant,
Some(chunk_tx.clone()),
)
.await;
let should_distill = if let Some(sub) = self.memory_subsystem.as_mut() {
sub.turns_since_open = sub.turns_since_open.saturating_add(1);
let n = sub.cfg.distill_every_n_turns;
sub.cfg.distill_enabled && n > 0 && sub.turns_since_open.is_multiple_of(n)
} else {
false
};
if should_distill {
let judge = LlmDistillJudge::new(self.provider.clone());
if let Err(e) = distill(&*store, &judge, &*embedder, &*clock, &cfg, &scope).await {
let _ = chunk_tx
.send(StreamPiece::Notice(format!(
"memory: distillation failed (non-fatal) — {}",
e
)))
.await;
}
}
return Ok(final_text);
}
let working = self.history.clone();
let (_full_text, final_text) = self
.run_tool_loop(working, &chunk_tx, &config, text)
.await?;
Ok(final_text)
}
async fn authorize_and_execute_tool(
&mut self,
id: &str,
name: &str,
input: &serde_json::Value,
config: &AgentRunConfig,
chunk_tx: &tokio::sync::mpsc::Sender<StreamPiece>,
) -> Content {
let approved = if let Some(observer) = config.observer.as_deref() {
observer.authorize(name)
} else {
let needs_approval = self
.tools
.iter()
.find(|t| t.name() == name)
.is_none_or(|t| t.requires_approval());
if needs_approval {
if let Some(ref tx) = self.approval_tx {
let (oneshot_tx, oneshot_rx) = oneshot::channel();
let _ = tx
.send(ApprovalRequest {
tool_name: name.to_string(),
input: input.clone(),
tx: oneshot_tx,
})
.await;
match timeout(Duration::from_secs(APPROVAL_TIMEOUT_SECS), oneshot_rx).await {
Ok(Ok(res)) => res,
_ => false,
}
} else {
true
}
} else {
if let Some(tool) = self.tools.iter().find(|t| t.name() == name) {
if let Some(notice) = tool.approval_notice() {
let _ = chunk_tx.send(StreamPiece::Notice(notice)).await;
}
}
true
}
};
if !approved {
let denial_msg = if config.observer.is_some() {
format!("Tool '{name}' denied: not authorized in the current authorization tier")
} else {
"Execution denied or timed out.".to_string()
};
if let Some(observer) = config.observer.as_deref() {
observer.on_tool_call(id, name, input, &denial_msg, false, 0);
}
return Content::ToolResult {
tool_use_id: id.to_string(),
content: denial_msg,
is_error: true,
};
}
let started = Instant::now();
let (result_content, is_error) =
if let Some(tool) = self.tools.iter().find(|t| t.name() == name) {
match tool.execute(input.clone(), &config.cancel).await {
Ok(val) => (val.to_string(), false),
Err(e) => (e.to_string(), true),
}
} else {
(format!("Tool '{}' not found", name), true)
};
let elapsed_ms = u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX);
if let Some(observer) = config.observer.as_deref() {
observer.on_tool_call(id, name, input, &result_content, !is_error, elapsed_ms);
}
Content::ToolResult {
tool_use_id: id.to_string(),
content: result_content,
is_error,
}
}
async fn run_tool_loop(
&mut self,
mut working: Vec<Message>,
chunk_tx: &tokio::sync::mpsc::Sender<StreamPiece>,
config: &AgentRunConfig,
prompt: &str,
) -> Result<(String, String)> {
let mut tool_call_count = 0;
let mut last_normalized_tool: Option<(String, String)> = None;
let mut repeat_count = 0;
let mut forced_consult_done = false;
if config.force_consult {
let input = serde_json::json!({ "query": prompt });
tool_call_count += 1;
if tool_call_count > config.max_tool_calls {
return Err(anyhow::anyhow!(MAX_TOOL_CALLS_ERROR));
}
let result_content = self
.authorize_and_execute_tool(
FORCED_CONSULT_TOOL_USE_ID,
CONSULT_TOOL_NAME,
&input,
config,
chunk_tx,
)
.await;
let tool_use_msg = Message {
role: Role::Assistant,
content: vec![Content::ToolUse {
id: FORCED_CONSULT_TOOL_USE_ID.to_string(),
name: CONSULT_TOOL_NAME.to_string(),
input: input.clone(),
}],
};
self.history.push(tool_use_msg.clone());
working.push(tool_use_msg.clone());
if let (Some(memory), Some(sid)) = (&self.memory, &self.session_id) {
memory.add_message(sid, &tool_use_msg).await?;
}
let tool_res_msg = Message {
role: Role::User,
content: vec![result_content],
};
self.history.push(tool_res_msg.clone());
working.push(tool_res_msg.clone());
if let (Some(memory), Some(sid)) = (&self.memory, &self.session_id) {
memory.add_message(sid, &tool_res_msg).await?;
}
forced_consult_done = true;
}
loop {
let mut stream = self
.provider
.stream_messages(&working, &self.tools, config.system.as_deref())
.await?;
let mut full_text = String::new();
let mut turn_text_blocks = 0usize;
let mut last_message: Option<Message> = None;
while let Some(chunk_result) = stream.next().await {
match chunk_result? {
ResponseChunk::TextDelta(delta) => {
let sanitized = Self::sanitize_text(&delta);
turn_text_blocks += 1;
full_text.push_str(&sanitized);
if chunk_tx
.send(StreamPiece::Content(sanitized))
.await
.is_err()
{
return Err(anyhow::anyhow!("TUI connection closed during streaming"));
}
}
ResponseChunk::ReasoningDelta(delta) => {
let sanitized = Self::sanitize_text(&delta);
if chunk_tx
.send(StreamPiece::Reasoning(sanitized))
.await
.is_err()
{
return Err(anyhow::anyhow!("TUI connection closed during streaming"));
}
}
ResponseChunk::MessageDone(msg) => {
last_message = Some(msg);
}
ResponseChunk::Usage {
input_tokens,
output_tokens,
} => {
if let Some(observer) = config.observer.as_deref() {
observer.on_usage(input_tokens, output_tokens);
}
}
ResponseChunk::ToolUseInputDelta { .. } => {}
}
}
let response =
last_message.ok_or_else(|| anyhow::anyhow!("Stream ended without MessageDone"))?;
self.history.push(response.clone());
working.push(response.clone());
if let (Some(memory), Some(sid)) = (&self.memory, &self.session_id) {
memory.add_message(sid, &response).await?;
}
let mut tool_results = Vec::new();
let mut requested_tool = false;
for content in &response.content {
if let Content::ToolUse { id, name, input } = content {
requested_tool = true;
tool_call_count += 1;
if tool_call_count > config.max_tool_calls {
return Err(anyhow::anyhow!(MAX_TOOL_CALLS_ERROR));
}
let normalized_input = Self::normalize_input(input, 0)?;
if let Some((ref last_name, ref last_norm_input)) = last_normalized_tool {
if last_name == name && last_norm_input == &normalized_input {
repeat_count += 1;
if repeat_count >= 3 && !config.disable_repetitive_guard {
return Err(anyhow::anyhow!("Repetitive tool call detected"));
}
} else {
repeat_count = 0;
}
}
last_normalized_tool = Some((name.clone(), normalized_input));
let result_content =
if config.force_consult && forced_consult_done && name == CONSULT_TOOL_NAME
{
Content::ToolResult {
tool_use_id: id.clone(),
content: CONSULT_ALREADY_FORCED_MESSAGE.to_string(),
is_error: true,
}
} else {
self.authorize_and_execute_tool(id, name, input, config, chunk_tx)
.await
};
tool_results.push(result_content);
}
}
if requested_tool {
let tool_res_msg = Message {
role: Role::User,
content: tool_results,
};
self.history.push(tool_res_msg.clone());
working.push(tool_res_msg.clone());
if let (Some(memory), Some(sid)) = (&self.memory, &self.session_id) {
memory.add_message(sid, &tool_res_msg).await?;
}
} else {
if let Some(observer) = config.observer.as_deref() {
observer.on_final_turn(turn_text_blocks);
}
let final_text = response
.content
.iter()
.rev()
.find_map(|c| {
if let Content::Text { text } = c {
Some(text.clone())
} else {
None
}
})
.unwrap_or_default();
return Ok((full_text, final_text));
}
}
}
}
fn summarize_assembly_error(e: &MemoryError) -> String {
let raw = e.to_string();
const NOTICE_MAX: usize = 80;
if raw.len() <= NOTICE_MAX {
format!(
"memory: context assembly failed — using full history ({})",
raw
)
} else {
let truncated: String = raw.chars().take(NOTICE_MAX).collect();
format!(
"memory: context assembly failed — using full history ({}…)",
truncated
)
}
}
#[allow(clippy::too_many_arguments)]
async fn write_turn_to_memory(
store: &SqliteVectorStore,
embedder: &dyn EmbeddingProvider,
clock: &dyn Clock,
cfg: &MemoryConfig,
scope: &str,
session_id: &str,
text: &str,
role: Role,
notice_tx: Option<tokio::sync::mpsc::Sender<StreamPiece>>,
) {
if text.trim().is_empty() {
return;
}
let now = clock.now();
let salience = assign_salience(MemoryKind::Episodic, text, role.clone(), cfg);
let doc_prefix = embedder.document_prefix();
let prefixed = if doc_prefix.is_empty() {
text.to_string()
} else {
format!("{doc_prefix}{text}")
};
let (embedding, model_id, dim) = match embedder.embed(&[prefixed]).await {
Ok(v) => match v.into_iter().next() {
Some(first) => {
let d = first.len();
(first, embedder.model_id().to_string(), d)
}
None => (Vec::new(), String::new(), 0),
},
Err(_) => (Vec::new(), String::new(), 0),
};
let seq = TURN_SEQ.fetch_add(1, Ordering::Relaxed);
let id = format!(
"turn:{:x}",
Sha256::digest(format!("{now}:{role:?}:{seq}:{text}").as_bytes())
);
let m = Memory {
id,
session_id: session_id.to_string(),
kind: MemoryKind::Episodic,
text: text.to_string(),
embedding,
model_id,
dim,
created_at: now,
salience,
access_count: 0,
last_accessed_at: now,
superseded_by: None,
evicted_at: None,
scope: scope.to_string(),
distilled_at: None,
};
if let Err(e) = store.insert(&m).await {
if let Some(tx) = notice_tx {
let _ = tx
.send(StreamPiece::Notice(
"memory: turn not persisted (non-fatal)".to_string(),
))
.await;
} else {
eprintln!("WARN [magi-rs]: memory insert failed (non-fatal): {e}");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tools::ToolResult;
use anyhow::Result;
use async_trait::async_trait;
use futures::stream::{self, BoxStream};
use serde_json::{json, Value};
pub struct MockProvider;
#[async_trait]
impl Provider for MockProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
let chunks = vec![
Ok(ResponseChunk::TextDelta("Summary content.".to_string())),
Ok(ResponseChunk::MessageDone(Message::assistant(
"Summary content.",
))),
];
Ok(Box::pin(stream::iter(chunks)))
}
async fn send_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<Message> {
Ok(Message::assistant("Summary content."))
}
}
#[tokio::test]
async fn test_agent_normalization_depth_limit() {
let mut deep_json = json!({"path": "."});
for i in 0..20 {
deep_json = json!({format!("level_{}", i): deep_json});
}
let result = Agent::normalize_input(&deep_json, 0);
assert!(result.is_err());
}
#[tokio::test]
async fn test_agent_text_sanitization() {
let input = "\x1B[31mDangerous\x1B[0m Text\x07";
assert_eq!(Agent::sanitize_text(input), "Dangerous Text");
}
#[tokio::test]
async fn test_agent_encrypted_persistence_integration() {
let tmp_dir = tempfile::tempdir().unwrap();
let db_path = tmp_dir.path().join("test_persist.db");
let password = "master_password_123".to_string();
let msg_text = "Persist this message";
let _sid = "session_1".to_string();
let sid = {
let mut agent = Agent::new(Arc::new(MockProvider));
let memory = Arc::new(
crate::system::database::EncryptedSqliteMemory::new(
db_path.clone(),
zeroize::Zeroizing::new(password.clone()),
)
.unwrap(),
);
let id = memory.create_session("test_proj").await.unwrap();
agent.set_memory(memory.clone(), id.clone());
let user_msg = Message::user(msg_text);
agent.history.push(user_msg.clone());
memory.add_message(&id, &user_msg).await.unwrap();
id
};
{
let mut agent = Agent::new(Arc::new(MockProvider));
let memory = Arc::new(
crate::system::database::EncryptedSqliteMemory::new(
db_path,
zeroize::Zeroizing::new(password),
)
.unwrap(),
);
agent.set_memory(memory, sid);
agent.load_history().await.unwrap();
assert_eq!(agent.history.len(), 1);
if let Content::Text { text } = &agent.history[0].content[0] {
assert_eq!(text, msg_text);
} else {
panic!("Expected text content");
}
}
}
#[tokio::test]
async fn test_query_streaming_forwards_deltas_to_channel_before_final() {
use tokio::sync::mpsc;
struct TwoDeltaProvider;
#[async_trait]
impl Provider for TwoDeltaProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
let chunks = vec![
Ok(ResponseChunk::ReasoningDelta("thinking".to_string())),
Ok(ResponseChunk::TextDelta("Hello ".to_string())),
Ok(ResponseChunk::TextDelta("world".to_string())),
Ok(ResponseChunk::MessageDone(Message::assistant(
"Hello world",
))),
];
Ok(Box::pin(stream::iter(chunks)))
}
}
let mut agent = Agent::new(Arc::new(TwoDeltaProvider));
let (chunk_tx, mut chunk_rx) = mpsc::channel::<StreamPiece>(8);
let collector = tokio::spawn(async move {
let mut received = Vec::new();
while let Some(piece) = chunk_rx.recv().await {
received.push(piece);
}
received
});
let final_text = agent
.query_streaming("Hi", chunk_tx, AgentRunConfig::default())
.await
.unwrap();
let received = collector.await.unwrap();
assert_eq!(
received,
vec![
StreamPiece::Reasoning("thinking".to_string()),
StreamPiece::Content("Hello ".to_string()),
StreamPiece::Content("world".to_string()),
]
);
assert_eq!(final_text, "Hello world");
}
#[tokio::test]
async fn test_set_provider_swaps_the_active_provider() {
struct FixedProvider(&'static str);
#[async_trait]
impl Provider for FixedProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
let text = self.0.to_string();
Ok(Box::pin(stream::iter(vec![Ok(
ResponseChunk::MessageDone(Message::assistant(&text)),
)])))
}
}
let mut agent = Agent::new(Arc::new(FixedProvider("from-A")));
agent.set_provider(Arc::new(FixedProvider("from-B")));
let (tx, _rx) = tokio::sync::mpsc::channel::<StreamPiece>(8);
let out = agent
.query_streaming("hi", tx, AgentRunConfig::default())
.await
.unwrap();
assert_eq!(out, "from-B", "set_provider must swap the active provider");
}
#[tokio::test]
async fn test_provider_is_static_reflects_provider() {
let static_agent = Agent::new(Arc::new(crate::agent::provider::StaticProvider));
assert!(static_agent.provider_is_static());
let real_agent = Agent::new(Arc::new(MockProvider));
assert!(!real_agent.provider_is_static());
}
struct SystemCapturingProvider {
seen: Arc<std::sync::Mutex<Vec<Option<String>>>>,
}
#[async_trait]
impl Provider for SystemCapturingProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
self.seen.lock().unwrap().push(system.map(str::to_string));
Ok(Box::pin(stream::iter(vec![Ok(
ResponseChunk::MessageDone(Message::assistant("ok")),
)])))
}
}
#[tokio::test]
async fn test_query_streaming_forwards_configured_system_to_provider() {
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let mut agent = Agent::new(Arc::new(SystemCapturingProvider { seen: seen.clone() }));
let (tx, _rx) = tokio::sync::mpsc::channel::<StreamPiece>(8);
let config = AgentRunConfig {
system: Some("You are a headless test assistant.".to_string()),
..AgentRunConfig::default()
};
agent.query_streaming("hi", tx, config).await.unwrap();
assert_eq!(
seen.lock().unwrap().as_slice(),
&[Some("You are a headless test assistant.".to_string())]
);
}
#[tokio::test]
async fn test_query_streaming_default_config_sends_no_system() {
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let mut agent = Agent::new(Arc::new(SystemCapturingProvider { seen: seen.clone() }));
let (tx, _rx) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("hi", tx, AgentRunConfig::default())
.await
.unwrap();
assert_eq!(seen.lock().unwrap().as_slice(), &[None]);
}
struct UsageScriptedProvider {
turns: std::sync::Mutex<std::collections::VecDeque<(u64, u64)>>,
}
#[async_trait]
impl Provider for UsageScriptedProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
let (input_tokens, output_tokens) =
self.turns.lock().unwrap().pop_front().unwrap_or((0, 0));
let is_last = self.turns.lock().unwrap().is_empty();
let mut chunks = vec![Ok(ResponseChunk::Usage {
input_tokens,
output_tokens,
})];
if is_last {
chunks.push(Ok(ResponseChunk::MessageDone(Message::assistant("done"))));
} else {
chunks.push(Ok(ResponseChunk::MessageDone(Message {
role: Role::Assistant,
content: vec![Content::ToolUse {
id: "u1".to_string(),
name: "no-such-tool".to_string(),
input: serde_json::json!({}),
}],
})));
}
Ok(Box::pin(stream::iter(chunks)))
}
}
#[derive(Default)]
struct UsageSpyObserver {
total: std::sync::Mutex<(u64, u64)>,
}
impl RunObserver for UsageSpyObserver {
fn authorize(&self, _tool_name: &str) -> bool {
true
}
fn on_tool_call(
&self,
_id: &str,
_name: &str,
_input: &serde_json::Value,
_result: &str,
_ok: bool,
_ms: u64,
) {
}
fn on_final_turn(&self, _text_block_count: usize) {}
fn on_usage(&self, input_tokens: u64, output_tokens: u64) {
let mut t = self.total.lock().unwrap();
t.0 += input_tokens;
t.1 += output_tokens;
}
}
#[tokio::test]
async fn test_run_tool_loop_accumulates_usage_via_observer_across_turns() {
let provider = Arc::new(UsageScriptedProvider {
turns: std::sync::Mutex::new(std::collections::VecDeque::from([(10, 2), (5, 3)])),
});
let mut agent = Agent::new(provider);
let observer = Arc::new(UsageSpyObserver::default());
let config = AgentRunConfig {
observer: Some(observer.clone() as Arc<dyn RunObserver>),
..AgentRunConfig::default()
};
let (tx, _rx) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent.query_streaming("hi", tx, config).await.unwrap();
assert_eq!(*observer.total.lock().unwrap(), (15, 5));
}
#[tokio::test]
async fn test_on_usage_default_impl_is_a_no_op_for_observers_that_ignore_it() {
struct MinimalObserver;
impl RunObserver for MinimalObserver {
fn authorize(&self, _tool_name: &str) -> bool {
true
}
fn on_tool_call(
&self,
_id: &str,
_name: &str,
_input: &serde_json::Value,
_result: &str,
_ok: bool,
_ms: u64,
) {
}
fn on_final_turn(&self, _text_block_count: usize) {}
}
MinimalObserver.on_usage(1, 1);
}
struct CapturingProvider {
calls: Arc<std::sync::Mutex<Vec<Vec<Message>>>>,
}
impl CapturingProvider {
fn new() -> (Self, Arc<std::sync::Mutex<Vec<Vec<Message>>>>) {
let calls = Arc::new(std::sync::Mutex::new(Vec::<Vec<Message>>::new()));
(
Self {
calls: calls.clone(),
},
calls,
)
}
}
#[async_trait]
impl Provider for CapturingProvider {
async fn stream_messages(
&self,
messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
self.calls.lock().unwrap().push(messages.to_vec());
let chunks = vec![
Ok(ResponseChunk::TextDelta("ok".to_string())),
Ok(ResponseChunk::MessageDone(Message::assistant("ok"))),
];
Ok(Box::pin(stream::iter(chunks)))
}
}
fn bow(text: &str, dim: usize) -> Vec<f32> {
let mut v = vec![0f32; dim];
for w in text.to_lowercase().split_whitespace() {
let h = w
.bytes()
.fold(0usize, |a, b| a.wrapping_mul(31).wrapping_add(b as usize))
% dim;
v[h] += 1.0;
}
let n = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if n > 0.0 {
for x in &mut v {
*x /= n;
}
}
v
}
struct FakeEmbedder {
dim: usize,
model: String,
}
#[async_trait]
impl EmbeddingProvider for FakeEmbedder {
async fn embed(
&self,
texts: &[String],
) -> std::result::Result<Vec<Vec<f32>>, crate::memory::error::EmbeddingError> {
Ok(texts.iter().map(|t| bow(t, self.dim)).collect())
}
fn model_id(&self) -> &str {
&self.model
}
fn dim(&self) -> usize {
self.dim
}
fn query_prefix(&self) -> &str {
""
}
fn document_prefix(&self) -> &str {
""
}
}
struct ErrorEmbedder;
#[async_trait]
impl EmbeddingProvider for ErrorEmbedder {
async fn embed(
&self,
_texts: &[String],
) -> std::result::Result<Vec<Vec<f32>>, crate::memory::error::EmbeddingError> {
Err(crate::memory::error::EmbeddingError::Auth)
}
fn model_id(&self) -> &str {
"error-model"
}
fn dim(&self) -> usize {
16
}
fn query_prefix(&self) -> &str {
""
}
fn document_prefix(&self) -> &str {
""
}
}
#[tokio::test]
async fn test_selective_mode_falls_back_to_history_on_embedder_error() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
use tokio::sync::mpsc;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
..MemoryConfig::default()
};
let embedder = Arc::new(ErrorEmbedder);
let (cap, _calls) = CapturingProvider::new();
let mut agent = Agent::new(Arc::new(cap));
agent.set_memory_subsystem(vstore, embedder, clock, cfg);
let (tx, mut rx) = mpsc::channel::<StreamPiece>(16);
let result = agent
.query_streaming("hello", tx, AgentRunConfig::default())
.await;
assert!(
result.is_ok(),
"D1: selective mode with embedder error must fall back, not propagate Err: {result:?}"
);
let mut got_response = false;
while let Ok(piece) = rx.try_recv() {
if let StreamPiece::Content(text) = piece {
if !text.is_empty() {
got_response = true;
}
}
}
assert!(
got_response,
"D1: a response must be delivered even after embedder error (fallback path)"
);
}
#[tokio::test]
async fn test_selective_mode_forwards_assembler_notices() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
use tokio::sync::mpsc;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 16,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
context_budget_tokens: 20, response_headroom_tokens: 0,
safety_margin_ratio: 0.0,
oversized_turn_policy: "truncate".into(),
..MemoryConfig::default()
};
let (cap, _calls) = CapturingProvider::new();
let mut agent = Agent::new(Arc::new(cap));
agent.set_memory_subsystem(vstore, embedder, clock, cfg);
let long_input = "x".repeat(500);
let (tx, mut rx) = mpsc::channel::<StreamPiece>(32);
agent
.query_streaming(&long_input, tx, AgentRunConfig::default())
.await
.unwrap();
let mut pieces = Vec::new();
while let Ok(p) = rx.try_recv() {
pieces.push(p);
}
let notice_found = pieces.iter().any(|p| matches!(p, StreamPiece::Notice(_)));
assert!(
notice_found,
"D2: a truncation notice must be forwarded to chunk_tx as StreamPiece::Notice when \
the turn exceeds the budget; got pieces: {pieces:?}"
);
}
#[tokio::test]
async fn test_selective_d1_fallback_routes_notice_not_eprintln() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
use tokio::sync::mpsc;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
..MemoryConfig::default()
};
let (cap, _calls) = CapturingProvider::new();
let mut agent = Agent::new(Arc::new(cap));
agent.set_memory_subsystem(vstore, Arc::new(ErrorEmbedder), clock, cfg);
let (tx, mut rx) = mpsc::channel::<StreamPiece>(16);
let result = agent
.query_streaming("hello", tx, AgentRunConfig::default())
.await;
assert!(
result.is_ok(),
"D1-notice: selective mode with embedder error must fall back, not Err: {result:?}"
);
let mut pieces = Vec::new();
while let Ok(p) = rx.try_recv() {
pieces.push(p);
}
let got_notice = pieces.iter().any(|p| matches!(p, StreamPiece::Notice(_)));
assert!(
got_notice,
"D1-notice: a StreamPiece::Notice must be emitted through chunk_tx on \
assembly failure (not eprintln!); got pieces: {pieces:?}"
);
let got_content = pieces
.iter()
.any(|p| matches!(p, StreamPiece::Content(s) if !s.is_empty()));
assert!(
got_content,
"D1-notice: provider response (Content piece) must arrive after fallback; \
got pieces: {pieces:?}"
);
}
#[tokio::test]
async fn test_load_all_mode_is_unchanged() {
let mut agent = Agent::new(Arc::new(MockProvider));
let (tx1, _rx1) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("hello", tx1, AgentRunConfig::default())
.await
.unwrap();
assert_eq!(
agent.history.len(),
2,
"load_all: history must have 2 messages after 1 turn"
);
let (tx2, _rx2) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("world", tx2, AgentRunConfig::default())
.await
.unwrap();
assert_eq!(
agent.history.len(),
4,
"load_all: history must grow to 4 messages after 2 turns"
);
let has_hello = agent.history.iter().any(|m| {
m.role == Role::User
&& m.content
.iter()
.any(|c| matches!(c, Content::Text { text } if text == "hello"))
});
assert!(
has_hello,
"load_all: full history must contain the first user turn"
);
}
#[tokio::test]
async fn test_selective_mode_sends_assembled_context_not_full_history() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 16,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
..MemoryConfig::default()
};
let (cap, calls) = CapturingProvider::new();
let mut agent = Agent::new(Arc::new(cap));
agent.set_memory_subsystem(vstore, embedder, clock, cfg);
let (tx1, _rx1) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("first query", tx1, AgentRunConfig::default())
.await
.unwrap();
let (tx2, _rx2) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("second query", tx2, AgentRunConfig::default())
.await
.unwrap();
let all_calls = calls.lock().unwrap();
assert!(
all_calls.len() >= 2,
"provider must have been called at least twice"
);
let turn2_msgs = &all_calls[1];
let has_assistant = turn2_msgs.iter().any(|m| m.role == Role::Assistant);
assert!(
!has_assistant,
"SC-18/SC-32: selective mode must send assembled context without \
Assistant messages (got {} messages in turn 2)",
turn2_msgs.len()
);
}
#[tokio::test]
async fn test_fact_written_in_one_turn_is_recalled_in_a_later_turn() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::{SqliteVectorStore, VectorStore};
use crate::system::database::EncryptedSqliteMemory;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 32,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
context_budget_tokens: 2000,
response_headroom_tokens: 0,
safety_margin_ratio: 0.0,
top_k: 10,
..MemoryConfig::default()
};
let (cap, calls) = CapturingProvider::new();
let mut agent = Agent::new(Arc::new(cap));
agent.set_memory_subsystem(vstore.clone(), embedder, clock, cfg);
let (tx1, _rx1) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("favorite color is blue", tx1, AgentRunConfig::default())
.await
.unwrap();
let mems = vstore.active("root").await.unwrap();
assert!(
!mems.is_empty(),
"SC-22 (write path): vector store must be non-empty after turn 1 \
(write_turn_to_memory stub → fails in RED)"
);
let has_fact = mems.iter().any(|m| m.text.contains("blue"));
assert!(
has_fact,
"SC-22: the 'blue' fact from turn 1 must be in the vector store"
);
let (tx2, _rx2) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("what is favorite color", tx2, AgentRunConfig::default())
.await
.unwrap();
let all_calls = calls.lock().unwrap();
assert!(
all_calls.len() >= 2,
"provider must have been called at least twice"
);
let turn2_msgs = &all_calls[1];
let has_assistant = turn2_msgs.iter().any(|m| m.role == Role::Assistant);
assert!(
!has_assistant,
"SC-22: turn 2 must use assembled context (no Assistant messages)"
);
}
#[tokio::test]
async fn test_on_session_open_is_noop_when_store_is_empty() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 8,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig::default();
let mut agent = Agent::new(Arc::new(MockProvider));
agent.set_memory_subsystem(vstore, embedder, clock, cfg);
agent.on_session_open().await.unwrap();
}
#[tokio::test]
async fn test_agent_history_resilience_to_key_rotation() {
let tmp_dir = tempfile::tempdir().unwrap();
let db_path = tmp_dir.path().join("resilient_test.db");
let key_a = "api_key_alpha".to_string();
let key_b = "api_key_beta".to_string();
let msg_text = "Secrets are safe";
let sid = {
let memory = crate::system::database::EncryptedSqliteMemory::new(
db_path.clone(),
zeroize::Zeroizing::new(key_a),
)
.unwrap();
let id = memory.create_session("test").await.unwrap();
memory
.add_message(&id, &Message::user(msg_text))
.await
.unwrap();
id
};
{
let result = crate::system::database::EncryptedSqliteMemory::new(
db_path.clone(),
zeroize::Zeroizing::new(key_b),
);
assert!(
result.is_err(),
"opening with a different DB master must fail (wrong KEK cannot unwrap the DEK)"
);
}
{
let key_a_again = "api_key_alpha".to_string();
let memory = crate::system::database::EncryptedSqliteMemory::new(
db_path,
zeroize::Zeroizing::new(key_a_again),
)
.unwrap();
let msgs = memory.get_messages(&sid).await.unwrap();
assert_eq!(
msgs,
vec![Message::user(msg_text)],
"the correct master still recovers the data after a failed wrong-key open"
);
}
}
#[tokio::test]
async fn test_distill_triggers_after_n_turns() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::{SqliteVectorStore, VectorStore};
use crate::system::database::EncryptedSqliteMemory;
use tokio::sync::mpsc;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 8,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
distill_every_n_turns: 2,
distill_enabled: true,
..MemoryConfig::default()
};
let mut agent = Agent::new(Arc::new(MockProvider));
agent.set_memory_subsystem(vstore.clone(), embedder, clock, cfg);
let (tx1, _rx1) = mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("first turn", tx1, AgentRunConfig::default())
.await
.unwrap();
let mems_after_1 = vstore.active("root").await.unwrap();
let distilled_after_1 = mems_after_1
.iter()
.filter(|m| m.distilled_at.is_some())
.count();
assert_eq!(
distilled_after_1, 0,
"no memories should be distilled after turn 1 (trigger fires at 2)"
);
let (tx2, _rx2) = mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("second turn", tx2, AgentRunConfig::default())
.await
.unwrap();
let mems_after_2 = vstore.active("root").await.unwrap();
let distilled_after_2 = mems_after_2
.iter()
.filter(|m| m.distilled_at.is_some())
.count();
assert!(
distilled_after_2 > 0,
"at least one memory must be distilled after turn 2 \
(distill_every_n_turns = 2); got 0 distilled"
);
}
#[tokio::test]
#[ignore = "placeholder: wired in Task 13b"]
async fn test_promoted_preference_placeholder() {}
#[tokio::test]
async fn test_selective_mode_current_turn_not_self_recalled() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 32,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
context_budget_tokens: 4000,
response_headroom_tokens: 0,
safety_margin_ratio: 0.0,
top_k: 5,
..MemoryConfig::default()
};
let (cap, calls) = CapturingProvider::new();
let mut agent = Agent::new(Arc::new(cap));
agent.set_memory_subsystem(vstore, embedder, clock, cfg);
let distinctive = "alpha bravo charlie unique self recall prevention test";
let (tx, _rx) = tokio::sync::mpsc::channel::<StreamPiece>(8);
agent
.query_streaming(distinctive, tx, AgentRunConfig::default())
.await
.unwrap();
let locked = calls.lock().unwrap();
assert!(!locked.is_empty(), "G1: provider must have been called");
let turn1 = &locked[0];
if turn1.len() >= 2 {
let preamble_text: String = turn1[..turn1.len() - 1]
.iter()
.flat_map(|m| m.content.iter())
.filter_map(|c| {
if let Content::Text { text } = c {
Some(text.as_str())
} else {
None
}
})
.collect::<Vec<_>>()
.join(" ");
assert!(
!preamble_text.contains("alpha bravo charlie unique"),
"G1: current turn text must not appear in the recall preamble \
on its own first turn (self-recall); preamble was:\n{preamble_text}"
);
}
}
#[tokio::test]
async fn test_write_turn_produces_distinct_ids_for_same_text_in_same_second() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::{SqliteVectorStore, VectorStore};
use crate::system::database::EncryptedSqliteMemory;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = FakeEmbedder {
dim: 8,
model: "fake".into(),
};
let clock = FixedClock::new(1_000_000); let cfg = MemoryConfig::default();
write_turn_to_memory(
&vstore,
&embedder,
&clock,
&cfg,
"root",
"session1",
"same text same second",
Role::User,
None,
)
.await;
write_turn_to_memory(
&vstore,
&embedder,
&clock,
&cfg,
"root",
"session1",
"same text same second",
Role::User,
None,
)
.await;
let all = vstore.active("root").await.unwrap();
assert_eq!(
all.len(),
2,
"G2: two identical write_turn_to_memory calls in the same FixedClock second \
must produce 2 distinct stored records (not 1 deduped record)"
);
}
struct SingleToolCallProvider {
tool_name: String,
call_count: Arc<std::sync::Mutex<u32>>,
}
impl SingleToolCallProvider {
fn new(tool_name: &str) -> (Self, Arc<std::sync::Mutex<u32>>) {
let count = Arc::new(std::sync::Mutex::new(0u32));
(
Self {
tool_name: tool_name.to_string(),
call_count: count.clone(),
},
count,
)
}
}
#[async_trait]
impl Provider for SingleToolCallProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
let mut n = self.call_count.lock().unwrap();
*n += 1;
let call = *n;
drop(n);
if call == 1 {
let msg = Message {
role: Role::Assistant,
content: vec![Content::ToolUse {
id: "approval-policy-test-1".to_string(),
name: self.tool_name.clone(),
input: serde_json::json!({}),
}],
};
Ok(Box::pin(stream::iter(vec![Ok(
ResponseChunk::MessageDone(msg),
)])))
} else {
Ok(Box::pin(stream::iter(vec![
Ok(ResponseChunk::TextDelta("done".to_string())),
Ok(ResponseChunk::MessageDone(Message::assistant("done"))),
])))
}
}
}
struct TrackingTool {
name_str: String,
executed: Arc<std::sync::Mutex<bool>>,
safe: bool,
notice: Option<String>,
}
impl TrackingTool {
fn new(name: &str, safe: bool) -> (Self, Arc<std::sync::Mutex<bool>>) {
let executed = Arc::new(std::sync::Mutex::new(false));
(
Self {
name_str: name.to_string(),
executed: executed.clone(),
safe,
notice: None,
},
executed,
)
}
fn with_notice(name: &str, notice: &str) -> (Self, Arc<std::sync::Mutex<bool>>) {
let executed = Arc::new(std::sync::Mutex::new(false));
(
Self {
name_str: name.to_string(),
executed: executed.clone(),
safe: true, notice: Some(notice.to_string()),
},
executed,
)
}
}
#[async_trait]
impl Tool for TrackingTool {
fn name(&self) -> &str {
&self.name_str
}
fn description(&self) -> &str {
"tracking tool for approval-policy tests"
}
fn input_schema(&self) -> Value {
json!({"type": "object", "properties": {}})
}
async fn execute(&self, _args: Value, _cancel: &CancellationToken) -> ToolResult<Value> {
*self.executed.lock().unwrap() = true;
Ok(json!({"executed": true}))
}
fn requires_approval(&self) -> bool {
!self.safe
}
fn approval_notice(&self) -> Option<String> {
self.notice.clone()
}
}
#[tokio::test]
async fn test_safe_tool_auto_approved_despite_blanket_denier() {
use tokio::sync::mpsc;
let (approval_tx, mut approval_rx) = mpsc::channel::<ApprovalRequest>(8);
let (tool, executed) = TrackingTool::new("safe_op", true );
let (provider, _count) = SingleToolCallProvider::new("safe_op");
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(tool));
agent.set_approval_channel(approval_tx);
let (chunk_tx, _rx) = mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("do the safe thing", chunk_tx, AgentRunConfig::default())
.await
.unwrap();
assert!(
*executed.lock().unwrap(),
"safe tool (requires_approval=false) must execute \
even when a blanket-denier approval_tx is connected"
);
assert!(
approval_rx.try_recv().is_err(),
"safe tool must NOT emit an ApprovalRequest (no-prompt guarantee)"
);
}
#[tokio::test]
async fn test_dangerous_tool_denied_by_blanket_denier() {
use tokio::sync::mpsc;
let (approval_tx, mut approval_rx) = mpsc::channel::<ApprovalRequest>(8);
tokio::spawn(async move {
while let Some(req) = approval_rx.recv().await {
let _ = req.tx.send(false);
}
});
let (tool, executed) = TrackingTool::new("dangerous_op", false );
let (provider, _count) = SingleToolCallProvider::new("dangerous_op");
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(tool));
agent.set_approval_channel(approval_tx);
let (chunk_tx, _rx) = mpsc::channel::<StreamPiece>(8);
agent
.query_streaming(
"do the dangerous thing",
chunk_tx,
AgentRunConfig::default(),
)
.await
.unwrap();
assert!(
!*executed.lock().unwrap(),
"dangerous tool (requires_approval=true) must NOT execute \
when the approval denier rejects the request"
);
}
#[tokio::test]
async fn test_auto_approved_tool_with_notice_emits_stream_notice() {
use tokio::sync::mpsc;
const NOTICE_TEXT: &str = "auto-launch notice: consensus in progress";
let (tool, _executed) = TrackingTool::with_notice("notice_op", NOTICE_TEXT);
let (provider, _count) = SingleToolCallProvider::new("notice_op");
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(tool));
let (chunk_tx, mut chunk_rx) = mpsc::channel::<StreamPiece>(32);
agent
.query_streaming("do the notice thing", chunk_tx, AgentRunConfig::default())
.await
.unwrap();
let mut pieces = Vec::new();
while let Ok(p) = chunk_rx.try_recv() {
pieces.push(p);
}
let has_notice = pieces.iter().any(|p| {
if let StreamPiece::Notice(msg) = p {
msg.contains(NOTICE_TEXT)
} else {
false
}
});
assert!(
has_notice,
"auto-approved tool with approval_notice must emit StreamPiece::Notice \
with the notice text before execution; pieces: {pieces:?}"
);
}
#[tokio::test]
async fn test_auto_approved_tool_without_notice_emits_no_stream_notice() {
use tokio::sync::mpsc;
let (tool, _executed) = TrackingTool::new("silent_op", true );
let (provider, _count) = SingleToolCallProvider::new("silent_op");
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(tool));
let (chunk_tx, mut chunk_rx) = mpsc::channel::<StreamPiece>(32);
agent
.query_streaming("do the silent thing", chunk_tx, AgentRunConfig::default())
.await
.unwrap();
let mut pieces = Vec::new();
while let Ok(p) = chunk_rx.try_recv() {
pieces.push(p);
}
let has_notice = pieces.iter().any(|p| matches!(p, StreamPiece::Notice(_)));
assert!(
!has_notice,
"auto-approved tool with approval_notice=None must NOT emit StreamPiece::Notice; \
pieces: {pieces:?}"
);
}
#[tokio::test]
async fn test_gated_tool_approval_notice_not_emitted_even_when_approved() {
use tokio::sync::mpsc;
struct DangerousNoticeeTool;
#[async_trait]
impl Tool for DangerousNoticeeTool {
fn name(&self) -> &str {
"dangerous_noticeee"
}
fn description(&self) -> &str {
"dangerous tool that also has a notice field"
}
fn input_schema(&self) -> Value {
json!({"type": "object", "properties": {}})
}
async fn execute(
&self,
_args: Value,
_cancel: &CancellationToken,
) -> ToolResult<Value> {
Ok(json!({"executed": true}))
}
fn requires_approval(&self) -> bool {
true }
fn approval_notice(&self) -> Option<String> {
Some("this should NOT appear — tool is gated".into())
}
}
let (approval_tx, mut approval_rx) = mpsc::channel::<ApprovalRequest>(8);
tokio::spawn(async move {
while let Some(req) = approval_rx.recv().await {
let _ = req.tx.send(true);
}
});
let (provider, _count) = SingleToolCallProvider::new("dangerous_noticeee");
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(DangerousNoticeeTool));
agent.set_approval_channel(approval_tx);
let (chunk_tx, mut chunk_rx) = mpsc::channel::<StreamPiece>(32);
agent
.query_streaming(
"do the dangerous noticeee thing",
chunk_tx,
AgentRunConfig::default(),
)
.await
.unwrap();
let mut pieces = Vec::new();
while let Ok(p) = chunk_rx.try_recv() {
pieces.push(p);
}
let has_notice = pieces.iter().any(|p| matches!(p, StreamPiece::Notice(_)));
assert!(
!has_notice,
"gated tool (requires_approval=true) must NOT emit StreamPiece::Notice \
even if it has an approval_notice — notice is only for auto-approved launches; \
pieces: {pieces:?}"
);
}
struct AlwaysSameToolProvider {
tool_name: String,
}
#[async_trait]
impl Provider for AlwaysSameToolProvider {
async fn stream_messages(
&self,
_messages: &[Message],
_tools: &[Box<dyn Tool>],
_system: Option<&str>,
) -> Result<BoxStream<'static, Result<ResponseChunk>>> {
let msg = Message {
role: Role::Assistant,
content: vec![Content::ToolUse {
id: "repeat-id".to_string(),
name: self.tool_name.clone(),
input: json!({"same": "input"}),
}],
};
Ok(Box::pin(stream::iter(vec![Ok(
ResponseChunk::MessageDone(msg),
)])))
}
}
#[tokio::test]
async fn test_interactive_path_regression_approval_gate_still_gates_dangerous_tool() {
use tokio::sync::mpsc;
let gated: Arc<std::sync::Mutex<Vec<String>>> = Arc::new(std::sync::Mutex::new(Vec::new()));
let (approval_tx, mut approval_rx) = mpsc::channel::<ApprovalRequest>(8);
let gated_seen = gated.clone();
tokio::spawn(async move {
while let Some(req) = approval_rx.recv().await {
gated_seen.lock().unwrap().push(req.tool_name.clone());
let _ = req.tx.send(true);
}
});
let (tool, executed) = TrackingTool::new("gated_op", false );
let (provider, _count) = SingleToolCallProvider::new("gated_op");
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(tool));
agent.set_approval_channel(approval_tx);
let (chunk_tx, _rx) = mpsc::channel::<StreamPiece>(8);
let response = agent
.query_streaming("use the gated tool", chunk_tx, AgentRunConfig::default())
.await
.expect("interactive run must complete once approval is granted");
assert_eq!(
*gated.lock().unwrap(),
vec!["gated_op".to_string()],
"interactive path must route a dangerous tool through the approval gate \
(one ApprovalRequest), NOT auto-approve it"
);
assert!(
*executed.lock().unwrap(),
"the approved dangerous tool must execute on the interactive path"
);
assert_eq!(
response, "done",
"interactive run must return the model's normal final text after the tool call"
);
}
#[tokio::test]
async fn test_interactive_path_regression_repetitive_guard_still_fires() {
use tokio::sync::mpsc;
let (tool, _executed) = TrackingTool::new("loop_op", true );
let provider = AlwaysSameToolProvider {
tool_name: "loop_op".to_string(),
};
let mut agent = Agent::new(Arc::new(provider));
agent.register_tool(Box::new(tool));
let (chunk_tx, _rx) = mpsc::channel::<StreamPiece>(8);
let err = agent
.query_streaming("loop forever", chunk_tx, AgentRunConfig::default())
.await
.expect_err("the repetitive-call guard must abort the run on the interactive path");
let msg = err.to_string();
assert!(
msg.contains("Repetitive tool call detected"),
"interactive path must fire the repetitive-call guard (not be silenced); got: {msg}"
);
assert!(
!msg.contains(MAX_TOOL_CALLS_ERROR),
"the guard must terminate the run before the max-tool-calls cap; got: {msg}"
);
}
#[tokio::test]
async fn test_promoted_preference_appears_in_assembled_context() {
use crate::memory::clock::FixedClock;
use crate::memory::config::MemoryConfig;
use crate::memory::store::SqliteVectorStore;
use crate::system::database::EncryptedSqliteMemory;
use tokio::sync::mpsc;
struct PrefCapProvider {
calls: Arc<std::sync::Mutex<Vec<Vec<Message>>>>,
}
#[async_trait]
impl Provider for PrefCapProvider {
async fn stream_messages(
&self,
messages: &[Message],
_tools: &[Box<dyn crate::tools::Tool>],
_system: Option<&str>,
) -> Result<futures::stream::BoxStream<'static, Result<ResponseChunk>>> {
self.calls.lock().unwrap().push(messages.to_vec());
Ok(Box::pin(futures::stream::iter(vec![
Ok(ResponseChunk::TextDelta("- use rust".to_string())),
Ok(ResponseChunk::MessageDone(Message::assistant("- use rust"))),
])))
}
}
let calls = Arc::new(std::sync::Mutex::new(Vec::<Vec<Message>>::new()));
let provider = Arc::new(PrefCapProvider {
calls: calls.clone(),
});
let tmp = tempfile::NamedTempFile::new().unwrap();
let mem = EncryptedSqliteMemory::new(
tmp.path().to_path_buf(),
zeroize::Zeroizing::new("pw".to_string()),
)
.unwrap();
let vstore =
Arc::new(SqliteVectorStore::new(mem.shared_conn(), mem.data_key().unwrap()).unwrap());
let embedder = Arc::new(FakeEmbedder {
dim: 8,
model: "fake".into(),
});
let clock = Arc::new(FixedClock::new(1_000_000));
let cfg = MemoryConfig {
mode: "selective".into(),
distill_every_n_turns: 1, distill_enabled: true,
context_budget_tokens: 4000,
response_headroom_tokens: 0,
safety_margin_ratio: 0.0,
top_k: 0, ..MemoryConfig::default()
};
let mut agent = Agent::new(provider);
agent.set_memory_subsystem(vstore.clone(), embedder, clock, cfg);
let (tx1, _rx1) = mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("hello", tx1, AgentRunConfig::default())
.await
.unwrap();
{
let calls_after_1 = calls.lock().unwrap().len();
let mems = vstore.active("root").await.unwrap();
let pref_count = mems
.iter()
.filter(|m| m.kind == crate::memory::MemoryKind::Preference)
.count();
assert!(
calls_after_1 >= 2,
"distill must fire a provider call after turn 1; got {calls_after_1} calls"
);
assert!(
pref_count > 0,
"promote_to_profile must have inserted a preference after turn 1; \
got {pref_count} preferences (total mems: {})",
mems.len()
);
}
let n_before_t2 = calls.lock().unwrap().len();
let (tx2, _rx2) = mpsc::channel::<StreamPiece>(8);
agent
.query_streaming("what are my prefs", tx2, AgentRunConfig::default())
.await
.unwrap();
let locked = calls.lock().unwrap();
let t2_main = locked
.get(n_before_t2)
.expect("turn 2's main stream_messages call must exist");
let all_text: String = t2_main
.iter()
.flat_map(|m| m.content.iter())
.filter_map(|c| {
if let Content::Text { text } = c {
Some(text.as_str())
} else {
None
}
})
.collect::<Vec<_>>()
.join("\n");
assert!(
all_text.contains("use rust"),
"turn 2's assembled context must contain the promoted preference 'use rust'; \
got:\n{all_text}"
);
}
}