use std::{
collections::VecDeque,
path::Path,
sync::{Arc, Mutex},
time::Duration,
};
use async_trait::async_trait;
use basis::{
ApprovalAnswer, ApprovalDecision, ApprovalRequest, Approver, CancellationToken, CollectingSink,
Event, RunConfig, RunOutcome, TurnOptions, approval::ApprovalGate, run::prepare_with_session,
};
use mentra::{
BuiltinProvider, ContentBlock, ModelInfo, Role, Runtime, RuntimePolicy, Session,
provider::{
Provider, ProviderDescriptor, ProviderError, ProviderEventStream, Request, Response,
provider_event_stream_from_response,
},
runtime::VolatileRuntimeStore,
};
use serde_json::json;
const PROMPTLY: Duration = Duration::from_secs(10);
struct ScriptedProvider {
model: ModelInfo,
turns: Mutex<VecDeque<Vec<ContentBlock>>>,
}
#[async_trait]
impl Provider for ScriptedProvider {
fn descriptor(&self) -> ProviderDescriptor {
ProviderDescriptor::new(self.model.provider.clone())
}
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
Ok(vec![self.model.clone()])
}
async fn stream(&self, _request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
let content = self
.turns
.lock()
.expect("not poisoned")
.pop_front()
.unwrap_or_else(|| vec![ContentBlock::text("done")]);
Ok(provider_event_stream_from_response(Response {
id: "scripted".to_string(),
model: self.model.id.clone(),
role: Role::Assistant,
content,
stop_reason: None,
usage: None,
}))
}
}
fn scripted_write(workspace: &Path) -> (Runtime, ModelInfo) {
let model = ModelInfo::new("scripted-model", BuiltinProvider::OpenAI);
let provider = ScriptedProvider {
model: model.clone(),
turns: Mutex::new(VecDeque::from(vec![
vec![ContentBlock::ToolUse {
id: "call-0".to_string(),
name: "files".to_string(),
input: json!({
"operations": [
{ "op": "create", "path": "made.txt", "content": "hi" }
]
}),
}],
vec![ContentBlock::text("all done")],
])),
};
let runtime = Runtime::builder()
.with_provider_instance(provider)
.with_store(VolatileRuntimeStore::new())
.with_policy(RuntimePolicy::workspace_bounded(workspace))
.with_tool_authorizer(ApprovalGate::new())
.build()
.expect("runtime builds");
(runtime, model)
}
fn session(runtime: &Runtime, workspace: &Path, model: ModelInfo) -> Session {
runtime
.create_session_with_config(
"test",
model,
mentra::agent::AgentConfig {
workspace: mentra::agent::WorkspaceConfig {
base_dir: workspace.to_path_buf(),
..Default::default()
},
..Default::default()
},
)
.expect("session")
}
fn workspace() -> tempfile::TempDir {
let dir = tempfile::tempdir().expect("tempdir");
std::fs::write(dir.path().join("AGENTS.md"), "house rules").expect("write AGENTS.md");
dir
}
fn config(workspace: &Path) -> RunConfig {
RunConfig::new(workspace, "make a file").with_context(basis::ContextConfig {
file_name: "AGENTS.md".to_string(),
global_dir: None,
walk_parents: false,
})
}
struct CancelsWhenAsked {
token: CancellationToken,
asked: Arc<Mutex<usize>>,
}
#[async_trait]
impl Approver for CancelsWhenAsked {
async fn approve(&mut self, _request: &ApprovalRequest) -> ApprovalAnswer {
*self.asked.lock().expect("not poisoned") += 1;
self.token.cancel();
ApprovalAnswer::new(ApprovalDecision::Allow)
}
}
#[tokio::test]
async fn a_turn_cancelled_mid_flight_reports_a_failed_run() {
let dir = workspace();
let (runtime, model) = scripted_write(dir.path());
let mut prepared = prepare_with_session(
session(&runtime, dir.path(), model),
&config(dir.path()),
"openai",
"scripted-model",
)
.expect("prepared");
let (options, token) = TurnOptions::cancellable();
let asked = Arc::new(Mutex::new(0));
let report = tokio::time::timeout(
PROMPTLY,
prepared.execute_with_approver_and_options(
CollectingSink::new(),
CancelsWhenAsked {
token,
asked: Arc::clone(&asked),
},
options,
),
)
.await
.expect("a cancelled turn must not run to completion")
.expect("cancelling ends the run, it does not break it");
assert_eq!(
*asked.lock().expect("not poisoned"),
1,
"the token was tripped while the turn was blocked on the approver, \
which is what makes this a mid-flight cancellation"
);
assert!(!report.succeeded());
assert_eq!(report.final_message, None);
assert_eq!(report.stopped_by, None);
let events = report.sink.into_events();
assert!(matches!(events.first(), Some(Event::RunStarted { .. })));
assert!(
matches!(
events.last(),
Some(Event::RunFinished {
outcome: RunOutcome::Error { .. },
..
})
),
"a cancelled turn must still close the stream a client is reading"
);
}
#[tokio::test]
async fn a_token_already_tripped_stops_the_turn_before_it_starts() {
let dir = workspace();
let (runtime, model) = scripted_write(dir.path());
let mut prepared = prepare_with_session(
session(&runtime, dir.path(), model),
&config(dir.path()),
"openai",
"scripted-model",
)
.expect("prepared");
let (options, token) = TurnOptions::cancellable();
token.cancel();
let report = tokio::time::timeout(
PROMPTLY,
prepared.execute_with_options(CollectingSink::new(), options),
)
.await
.expect("an already-cancelled turn must return at once")
.expect("cancelling ends the run, it does not break it");
assert!(!report.succeeded());
assert_eq!(report.stopped_by, None);
assert!(
!dir.path().join("made.txt").exists(),
"nothing the scripted turn would have done may happen"
);
}
#[tokio::test]
async fn a_second_turn_is_unaffected_by_the_first_turns_token() {
let dir = workspace();
let (runtime, model) = scripted_write(dir.path());
let mut prepared = prepare_with_session(
session(&runtime, dir.path(), model),
&config(dir.path()),
"openai",
"scripted-model",
)
.expect("prepared");
let (options, token) = TurnOptions::cancellable();
token.cancel();
let cancelled = tokio::time::timeout(
PROMPTLY,
prepared.execute_with_options(CollectingSink::new(), options),
)
.await
.expect("returns at once")
.expect("reports rather than erroring");
assert!(!cancelled.succeeded());
let second = tokio::time::timeout(
PROMPTLY,
prepared.send("try again", CollectingSink::new(), basis::AllowAll),
)
.await
.expect("the second turn must not inherit the first turn's stop button")
.expect("run completes");
assert!(second.succeeded());
}
struct StopsWhenAsked {
token: CancellationToken,
}
#[async_trait]
impl Approver for StopsWhenAsked {
async fn approve(&mut self, _request: &ApprovalRequest) -> ApprovalAnswer {
self.token.cancel();
ApprovalAnswer::new(ApprovalDecision::Allow)
}
}
#[tokio::test]
async fn a_graceful_stop_after_a_tool_round_keeps_its_work_but_reports_failure() {
let dir = workspace();
let (runtime, model) = scripted_write(dir.path());
let mut prepared = prepare_with_session(
session(&runtime, dir.path(), model),
&config(dir.path()),
"openai",
"scripted-model",
)
.expect("prepared");
let (options, token) = TurnOptions::stoppable();
let report = tokio::time::timeout(
PROMPTLY,
prepared.execute_with_approver_and_options(
CollectingSink::new(),
StopsWhenAsked { token },
options,
),
)
.await
.expect("a stopped turn must not run to completion")
.expect("stopping ends the run, it does not break it");
assert!(
dir.path().join("made.txt").exists(),
"a graceful stop keeps the work the run had already committed"
);
assert!(
!report.succeeded(),
"and today reports it as a failure anyway — see this test's docs"
);
assert_eq!(
report.stopped_by, None,
"stopping is not one of the run's own bounds"
);
}