use std::{
collections::VecDeque,
path::Path,
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
};
use async_trait::async_trait;
use basis::{
AllowAll, Bound, CollectingSink, Event, FnSink, OutputSpec, RunConfig, RunError, RunOutcome,
TurnOptions, run::prepare_with_session,
};
use mentra::{
BuiltinProvider, ContentBlock, ModelInfo, Role, Runtime, RuntimePolicy, Session, TokenUsage,
ToolChoice,
provider::{
Provider, ProviderDescriptor, ProviderError, ProviderEventStream, Request, Response,
provider_event_stream_from_response,
},
runtime::VolatileRuntimeStore,
};
use serde::Deserialize;
use serde_json::{Value, json};
#[derive(Debug, Deserialize, PartialEq, Eq)]
struct Review {
verdict: String,
findings: Vec<String>,
}
const SUBMIT_REVIEW: &str = "call this once you have read every changed file";
const WORKSPACE_FILE: &str = "AGENTS.md";
struct ForcedToolProvider {
model: ModelInfo,
payload: Option<Value>,
usage: Option<TokenUsage>,
calls: Arc<AtomicUsize>,
}
impl ForcedToolProvider {
fn answering(payload: Value) -> Self {
Self {
model: ModelInfo::new("typed-model", BuiltinProvider::Anthropic),
payload: Some(payload),
usage: None,
calls: Arc::new(AtomicUsize::new(0)),
}
}
fn ignoring_the_forced_tool() -> Self {
Self {
payload: None,
..Self::answering(json!({}))
}
}
fn reporting_usage(self, input: u64, output: u64) -> Self {
Self {
usage: Some(TokenUsage {
input_tokens: Some(input),
output_tokens: Some(output),
cache_read_input_tokens: Some(1),
cache_creation_input_tokens: Some(2),
..TokenUsage::default()
}),
..self
}
}
}
#[async_trait]
impl Provider for ForcedToolProvider {
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 call = self.calls.fetch_add(1, Ordering::SeqCst);
let forced = match request.tool_choice.clone() {
Some(ToolChoice::Tool { name }) => Some(name),
_ => None,
};
let (content, stop_reason) = match (forced, self.payload.clone()) {
(Some(name), Some(payload)) => (
vec![ContentBlock::ToolUse {
id: format!("terminal-{call}"),
name,
input: payload,
}],
Some("tool_use".to_string()),
),
_ => (vec![ContentBlock::text("I would rather explain")], None),
};
Ok(provider_event_stream_from_response(Response {
id: format!("typed-{call}"),
model: self.model.id.clone(),
role: Role::Assistant,
content,
stop_reason,
usage: self.usage.clone(),
}))
}
}
#[derive(Clone)]
enum Say {
Read,
Answer(Value),
Prose,
}
#[derive(Clone, Debug)]
struct Offer {
tools: Vec<String>,
terminal: Option<String>,
choice: Option<ToolChoice>,
}
impl Offer {
fn ordinary(&self) -> Vec<&String> {
self.tools
.iter()
.filter(|name| Some(*name) != self.terminal.as_ref())
.collect()
}
}
#[derive(Clone)]
struct ScriptedModel {
model: ModelInfo,
rounds: Arc<Mutex<VecDeque<Say>>>,
offers: Arc<Mutex<Vec<Offer>>>,
usage: Option<TokenUsage>,
}
impl ScriptedModel {
fn new(rounds: Vec<Say>) -> Self {
Self {
model: ModelInfo::new("typed-model", BuiltinProvider::Anthropic),
rounds: Arc::new(Mutex::new(VecDeque::from(rounds))),
offers: Arc::new(Mutex::new(Vec::new())),
usage: None,
}
}
fn spending(self, input: u64, output: u64) -> Self {
Self {
usage: Some(TokenUsage {
input_tokens: Some(input),
output_tokens: Some(output),
..TokenUsage::default()
}),
..self
}
}
fn offers(&self) -> Vec<Offer> {
self.offers.lock().expect("not poisoned").clone()
}
}
#[async_trait]
impl Provider for ScriptedModel {
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 terminal = request
.tools
.iter()
.find(|tool| tool.description.as_deref() == Some(SUBMIT_REVIEW))
.map(|tool| tool.name.clone());
let round = {
let mut offers = self.offers.lock().expect("not poisoned");
offers.push(Offer {
tools: request.tools.iter().map(|tool| tool.name.clone()).collect(),
terminal: terminal.clone(),
choice: request.tool_choice.clone(),
});
offers.len()
};
let say = self
.rounds
.lock()
.expect("not poisoned")
.pop_front()
.unwrap_or_else(|| panic!("the model was asked for an unscripted round {round}"));
let content = match say {
Say::Read => vec![ContentBlock::ToolUse {
id: format!("read-{round}"),
name: "files".to_string(),
input: json!({ "operations": [{ "op": "read", "path": WORKSPACE_FILE }] }),
}],
Say::Answer(payload) => vec![ContentBlock::ToolUse {
id: format!("answer-{round}"),
name: terminal.expect("a typed turn's request carries the terminal tool"),
input: payload,
}],
Say::Prose => vec![ContentBlock::text("I read it, and it looks fine to me")],
};
let calls_a_tool = content
.iter()
.any(|block| matches!(block, ContentBlock::ToolUse { .. }));
Ok(provider_event_stream_from_response(Response {
id: format!("scripted-{round}"),
model: self.model.id.clone(),
role: Role::Assistant,
content,
stop_reason: calls_a_tool.then(|| "tool_use".to_string()),
usage: self.usage.clone(),
}))
}
}
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, "review this diff").with_context(basis::ContextConfig {
file_name: "AGENTS.md".to_string(),
global_dir: None,
walk_parents: false,
})
}
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 prepared(
dir: &tempfile::TempDir,
provider: ForcedToolProvider,
) -> (Runtime, basis::PreparedRun) {
let model = provider.model.clone();
prepared_with(dir, provider, model)
}
fn prepared_with<P: Provider + 'static>(
dir: &tempfile::TempDir,
provider: P,
model: ModelInfo,
) -> (Runtime, basis::PreparedRun) {
let runtime = Runtime::builder()
.with_provider_instance(provider)
.with_store(VolatileRuntimeStore::new())
.with_policy(RuntimePolicy::workspace_bounded(dir.path()))
.build()
.expect("runtime builds");
let run = prepare_with_session(
session(&runtime, dir.path(), model),
&config(dir.path()),
"anthropic",
"typed-model",
)
.expect("prepared");
(runtime, run)
}
fn review_spec() -> OutputSpec {
OutputSpec::new(
"submit_review",
SUBMIT_REVIEW,
json!({
"type": "object",
"properties": {
"verdict": { "type": "string", "description": "ship or hold" },
"findings": {
"type": "array",
"items": { "type": "string" },
"description": "one line per problem worth fixing"
}
},
"required": ["verdict", "findings"]
}),
)
}
#[tokio::test]
async fn a_typed_turn_hands_back_the_value_the_model_committed() {
let dir = workspace();
let (_runtime, mut run) = prepared(
&dir,
ForcedToolProvider::answering(json!({
"verdict": "hold",
"findings": ["the retry loop never gives up"]
})),
);
let output = run
.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect("the run produces a value");
assert_eq!(
output.value,
Review {
verdict: "hold".to_string(),
findings: vec!["the retry loop never gives up".to_string()],
}
);
assert!(output.report.succeeded());
}
#[tokio::test]
async fn a_typed_turn_streams_the_same_bookends_as_any_other() {
let dir = workspace();
let (_runtime, mut run) = prepared(
&dir,
ForcedToolProvider::answering(json!({ "verdict": "ship", "findings": [] })),
);
let output = run
.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect("the run produces a value");
let events = output.report.sink.into_events();
assert!(matches!(events.first(), Some(Event::RunStarted { .. })));
assert!(matches!(
events.last(),
Some(Event::RunFinished {
outcome: RunOutcome::Ok,
..
})
));
assert_eq!(
output.report.final_message, None,
"a typed turn's answer is the value, not prose"
);
assert!(
events.iter().any(|event| matches!(
event,
Event::ToolQueued { input, .. } if input["verdict"] == "ship"
)),
"the payload is on the stream as the terminal call's input"
);
}
#[tokio::test]
async fn an_answer_in_the_wrong_shape_is_told_apart_from_a_failed_run() {
let dir = workspace();
let (_runtime, mut run) = prepared(
&dir,
ForcedToolProvider::answering(json!({ "verdict": "hold", "findings": "lots" })),
);
let error = run
.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect_err("the value does not fit the type");
assert!(
matches!(error, RunError::OutputMismatch(_)),
"expected a mismatch, got {error:?}"
);
}
#[tokio::test]
async fn a_run_that_never_calls_the_terminal_tool_produces_no_value() {
let dir = workspace();
let (_runtime, mut run) = prepared(&dir, ForcedToolProvider::ignoring_the_forced_tool());
let error = run
.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect_err("prose is not an answer to a typed ask");
assert!(
matches!(error, RunError::Runtime(_)),
"expected a runtime failure, got {error:?}"
);
}
#[tokio::test]
async fn a_typed_turn_reports_what_it_spent() {
let dir = workspace();
let (_runtime, mut run) = prepared(
&dir,
ForcedToolProvider::answering(json!({ "verdict": "ship", "findings": [] }))
.reporting_usage(120, 34),
);
let output = run
.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect("the run produces a value");
assert_eq!(output.report.usage.input_tokens, 120);
assert_eq!(output.report.usage.output_tokens, 34);
assert_eq!(output.report.usage.total_tokens(), 154);
}
#[tokio::test]
async fn usage_is_summed_across_every_round_of_a_turn() {
let dir = workspace();
let calls = Arc::new(Mutex::new(0_usize));
struct TwoRounds {
inner: ForcedToolProvider,
rounds: Arc<Mutex<usize>>,
}
#[async_trait]
impl Provider for TwoRounds {
fn descriptor(&self) -> ProviderDescriptor {
self.inner.descriptor()
}
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
self.inner.list_models().await
}
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
let round = {
let mut rounds = self.rounds.lock().expect("not poisoned");
*rounds += 1;
*rounds
};
if round == 1 {
return Ok(provider_event_stream_from_response(Response {
id: "round-1".to_string(),
model: self.inner.model.id.clone(),
role: Role::Assistant,
content: vec![ContentBlock::ToolUse {
id: "call-0".to_string(),
name: "files".to_string(),
input: json!({ "operations": [{ "op": "list", "path": "." }] }),
}],
stop_reason: Some("tool_use".to_string()),
usage: self.inner.usage.clone(),
}));
}
self.inner.stream(request).await
}
}
let inner = ForcedToolProvider::answering(json!({ "verdict": "ship", "findings": [] }))
.reporting_usage(120, 34);
let model = inner.model.clone();
let runtime = Runtime::builder()
.with_provider_instance(TwoRounds {
inner,
rounds: Arc::clone(&calls),
})
.with_store(VolatileRuntimeStore::new())
.with_policy(RuntimePolicy::workspace_bounded(dir.path()))
.build()
.expect("runtime builds");
let mut run = prepare_with_session(
session(&runtime, dir.path(), model),
&config(dir.path()),
"anthropic",
"typed-model",
)
.expect("prepared");
let output = run
.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect("the run produces a value");
assert_eq!(*calls.lock().expect("not poisoned"), 2, "two rounds ran");
assert_eq!(output.report.usage.input_tokens, 240);
assert_eq!(output.report.usage.output_tokens, 68);
assert_eq!(output.report.usage.cache_read_tokens, 2);
assert_eq!(output.report.usage.cache_creation_tokens, 4);
}
#[tokio::test]
async fn a_plain_turn_reports_what_it_spent_too() {
let dir = workspace();
let (_runtime, mut run) = prepared(
&dir,
ForcedToolProvider::ignoring_the_forced_tool().reporting_usage(90, 10),
);
let report = run
.execute(CollectingSink::new())
.await
.expect("run completes");
assert!(report.succeeded());
assert_eq!(report.usage.total_tokens(), 100);
}
#[tokio::test]
async fn a_working_typed_turn_reads_a_file_and_answers_in_the_same_call() {
let dir = workspace();
let provider = ScriptedModel::new(vec![
Say::Read,
Say::Answer(json!({ "verdict": "hold", "findings": ["the house rules are unenforced"] })),
]);
let handle = provider.clone();
let model = provider.model.clone();
let (_runtime, mut run) = prepared_with(&dir, provider, model);
let output = run
.output::<Review, _, _>(
"read AGENTS.md, then review this diff",
review_spec().with_tools(),
CollectingSink::new(),
AllowAll,
)
.await
.expect("a working turn answers");
assert_eq!(
output.value,
Review {
verdict: "hold".to_string(),
findings: vec!["the house rules are unenforced".to_string()],
}
);
let offers = handle.offers();
assert_eq!(offers.len(), 2, "the turn worked a round, then answered");
for (round, offer) in offers.iter().enumerate() {
assert!(
offer.terminal.is_some(),
"round {round} can still end the turn: {:?}",
offer.tools
);
assert!(
offer.ordinary().iter().any(|name| *name == "files"),
"round {round} keeps the ordinary toolset: {:?}",
offer.tools
);
assert!(
!matches!(offer.choice, Some(ToolChoice::Tool { .. })),
"round {round} forces nothing — a forced choice would preclude \
either the working rounds or the call that ends them, got {:?}",
offer.choice
);
}
let events = output.report.sink.into_events();
assert!(
events.iter().any(|event| matches!(
event,
Event::ToolCompleted { tool_name, is_error: false, .. } if tool_name == "files"
)),
"the file was read on the turn that answered: {events:#?}"
);
}
#[tokio::test]
async fn a_shaping_turn_is_still_handed_one_tool_and_told_to_call_it() {
let dir = workspace();
let provider = ScriptedModel::new(vec![Say::Answer(
json!({ "verdict": "ship", "findings": [] }),
)]);
let handle = provider.clone();
let model = provider.model.clone();
let (_runtime, mut run) = prepared_with(&dir, provider, model);
run.output::<Review, _, _>(
"review this diff",
review_spec(),
CollectingSink::new(),
AllowAll,
)
.await
.expect("a shaping turn answers");
let offers = handle.offers();
assert_eq!(offers.len(), 1, "one round decides a shape");
assert!(
offers[0].terminal.is_some() && offers[0].ordinary().is_empty(),
"the terminal tool is the only tool: {:?}",
offers[0].tools
);
assert!(
matches!(offers[0].choice, Some(ToolChoice::Tool { .. })),
"and the model is told to call it, got {:?}",
offers[0].choice
);
}
#[tokio::test]
async fn a_working_turn_that_settles_for_prose_produces_no_value() {
let dir = workspace();
let provider = ScriptedModel::new(vec![Say::Read, Say::Prose]);
let handle = provider.clone();
let model = provider.model.clone();
let (_runtime, mut run) = prepared_with(&dir, provider, model);
let error = run
.output::<Review, _, _>(
"read AGENTS.md, then review this diff",
review_spec().with_tools(),
CollectingSink::new(),
AllowAll,
)
.await
.expect_err("prose is not an answer to a typed ask");
assert!(
matches!(error, RunError::Runtime(_)),
"expected a runtime failure, got {error:?}"
);
let offers = handle.offers();
assert_eq!(offers.len(), 2);
assert!(
offers
.iter()
.all(|offer| offer.ordinary().iter().any(|name| *name == "files")),
"a working turn ran: {offers:?}"
);
}
#[tokio::test]
async fn a_working_turn_out_of_budget_says_so_on_the_stream() {
let dir = workspace();
let provider = ScriptedModel::new(vec![
Say::Read,
Say::Answer(json!({ "verdict": "ship", "findings": [] })),
])
.spending(60, 40);
let handle = provider.clone();
let model = provider.model.clone();
let (_runtime, mut run) = prepared_with(&dir, provider, model);
let events = Arc::new(Mutex::new(Vec::new()));
let recorded = Arc::clone(&events);
let Err(error) = run
.output_with_options::<Review, _, _>(
"read AGENTS.md, then review this diff",
review_spec().with_tools(),
FnSink::new(move |event| {
recorded.lock().expect("not poisoned").push(event);
Ok(())
}),
AllowAll,
TurnOptions::default().with_token_budget(100),
)
.await
else {
panic!("a turn stopped before the terminal call has no value");
};
assert!(
matches!(error, RunError::Runtime(_)),
"expected a runtime failure, got {error:?}"
);
let offers = handle.offers();
assert_eq!(
offers.len(),
1,
"the budget ended the turn before the answering round"
);
assert!(
offers[0].ordinary().iter().any(|name| *name == "files"),
"and it was a working turn that got cut off: {offers:?}"
);
let events = events.lock().expect("not poisoned").clone();
assert!(
matches!(
events.last(),
Some(Event::RunFinished {
stopped_by: Some(Bound::TokenBudget),
..
})
),
"the allowance, not the provider, is what ended it: {:?}",
events.last()
);
}