Add operator Mentra memory browser
This commit is contained in:
562
vendor/mentra/tests/agent_runtime.rs
vendored
Normal file
562
vendor/mentra/tests/agent_runtime.rs
vendored
Normal file
@@ -0,0 +1,562 @@
|
||||
use std::{
|
||||
collections::VecDeque,
|
||||
fs,
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicU64, Ordering},
|
||||
},
|
||||
time::{Duration, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use mentra::runtime::{ProviderRetry, RunOptions, SqliteRuntimeStore};
|
||||
use mentra::{
|
||||
AgentConfig, BuiltinProvider, ContentBlock, Message, Role, Runtime,
|
||||
agent::{AgentEvent, AgentStatus, RoundContext, RoundDecision, RoundStrategy},
|
||||
error::RuntimeError,
|
||||
provider::{
|
||||
ContentBlockDelta, ContentBlockStart, ModelInfo, Provider, ProviderDescriptor,
|
||||
ProviderError, ProviderEvent, ProviderEventStream, ProviderId, Request,
|
||||
},
|
||||
};
|
||||
use reqwest::StatusCode;
|
||||
use tokio::sync::{Mutex, broadcast, mpsc};
|
||||
|
||||
static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1);
|
||||
|
||||
enum StreamScript {
|
||||
Buffered(Vec<Result<ProviderEvent, ProviderError>>),
|
||||
Fail(ProviderError),
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ScriptedProvider {
|
||||
kind: ProviderId,
|
||||
models: Vec<ModelInfo>,
|
||||
scripts: std::sync::Arc<Mutex<VecDeque<StreamScript>>>,
|
||||
requests: std::sync::Arc<Mutex<Vec<Request<'static>>>>,
|
||||
}
|
||||
|
||||
impl ScriptedProvider {
|
||||
fn new(
|
||||
kind: impl Into<ProviderId>,
|
||||
models: Vec<ModelInfo>,
|
||||
scripts: Vec<StreamScript>,
|
||||
) -> Self {
|
||||
Self {
|
||||
kind: kind.into(),
|
||||
models,
|
||||
scripts: std::sync::Arc::new(Mutex::new(VecDeque::from(scripts))),
|
||||
requests: std::sync::Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn recorded_requests(&self) -> Vec<Request<'static>> {
|
||||
self.requests.lock().await.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for ScriptedProvider {
|
||||
fn descriptor(&self) -> ProviderDescriptor {
|
||||
ProviderDescriptor::new(self.kind.clone())
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
|
||||
Ok(self.models.clone())
|
||||
}
|
||||
|
||||
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
|
||||
self.requests.lock().await.push(request.into_owned());
|
||||
match self.scripts.lock().await.pop_front() {
|
||||
Some(StreamScript::Buffered(items)) => {
|
||||
let (tx, rx) = mpsc::unbounded_channel();
|
||||
for item in items {
|
||||
tx.send(item)
|
||||
.expect("test stream receiver dropped unexpectedly");
|
||||
}
|
||||
Ok(rx)
|
||||
}
|
||||
Some(StreamScript::Fail(error)) => Err(error),
|
||||
None => panic!("no scripted stream available"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn send_streamed_text_turn_emits_events_and_commits_history() {
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![text_stream(&model.id, "Hello")],
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime
|
||||
.spawn_with_config(
|
||||
"agent",
|
||||
model,
|
||||
AgentConfig {
|
||||
system: Some("system prompt".to_string()),
|
||||
..AgentConfig::default()
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let mut events = agent.subscribe_events();
|
||||
|
||||
let message = agent.send(vec![ContentBlock::text("hi")]).await.unwrap();
|
||||
|
||||
assert_eq!(message, Message::assistant(ContentBlock::text("Hello")));
|
||||
assert_eq!(agent.name(), "agent");
|
||||
assert_eq!(agent.model(), "model");
|
||||
assert_eq!(agent.history().len(), 2);
|
||||
assert_eq!(agent.config().system.as_deref(), Some("system prompt"));
|
||||
assert_eq!(
|
||||
agent.last_message(),
|
||||
Some(&Message::assistant(ContentBlock::text("Hello")))
|
||||
);
|
||||
|
||||
let events = collect_events(&mut events);
|
||||
assert!(events.contains(&AgentEvent::RunStarted));
|
||||
assert!(events.contains(&AgentEvent::TextDelta {
|
||||
delta: "Hello".to_string(),
|
||||
full_text: "Hello".to_string(),
|
||||
}));
|
||||
assert!(matches!(events.last(), Some(AgentEvent::RunFinished)));
|
||||
|
||||
let snapshot = agent.watch_snapshot();
|
||||
assert_eq!(snapshot.borrow().status, AgentStatus::Finished);
|
||||
assert_eq!(snapshot.borrow().history_len, 2);
|
||||
assert!(snapshot.borrow().current_text.is_empty());
|
||||
assert!(snapshot.borrow().pending_tool_uses.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn send_failure_rolls_history_back_and_emits_run_failed() {
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![
|
||||
text_stream(&model.id, "ok"),
|
||||
erroring_stream(
|
||||
vec![ProviderEvent::MessageStarted {
|
||||
id: "msg-2".to_string(),
|
||||
model: model.id.clone(),
|
||||
role: Role::Assistant,
|
||||
}],
|
||||
ProviderError::MalformedStream("boom".to_string()),
|
||||
),
|
||||
],
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).unwrap();
|
||||
agent.send(vec![ContentBlock::text("first")]).await.unwrap();
|
||||
let baseline = agent.history().to_vec();
|
||||
let mut events = agent.subscribe_events();
|
||||
|
||||
let result = agent.send(vec![ContentBlock::text("second")]).await;
|
||||
assert!(result.is_err());
|
||||
assert_eq!(agent.history(), baseline.as_slice());
|
||||
|
||||
let events = collect_events(&mut events);
|
||||
assert!(matches!(events.last(), Some(AgentEvent::RunFailed { .. })));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn send_retries_transient_provider_error_before_streaming() {
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![
|
||||
failed_request(ProviderError::Http {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
body: "offline".to_string(),
|
||||
retry_after: None,
|
||||
}),
|
||||
text_stream(&model.id, "recovered"),
|
||||
],
|
||||
);
|
||||
let provider_handle = provider.clone();
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
|
||||
let message = agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("send should retry");
|
||||
|
||||
assert_eq!(message.text(), "recovered");
|
||||
assert_eq!(provider_handle.recorded_requests().await.len(), 2);
|
||||
assert_eq!(
|
||||
agent.last_message(),
|
||||
Some(&Message::assistant(ContentBlock::text("recovered")))
|
||||
);
|
||||
}
|
||||
|
||||
/// Every delay a run announced before retrying, in the order it waited them.
|
||||
///
|
||||
/// `RetryAttempt` reports the delay the runner is about to take, so this is the
|
||||
/// schedule as it was actually applied rather than as it was configured.
|
||||
fn announced_retry_delays(events: &[AgentEvent]) -> Vec<u64> {
|
||||
events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
AgentEvent::RetryAttempt { next_delay_ms, .. } => Some(*next_delay_ms),
|
||||
_ => None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Fails `count` times with `error`, then answers `text`.
|
||||
fn failing_then_text(count: usize, error: fn() -> ProviderError, text: &str) -> Vec<StreamScript> {
|
||||
let mut scripts: Vec<StreamScript> = (0..count).map(|_| failed_request(error())).collect();
|
||||
scripts.push(text_stream("model", text));
|
||||
scripts
|
||||
}
|
||||
|
||||
fn offline() -> ProviderError {
|
||||
ProviderError::Http {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
body: "offline".to_string(),
|
||||
retry_after: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn a_default_run_waits_the_schedule_mentra_has_always_waited() {
|
||||
// Nobody's timing changes until they ask: 500 ms, doubling.
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
failing_then_text(3, offline, "recovered"),
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
let mut events = agent.subscribe_events();
|
||||
|
||||
let message = agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("the run retries and then succeeds");
|
||||
|
||||
assert_eq!(message.text(), "recovered");
|
||||
assert_eq!(
|
||||
announced_retry_delays(&collect_events(&mut events)),
|
||||
vec![500, 1_000, 2_000]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn a_host_schedule_replaces_the_default_delays() {
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
failing_then_text(3, offline, "recovered"),
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
let mut events = agent.subscribe_events();
|
||||
|
||||
let message = agent
|
||||
.run(
|
||||
vec![ContentBlock::text("hello")],
|
||||
RunOptions::default().with_provider_retry(ProviderRetry {
|
||||
base_delay: Duration::from_secs(2),
|
||||
max_delay: Duration::from_secs(6),
|
||||
..ProviderRetry::default()
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.expect("the run retries and then succeeds");
|
||||
|
||||
assert_eq!(message.text(), "recovered");
|
||||
assert_eq!(
|
||||
announced_retry_delays(&collect_events(&mut events)),
|
||||
vec![2_000, 4_000, 6_000],
|
||||
"the host's base doubles to the host's ceiling"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn a_rate_limit_that_names_its_window_is_waited_out() {
|
||||
// The failure this exists for: a gateway answering 429 with a window far
|
||||
// longer than a blip-shaped backoff would ever reach.
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![
|
||||
failed_request(ProviderError::Http {
|
||||
status: StatusCode::TOO_MANY_REQUESTS,
|
||||
body: "rate limit exceeded".to_string(),
|
||||
retry_after: Some(Duration::from_secs(45)),
|
||||
}),
|
||||
text_stream("model", "recovered"),
|
||||
],
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
let mut events = agent.subscribe_events();
|
||||
|
||||
let message = agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("the run waits out the window and succeeds");
|
||||
|
||||
assert_eq!(message.text(), "recovered");
|
||||
assert_eq!(
|
||||
announced_retry_delays(&collect_events(&mut events)),
|
||||
vec![45_000],
|
||||
"the provider's window, not the schedule's 500 ms"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn a_server_cannot_park_a_run_for_an_hour() {
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![
|
||||
failed_request(ProviderError::Http {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
body: "come back later".to_string(),
|
||||
retry_after: Some(Duration::from_secs(3_600)),
|
||||
}),
|
||||
text_stream("model", "recovered"),
|
||||
],
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
let mut events = agent.subscribe_events();
|
||||
|
||||
agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("the run succeeds after the clamped wait");
|
||||
|
||||
assert_eq!(
|
||||
announced_retry_delays(&collect_events(&mut events)),
|
||||
vec![60_000],
|
||||
"clamped to the default one-minute ceiling"
|
||||
);
|
||||
}
|
||||
|
||||
/// Counters a [`RoundStrategy`] observed at the most recent round boundary,
|
||||
/// captured by [`CountingStrategy`].
|
||||
#[derive(Clone, Copy, Default)]
|
||||
struct RoundCounters {
|
||||
rounds_completed: usize,
|
||||
model_requests: usize,
|
||||
transport_retries: usize,
|
||||
}
|
||||
|
||||
/// A [`RoundStrategy`] that records [`RoundContext`]'s counters at each boundary
|
||||
/// it observes, always proceeding.
|
||||
struct CountingStrategy {
|
||||
last: Mutex<Option<RoundCounters>>,
|
||||
}
|
||||
|
||||
impl CountingStrategy {
|
||||
fn new() -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
last: Mutex::new(None),
|
||||
})
|
||||
}
|
||||
|
||||
async fn last_counters(&self) -> RoundCounters {
|
||||
self.last.lock().await.expect("strategy observed a round")
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl RoundStrategy for CountingStrategy {
|
||||
async fn on_round(&self, ctx: RoundContext<'_>) -> RoundDecision {
|
||||
*self.last.lock().await = Some(RoundCounters {
|
||||
rounds_completed: ctx.rounds_completed(),
|
||||
model_requests: ctx.model_requests(),
|
||||
transport_retries: ctx.transport_retries(),
|
||||
});
|
||||
RoundDecision::proceed()
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retry_and_round_counters_are_reported_distinctly() {
|
||||
// A connection-open retry, then a success: `model_requests` (today's
|
||||
// request-including-retries counter) must stay at 2, `rounds_completed` (the
|
||||
// logical-round counter) must land at 1 — the round that needed a retry still
|
||||
// counts as exactly one completed round — and the new `transport_retries`
|
||||
// counter must isolate the one retry, distinct from both.
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![
|
||||
failed_request(ProviderError::Http {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
body: "offline".to_string(),
|
||||
retry_after: None,
|
||||
}),
|
||||
text_stream(&model.id, "recovered"),
|
||||
],
|
||||
);
|
||||
let provider_handle = provider.clone();
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
|
||||
let strategy = CountingStrategy::new();
|
||||
let message = agent
|
||||
.run(
|
||||
vec![ContentBlock::text("hello")],
|
||||
RunOptions::default().with_round_strategy(strategy.clone()),
|
||||
)
|
||||
.await
|
||||
.expect("run should retry then succeed");
|
||||
|
||||
assert_eq!(message.text(), "recovered");
|
||||
assert_eq!(provider_handle.recorded_requests().await.len(), 2);
|
||||
|
||||
let counters = strategy.last_counters().await;
|
||||
assert_eq!(counters.rounds_completed, 1, "one logical round completed");
|
||||
assert_eq!(
|
||||
counters.model_requests, 2,
|
||||
"model_requests keeps today's semantics: it counts the retry and the success"
|
||||
);
|
||||
assert_eq!(
|
||||
counters.transport_retries, 1,
|
||||
"exactly one transient retry, isolated from rounds_completed and reported distinctly"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resume_replays_last_failed_turn() {
|
||||
let model = model_info("model", BuiltinProvider::Anthropic);
|
||||
let provider = ScriptedProvider::new(
|
||||
BuiltinProvider::Anthropic,
|
||||
vec![model.clone()],
|
||||
vec![
|
||||
erroring_stream(
|
||||
vec![ProviderEvent::MessageStarted {
|
||||
id: "msg-1".to_string(),
|
||||
model: model.id.clone(),
|
||||
role: Role::Assistant,
|
||||
}],
|
||||
ProviderError::MalformedStream("boom".to_string()),
|
||||
),
|
||||
text_stream(&model.id, "done"),
|
||||
],
|
||||
);
|
||||
|
||||
let runtime = test_runtime(provider);
|
||||
let mut agent = runtime.spawn("agent", model).expect("spawn agent");
|
||||
|
||||
let error = agent
|
||||
.send(vec![ContentBlock::text("retry me")])
|
||||
.await
|
||||
.expect_err("first send should fail");
|
||||
assert!(matches!(error, RuntimeError::FailedToStreamResponse(_)));
|
||||
assert!(agent.history().is_empty());
|
||||
|
||||
let resumed = agent
|
||||
.resume()
|
||||
.await
|
||||
.expect("resume should replay user turn");
|
||||
|
||||
assert_eq!(resumed.text(), "done");
|
||||
assert_eq!(agent.history().len(), 2);
|
||||
assert_eq!(
|
||||
agent.history()[0],
|
||||
Message::user(ContentBlock::text("retry me"))
|
||||
);
|
||||
assert_eq!(
|
||||
agent.history()[1],
|
||||
Message::assistant(ContentBlock::text("done"))
|
||||
);
|
||||
|
||||
let error = agent
|
||||
.resume()
|
||||
.await
|
||||
.expect_err("successful run clears resume state");
|
||||
assert!(matches!(error, RuntimeError::NoResumableTurn));
|
||||
}
|
||||
|
||||
fn collect_events(receiver: &mut broadcast::Receiver<AgentEvent>) -> Vec<AgentEvent> {
|
||||
let mut events = Vec::new();
|
||||
while let Ok(event) = receiver.try_recv() {
|
||||
events.push(event);
|
||||
}
|
||||
events
|
||||
}
|
||||
|
||||
fn test_runtime(provider: ScriptedProvider) -> Runtime {
|
||||
Runtime::empty_builder()
|
||||
.with_provider_instance(provider)
|
||||
.with_store(temp_store("agent-runtime"))
|
||||
.build()
|
||||
.expect("build runtime")
|
||||
}
|
||||
|
||||
fn temp_store(label: &str) -> SqliteRuntimeStore {
|
||||
let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed);
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.expect("system time")
|
||||
.as_nanos();
|
||||
let path = std::env::temp_dir().join(format!(
|
||||
"mentra-agent-runtime-{label}-{timestamp}-{unique}.sqlite"
|
||||
));
|
||||
if let Some(parent) = path.parent() {
|
||||
fs::create_dir_all(parent).expect("create temp dir");
|
||||
}
|
||||
SqliteRuntimeStore::new(path)
|
||||
}
|
||||
|
||||
fn model_info(id: &str, provider: impl Into<ProviderId>) -> ModelInfo {
|
||||
ModelInfo::new(id, provider)
|
||||
}
|
||||
|
||||
fn buffered_stream(events: Vec<ProviderEvent>) -> StreamScript {
|
||||
StreamScript::Buffered(events.into_iter().map(Ok).collect())
|
||||
}
|
||||
|
||||
fn erroring_stream(mut events: Vec<ProviderEvent>, error: ProviderError) -> StreamScript {
|
||||
let mut items = events.drain(..).map(Ok).collect::<Vec<_>>();
|
||||
items.push(Err(error));
|
||||
StreamScript::Buffered(items)
|
||||
}
|
||||
|
||||
fn failed_request(error: ProviderError) -> StreamScript {
|
||||
StreamScript::Fail(error)
|
||||
}
|
||||
|
||||
fn text_stream(model: &str, text: &str) -> StreamScript {
|
||||
buffered_stream(vec![
|
||||
ProviderEvent::MessageStarted {
|
||||
id: "msg-text".to_string(),
|
||||
model: model.to_string(),
|
||||
role: Role::Assistant,
|
||||
},
|
||||
ProviderEvent::ContentBlockStarted {
|
||||
index: 0,
|
||||
kind: ContentBlockStart::Text,
|
||||
},
|
||||
ProviderEvent::ContentBlockDelta {
|
||||
index: 0,
|
||||
delta: ContentBlockDelta::Text(text.to_string()),
|
||||
},
|
||||
ProviderEvent::ContentBlockStopped { index: 0 },
|
||||
ProviderEvent::MessageStopped,
|
||||
])
|
||||
}
|
||||
175
vendor/mentra/tests/branching.rs
vendored
Normal file
175
vendor/mentra/tests/branching.rs
vendored
Normal file
@@ -0,0 +1,175 @@
|
||||
//! Public-API tests for the transcript entry tree.
|
||||
//!
|
||||
//! Branching is what lets a conversation return to an earlier point and
|
||||
//! continue differently — "undo that exchange and try another instruction",
|
||||
//! editing a message and re-running, or exploring two approaches from a shared
|
||||
//! prefix — without starting a new session and replaying a prefix by hand.
|
||||
|
||||
use mentra::{
|
||||
AgentTranscript, ContentBlock, Message,
|
||||
transcript::{EntryId, TranscriptItem},
|
||||
};
|
||||
|
||||
fn user(text: &str) -> TranscriptItem {
|
||||
TranscriptItem::user_turn(Message::user(ContentBlock::text(text)))
|
||||
}
|
||||
|
||||
fn assistant(text: &str) -> TranscriptItem {
|
||||
TranscriptItem::assistant_turn(Message::assistant(ContentBlock::text(text)))
|
||||
}
|
||||
|
||||
fn transcript_of(texts: &[&str]) -> AgentTranscript {
|
||||
let mut transcript = AgentTranscript::default();
|
||||
for (index, text) in texts.iter().enumerate() {
|
||||
if index % 2 == 0 {
|
||||
transcript.push(user(text));
|
||||
} else {
|
||||
transcript.push(assistant(text));
|
||||
}
|
||||
}
|
||||
transcript
|
||||
}
|
||||
|
||||
fn texts(transcript: &AgentTranscript) -> Vec<String> {
|
||||
transcript
|
||||
.items()
|
||||
.iter()
|
||||
.map(TranscriptItem::text)
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn appending_hangs_each_entry_from_the_leaf() {
|
||||
let transcript = transcript_of(&["one", "two", "three"]);
|
||||
let items = transcript.items();
|
||||
|
||||
assert_eq!(items[0].parent_id, None, "the first entry is a root");
|
||||
assert_eq!(items[1].parent_id.as_ref(), Some(&items[0].id));
|
||||
assert_eq!(items[2].parent_id.as_ref(), Some(&items[1].id));
|
||||
assert_eq!(transcript.leaf(), Some(&items[2].id));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn branching_rewinds_the_active_path_without_deleting_anything() {
|
||||
let mut transcript = transcript_of(&["ask", "answer", "follow up"]);
|
||||
let rewind_to = transcript.items()[0].id.clone();
|
||||
|
||||
let moved = transcript
|
||||
.branch_from(&rewind_to)
|
||||
.expect("entry is on path");
|
||||
|
||||
assert_eq!(moved, 2, "two entries leave the active path");
|
||||
assert_eq!(texts(&transcript), vec!["ask"]);
|
||||
assert_eq!(transcript.leaf(), Some(&rewind_to));
|
||||
assert_eq!(
|
||||
transcript.archived().len(),
|
||||
2,
|
||||
"abandoned entries stay in the transcript"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_new_turn_after_branching_becomes_a_sibling() {
|
||||
let mut transcript = transcript_of(&["ask", "first answer"]);
|
||||
let fork_point = transcript.items()[0].id.clone();
|
||||
|
||||
transcript.branch_from(&fork_point).expect("on path");
|
||||
transcript.push(assistant("second answer"));
|
||||
|
||||
let children = transcript.children(&fork_point);
|
||||
assert_eq!(children.len(), 2, "the fork point now has two paths");
|
||||
|
||||
let child_texts: Vec<String> = children.iter().map(|item| item.text()).collect();
|
||||
assert!(child_texts.contains(&"first answer".to_string()));
|
||||
assert!(child_texts.contains(&"second answer".to_string()));
|
||||
|
||||
// Only the new path is live.
|
||||
assert_eq!(texts(&transcript), vec!["ask", "second answer"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_abandoned_branch_can_be_returned_to() {
|
||||
let mut transcript = transcript_of(&["ask", "first answer"]);
|
||||
let fork_point = transcript.items()[0].id.clone();
|
||||
let first_answer = transcript.items()[1].id.clone();
|
||||
|
||||
transcript.branch_from(&fork_point).expect("on path");
|
||||
transcript.push(assistant("second answer"));
|
||||
|
||||
// The abandoned entry is still addressable, which is what makes this a
|
||||
// branch rather than a truncation.
|
||||
let recovered = transcript.entry(&first_answer).expect("still present");
|
||||
assert_eq!(recovered.text(), "first answer");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn branching_to_an_unknown_entry_is_refused() {
|
||||
let mut transcript = transcript_of(&["ask"]);
|
||||
|
||||
let error = transcript
|
||||
.branch_from(&EntryId::new())
|
||||
.expect_err("an entry that was never appended is not a branch point");
|
||||
|
||||
assert!(error.to_string().contains("no entry"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn branching_to_the_leaf_changes_nothing() {
|
||||
let mut transcript = transcript_of(&["ask", "answer"]);
|
||||
let leaf = transcript.leaf().cloned().expect("a leaf");
|
||||
|
||||
let moved = transcript.branch_from(&leaf).expect("the leaf is on path");
|
||||
|
||||
assert_eq!(moved, 0);
|
||||
assert_eq!(texts(&transcript), vec!["ask", "answer"]);
|
||||
assert!(transcript.archived().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_tree_survives_a_round_trip() {
|
||||
let mut transcript = transcript_of(&["ask", "first answer"]);
|
||||
let fork_point = transcript.items()[0].id.clone();
|
||||
transcript.branch_from(&fork_point).expect("on path");
|
||||
transcript.push(assistant("second answer"));
|
||||
|
||||
let encoded = serde_json::to_string(&transcript).expect("serializes");
|
||||
let decoded: AgentTranscript = serde_json::from_str(&encoded).expect("deserializes");
|
||||
|
||||
assert_eq!(decoded, transcript);
|
||||
assert_eq!(
|
||||
decoded.children(&fork_point).len(),
|
||||
2,
|
||||
"both paths survive persistence"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_transcript_written_before_entries_had_ids_still_loads_linked() {
|
||||
// Build the shape mentra persisted previously by stripping the tree
|
||||
// fields back out of a current transcript, so the fixture cannot drift
|
||||
// from the real serialization format.
|
||||
let modern = transcript_of(&["ask", "answer"]);
|
||||
let mut encoded: serde_json::Value = serde_json::to_value(&modern).expect("serializes");
|
||||
for item in encoded["items"]
|
||||
.as_array_mut()
|
||||
.expect("items is an array")
|
||||
.iter_mut()
|
||||
{
|
||||
let object = item.as_object_mut().expect("each item is an object");
|
||||
object.remove("id");
|
||||
object.remove("parent_id");
|
||||
}
|
||||
|
||||
let transcript: AgentTranscript =
|
||||
serde_json::from_value(encoded).expect("a pre-tree transcript still deserializes");
|
||||
let items = transcript.items();
|
||||
|
||||
assert_eq!(items.len(), 2);
|
||||
assert_eq!(texts(&transcript), vec!["ask", "answer"]);
|
||||
assert_eq!(items[0].parent_id, None, "the first entry is a root");
|
||||
assert_eq!(
|
||||
items[1].parent_id.as_ref(),
|
||||
Some(&items[0].id),
|
||||
"migration links the chain, so a legacy transcript can be branched"
|
||||
);
|
||||
}
|
||||
54
vendor/mentra/tests/mcp_sse_smoke.rs
vendored
Normal file
54
vendor/mentra/tests/mcp_sse_smoke.rs
vendored
Normal file
@@ -0,0 +1,54 @@
|
||||
//! Optional manual smoke test against a real MCP HTTP+SSE server.
|
||||
//!
|
||||
//! This is ignored by default and never runs in ordinary CI: it needs a live
|
||||
//! endpoint, which no automated run should depend on. The transport itself is
|
||||
//! covered by the deterministic fixture tests in `mentra::mcp::sse`.
|
||||
//!
|
||||
//! Point it at any server speaking the 2024-11-05 HTTP+SSE transport:
|
||||
//!
|
||||
//! ```text
|
||||
//! MENTRA_MCP_SSE_URL=https://mcp.example.com/sse \
|
||||
//! MENTRA_MCP_SSE_TOKEN=<token> \
|
||||
//! cargo test -p mentra --test mcp_sse_smoke -- --ignored --nocapture
|
||||
//! ```
|
||||
//!
|
||||
//! `MENTRA_MCP_SSE_TOKEN` is optional. The test performs only `initialize` and
|
||||
//! `tools/list`; it never calls a tool, so it cannot cause a side effect on the
|
||||
//! server it is pointed at.
|
||||
|
||||
use mentra::{McpSseClient, McpSseServerConfig};
|
||||
|
||||
/// Environment variable naming the SSE endpoint to probe.
|
||||
const URL_VAR: &str = "MENTRA_MCP_SSE_URL";
|
||||
/// Environment variable carrying an optional bearer token.
|
||||
const TOKEN_VAR: &str = "MENTRA_MCP_SSE_TOKEN";
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "requires a live MCP server; set MENTRA_MCP_SSE_URL"]
|
||||
async fn initializes_and_lists_tools_against_a_live_server() {
|
||||
let url = std::env::var(URL_VAR)
|
||||
.unwrap_or_else(|_| panic!("set {URL_VAR} to the server's SSE endpoint"));
|
||||
|
||||
let mut config = McpSseServerConfig::new("smoke", &url);
|
||||
if let Ok(token) = std::env::var(TOKEN_VAR) {
|
||||
config = config.with_bearer_token(token);
|
||||
}
|
||||
|
||||
let client = McpSseClient::connect(&config)
|
||||
.await
|
||||
.expect("the handshake should complete");
|
||||
|
||||
let info = client
|
||||
.server_info()
|
||||
.expect("initialize should report server info");
|
||||
println!("connected to {} {}", info.name, info.version);
|
||||
|
||||
println!("{} tools advertised:", client.tools().len());
|
||||
for tool in client.tools() {
|
||||
println!(" {}", tool.name);
|
||||
}
|
||||
|
||||
// No tool is called: a smoke test must not cause side effects on whatever
|
||||
// server it happens to be pointed at.
|
||||
client.shutdown().await;
|
||||
}
|
||||
859
vendor/mentra/tests/public_api.rs
vendored
Normal file
859
vendor/mentra/tests/public_api.rs
vendored
Normal file
@@ -0,0 +1,859 @@
|
||||
use std::{
|
||||
collections::VecDeque,
|
||||
io::{Read, Write},
|
||||
net::TcpListener,
|
||||
sync::{Arc, Mutex},
|
||||
thread,
|
||||
time::{SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use mentra::{
|
||||
Agent, BuiltinProvider, ContentBlock, FileToolProfile, ModelInfo, ModelSelector, Runtime,
|
||||
error::RuntimeError,
|
||||
provider::{
|
||||
Provider, ProviderDescriptor, ProviderError, ProviderEventStream, ProviderId, Request,
|
||||
Response, Role, provider_event_stream_from_response,
|
||||
},
|
||||
runtime::{
|
||||
CommandOutput, CommandRequest, RuntimeExecutor, RuntimePolicy, VolatileRuntimeStore,
|
||||
},
|
||||
tool::{ParallelToolContext, ToolContext, ToolDefinition, ToolExecutor, ToolResult, ToolSpec},
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
#[derive(Debug)]
|
||||
enum Turn {
|
||||
Text(String),
|
||||
ToolCalls(Vec<ScriptedToolCall>),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct ScriptedToolCall {
|
||||
id: Option<String>,
|
||||
name: String,
|
||||
input: Value,
|
||||
}
|
||||
|
||||
impl ScriptedToolCall {
|
||||
fn new(name: impl Into<String>, input: Value) -> Self {
|
||||
Self {
|
||||
id: None,
|
||||
name: name.into(),
|
||||
input,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ScriptedProvider {
|
||||
kind: ProviderId,
|
||||
models: Vec<ModelInfo>,
|
||||
turns: Arc<Mutex<VecDeque<Turn>>>,
|
||||
requests: Arc<Mutex<Vec<Request<'static>>>>,
|
||||
}
|
||||
|
||||
impl ScriptedProvider {
|
||||
fn new(kind: ProviderId, models: Vec<ModelInfo>) -> Self {
|
||||
Self {
|
||||
kind,
|
||||
models,
|
||||
turns: Arc::new(Mutex::new(VecDeque::new())),
|
||||
requests: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn push_turns(&self, turns: Vec<Turn>) {
|
||||
let mut queue = self.turns.lock().expect("scripted turn queue poisoned");
|
||||
queue.extend(turns);
|
||||
}
|
||||
|
||||
fn recorded_requests(&self) -> Vec<Request<'static>> {
|
||||
self.requests
|
||||
.lock()
|
||||
.expect("scripted request log poisoned")
|
||||
.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for ScriptedProvider {
|
||||
fn descriptor(&self) -> ProviderDescriptor {
|
||||
ProviderDescriptor::new(self.kind.clone())
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
|
||||
Ok(self.models.clone())
|
||||
}
|
||||
|
||||
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
|
||||
self.requests
|
||||
.lock()
|
||||
.expect("scripted request log poisoned")
|
||||
.push(request.into_owned());
|
||||
|
||||
let turn = self
|
||||
.turns
|
||||
.lock()
|
||||
.expect("scripted turn queue poisoned")
|
||||
.pop_front()
|
||||
.unwrap_or_else(|| panic!("no scripted turn remaining for public API test"));
|
||||
|
||||
match turn {
|
||||
Turn::Text(text) => Ok(provider_event_stream_from_response(Response {
|
||||
id: format!("public-response-{}", now_nanos()),
|
||||
model: self.models[0].id.clone(),
|
||||
role: Role::Assistant,
|
||||
content: vec![ContentBlock::text(text)],
|
||||
stop_reason: None,
|
||||
usage: None,
|
||||
})),
|
||||
Turn::ToolCalls(calls) => Ok(provider_event_stream_from_response(Response {
|
||||
id: format!("public-response-{}", now_nanos()),
|
||||
model: self.models[0].id.clone(),
|
||||
role: Role::Assistant,
|
||||
content: calls
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(index, call)| ContentBlock::ToolUse {
|
||||
id: call.id.unwrap_or_else(|| format!("tool-{}", index + 1)),
|
||||
name: call.name,
|
||||
input: call.input,
|
||||
})
|
||||
.collect(),
|
||||
stop_reason: Some("tool_use".to_string()),
|
||||
usage: None,
|
||||
})),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct Harness {
|
||||
runtime: Runtime,
|
||||
provider: ScriptedProvider,
|
||||
model: ModelInfo,
|
||||
}
|
||||
|
||||
impl Harness {
|
||||
fn new(turns: Vec<Turn>) -> Self {
|
||||
let runtime_id = format!("public-api-{}", now_nanos());
|
||||
let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI);
|
||||
let provider = ScriptedProvider::new(model.provider.clone(), vec![model.clone()]);
|
||||
provider.push_turns(turns);
|
||||
|
||||
let runtime = Runtime::builder()
|
||||
.with_runtime_identifier(runtime_id)
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider.clone())
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
Self {
|
||||
runtime,
|
||||
provider,
|
||||
model,
|
||||
}
|
||||
}
|
||||
|
||||
fn spawn(&self, name: &str) -> Agent {
|
||||
self.runtime
|
||||
.spawn(name, self.model.clone())
|
||||
.expect("spawn test agent")
|
||||
}
|
||||
|
||||
async fn recorded_requests(&self) -> Vec<Request<'static>> {
|
||||
self.provider.recorded_requests()
|
||||
}
|
||||
}
|
||||
|
||||
struct EchoTool;
|
||||
|
||||
struct AlphaTool;
|
||||
|
||||
struct EndTurnTool;
|
||||
|
||||
struct SubagentSummaryTool;
|
||||
|
||||
#[async_trait]
|
||||
impl ToolDefinition for EchoTool {
|
||||
fn descriptor(&self) -> ToolSpec {
|
||||
ToolSpec::builder("echo_tool")
|
||||
.description("Echo a canned result")
|
||||
.input_schema(json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}))
|
||||
.build()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolExecutor for EchoTool {
|
||||
async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult {
|
||||
Ok("echoed".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolDefinition for AlphaTool {
|
||||
fn descriptor(&self) -> ToolSpec {
|
||||
ToolSpec::builder("alpha_tool")
|
||||
.description("Return a canned alpha result")
|
||||
.input_schema(json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}))
|
||||
.build()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolExecutor for AlphaTool {
|
||||
async fn execute(&self, _ctx: ParallelToolContext, _input: Value) -> ToolResult {
|
||||
Ok("alpha".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolDefinition for EndTurnTool {
|
||||
fn descriptor(&self) -> ToolSpec {
|
||||
ToolSpec::builder("stop_here")
|
||||
.description("End the current turn without a follow-up assistant message")
|
||||
.input_schema(json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}))
|
||||
.build()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolExecutor for EndTurnTool {
|
||||
async fn execute_mut(&self, mut ctx: ToolContext<'_>, _input: Value) -> ToolResult {
|
||||
ctx.request_idle();
|
||||
Ok("stopping now".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolDefinition for SubagentSummaryTool {
|
||||
fn descriptor(&self) -> ToolSpec {
|
||||
ToolSpec::builder("subagent_summary")
|
||||
.description("Spawn a disposable subagent and return its summary")
|
||||
.input_schema(json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"prompt": { "type": "string" }
|
||||
},
|
||||
"required": ["prompt"]
|
||||
}))
|
||||
.build()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolExecutor for SubagentSummaryTool {
|
||||
async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult {
|
||||
let prompt = input
|
||||
.get("prompt")
|
||||
.and_then(|value| value.as_str())
|
||||
.ok_or_else(|| "prompt is required".to_string())?;
|
||||
let mut child = ctx.spawn_subagent().map_err(|error| error.to_string())?;
|
||||
// The public pattern for a custom tool that spawns work: the child runs
|
||||
// under the parent run's derived bounds, not a fresh unbounded set.
|
||||
let message = child
|
||||
.run(vec![ContentBlock::text(prompt)], ctx.child_run_options())
|
||||
.await
|
||||
.map_err(|error| format!("child failed: {error}"))?;
|
||||
Ok(message.text())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn send_returns_final_message_after_tool_execution() {
|
||||
let harness = Harness::new(vec![
|
||||
Turn::ToolCalls(vec![ScriptedToolCall::new("echo_tool", json!({}))]),
|
||||
Turn::Text("done".to_string()),
|
||||
]);
|
||||
harness.runtime.register_tool(EchoTool);
|
||||
let mut agent = harness.spawn("tool-agent");
|
||||
|
||||
let message = agent
|
||||
.send(vec![ContentBlock::text("run the tool")])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(message.role, Role::Assistant);
|
||||
assert_eq!(message.text(), "done");
|
||||
assert_eq!(harness.recorded_requests().await.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_exposes_registered_tool_descriptors() {
|
||||
let runtime_id = format!("public-api-{}", now_nanos());
|
||||
let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI);
|
||||
let provider = ScriptedProvider::new(model.provider.clone(), vec![model.clone()]);
|
||||
|
||||
let runtime = Runtime::empty_builder()
|
||||
.with_runtime_identifier(runtime_id)
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider)
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
runtime.register_tool(EchoTool);
|
||||
runtime.register_tool(AlphaTool);
|
||||
|
||||
assert_eq!(
|
||||
runtime.tools(),
|
||||
vec![AlphaTool.descriptor(), EchoTool.descriptor()]
|
||||
);
|
||||
assert_eq!(
|
||||
runtime.tool_descriptor("echo_tool"),
|
||||
Some(EchoTool.descriptor())
|
||||
);
|
||||
assert_eq!(
|
||||
runtime.tool_descriptor("alpha_tool"),
|
||||
Some(AlphaTool.descriptor())
|
||||
);
|
||||
assert_eq!(runtime.tool_descriptor("missing_tool"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_builder_publicly_selects_split_file_tools() {
|
||||
let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI);
|
||||
let provider = ScriptedProvider::new(model.provider.clone(), vec![model]);
|
||||
let runtime = Runtime::builder()
|
||||
.with_file_tools(FileToolProfile::Split)
|
||||
.with_provider_instance(provider)
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
let names = runtime
|
||||
.tools()
|
||||
.into_iter()
|
||||
.map(|tool| tool.provider.name)
|
||||
.collect::<std::collections::BTreeSet<_>>();
|
||||
|
||||
for name in ["read", "ls", "grep", "glob", "write", "edit"] {
|
||||
assert!(names.contains(name), "missing split tool {name}");
|
||||
}
|
||||
assert!(!names.contains("files"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn parallel_tool_context_can_spawn_subagents_from_public_api() {
|
||||
let harness = Harness::new(vec![
|
||||
Turn::ToolCalls(vec![ScriptedToolCall::new(
|
||||
"subagent_summary",
|
||||
json!({ "prompt": "summarize the delegated work" }),
|
||||
)]),
|
||||
Turn::Text("child summary".to_string()),
|
||||
Turn::Text("parent complete".to_string()),
|
||||
]);
|
||||
harness.runtime.register_tool(SubagentSummaryTool);
|
||||
let mut agent = harness.spawn("parent-agent");
|
||||
|
||||
let message = agent
|
||||
.send(vec![ContentBlock::text("delegate that")])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(message.role, Role::Assistant);
|
||||
assert_eq!(message.text(), "parent complete");
|
||||
assert_eq!(harness.recorded_requests().await.len(), 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_assistant_response_preserves_committed_tool_results() {
|
||||
let harness = Harness::new(vec![Turn::ToolCalls(vec![ScriptedToolCall::new(
|
||||
"stop_here",
|
||||
json!({}),
|
||||
)])]);
|
||||
harness.runtime.register_tool(EndTurnTool);
|
||||
let mut agent = harness.spawn("idle-agent");
|
||||
|
||||
let error = agent
|
||||
.send(vec![ContentBlock::text("stop after the tool")])
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, RuntimeError::EmptyAssistantResponse));
|
||||
assert_eq!(harness.recorded_requests().await.len(), 1);
|
||||
assert_eq!(agent.history().len(), 3);
|
||||
match &agent.history()[2].content[0] {
|
||||
ContentBlock::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
is_error,
|
||||
} => {
|
||||
assert_eq!(tool_use_id, "tool-1");
|
||||
assert_eq!(content, "stopping now");
|
||||
assert!(!is_error);
|
||||
}
|
||||
other => panic!("expected tool result block, found {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_returns_explicit_id_without_listing_models() {
|
||||
let runtime_id = format!("public-api-{}", now_nanos());
|
||||
let provider = FailingListModelsProvider {
|
||||
kind: BuiltinProvider::Anthropic.into(),
|
||||
};
|
||||
|
||||
let runtime = Runtime::builder()
|
||||
.with_runtime_identifier(runtime_id)
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider)
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let model = runtime
|
||||
.resolve_model(
|
||||
BuiltinProvider::Anthropic,
|
||||
ModelSelector::Id("claude-custom".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect("resolve explicit model");
|
||||
|
||||
assert_eq!(
|
||||
model,
|
||||
ModelInfo::new("claude-custom", BuiltinProvider::Anthropic)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_selects_newest_available_then_breaks_ties_by_id() {
|
||||
let runtime_id = format!("public-api-{}", now_nanos());
|
||||
let provider = ModelListingProvider {
|
||||
kind: BuiltinProvider::OpenAI.into(),
|
||||
models: vec![
|
||||
model_with_created_at("zeta", BuiltinProvider::OpenAI, 1_700_000_100),
|
||||
model_with_created_at("alpha", BuiltinProvider::OpenAI, 1_700_000_100),
|
||||
model_with_created_at("older", BuiltinProvider::OpenAI, 1_700_000_000),
|
||||
],
|
||||
};
|
||||
|
||||
let runtime = Runtime::builder()
|
||||
.with_runtime_identifier(runtime_id)
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider)
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let model = runtime
|
||||
.resolve_model(BuiltinProvider::OpenAI, ModelSelector::NewestAvailable)
|
||||
.await
|
||||
.expect("resolve newest model");
|
||||
|
||||
assert_eq!(
|
||||
model,
|
||||
model_with_created_at("alpha", BuiltinProvider::OpenAI, 1_700_000_100)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_reports_empty_provider_listing() {
|
||||
let runtime_id = format!("public-api-{}", now_nanos());
|
||||
let provider = ModelListingProvider {
|
||||
kind: BuiltinProvider::Gemini.into(),
|
||||
models: Vec::new(),
|
||||
};
|
||||
|
||||
let runtime = Runtime::builder()
|
||||
.with_runtime_identifier(runtime_id)
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider)
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let error = runtime
|
||||
.resolve_model(BuiltinProvider::Gemini, ModelSelector::NewestAvailable)
|
||||
.await
|
||||
.expect_err("empty listing should fail");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
RuntimeError::NoModelsAvailable(provider) if provider == BuiltinProvider::Gemini.into()
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_supports_openrouter_provider() {
|
||||
let runtime_id = format!("public-api-{}", now_nanos());
|
||||
let provider = ModelListingProvider {
|
||||
kind: BuiltinProvider::OpenRouter.into(),
|
||||
models: vec![model_with_created_at(
|
||||
"openai/gpt-4.1-mini",
|
||||
BuiltinProvider::OpenRouter,
|
||||
1_741_049_700,
|
||||
)],
|
||||
};
|
||||
|
||||
let runtime = Runtime::builder()
|
||||
.with_runtime_identifier(runtime_id)
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider)
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let model = runtime
|
||||
.resolve_model(BuiltinProvider::OpenRouter, ModelSelector::NewestAvailable)
|
||||
.await
|
||||
.expect("resolve newest model");
|
||||
|
||||
assert_eq!(
|
||||
model,
|
||||
model_with_created_at(
|
||||
"openai/gpt-4.1-mini",
|
||||
BuiltinProvider::OpenRouter,
|
||||
1_741_049_700,
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_supports_ollama_provider_registration() {
|
||||
let runtime = Runtime::empty_builder()
|
||||
.with_ollama()
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let model = runtime
|
||||
.resolve_model(
|
||||
BuiltinProvider::Ollama,
|
||||
ModelSelector::Id("qwen2.5-coder".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect("resolve explicit model");
|
||||
|
||||
assert_eq!(
|
||||
model,
|
||||
ModelInfo::new("qwen2.5-coder", BuiltinProvider::Ollama)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_supports_lmstudio_provider_registration() {
|
||||
let runtime = Runtime::empty_builder()
|
||||
.with_lmstudio()
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let model = runtime
|
||||
.resolve_model(
|
||||
BuiltinProvider::LmStudio,
|
||||
ModelSelector::Id("local-model".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect("resolve explicit model");
|
||||
|
||||
assert_eq!(
|
||||
model,
|
||||
ModelInfo::new("local-model", BuiltinProvider::LmStudio)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolve_model_reports_missing_provider() {
|
||||
let harness = Harness::new(vec![Turn::Text("unused".to_string())]);
|
||||
|
||||
let error = harness
|
||||
.runtime
|
||||
.resolve_model(
|
||||
BuiltinProvider::Gemini,
|
||||
ModelSelector::Id("gemini-2.5-pro".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect_err("missing provider should fail");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
RuntimeError::ProviderNotFound(Some(provider))
|
||||
if provider == BuiltinProvider::Gemini.into()
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn runtime_accepts_provider_core_openai_compatible_instances() {
|
||||
let provider_id = ProviderId::new("custom-openai-compatible");
|
||||
let (base_url, handle) = spawn_models_server(
|
||||
r#"{"data":[{"id":"compat-model","name":"Compat Model","created":1}]}"#,
|
||||
);
|
||||
|
||||
let mut definition = mentra::provider_core::responses::openai_definition();
|
||||
definition.descriptor.id = provider_id.clone();
|
||||
definition.descriptor.display_name = Some("Custom OpenAI-Compatible".to_string());
|
||||
definition.base_url = Some(base_url);
|
||||
|
||||
let runtime = Runtime::empty_builder()
|
||||
.with_registered_provider(mentra::provider_core::responses::ResponsesProvider::new(
|
||||
definition,
|
||||
mentra::provider_core::StaticCredentialSource::new("test-key"),
|
||||
))
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
|
||||
let model = runtime
|
||||
.resolve_model(provider_id.clone(), ModelSelector::NewestAvailable)
|
||||
.await
|
||||
.expect("resolve model from provider-core instance");
|
||||
|
||||
assert_eq!(model.provider, provider_id);
|
||||
assert_eq!(model.id, "compat-model");
|
||||
|
||||
let captured = handle.join().expect("capture request");
|
||||
let captured_lower = captured.to_ascii_lowercase();
|
||||
assert!(captured.starts_with("GET /v1/models HTTP/1.1\r\n"));
|
||||
assert!(captured_lower.contains("authorization: bearer test-key\r\n"));
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct FailingListModelsProvider {
|
||||
kind: ProviderId,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for FailingListModelsProvider {
|
||||
fn descriptor(&self) -> ProviderDescriptor {
|
||||
ProviderDescriptor::new(self.kind.clone())
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
|
||||
Err(ProviderError::InvalidResponse(
|
||||
"list_models should not be called".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn stream(&self, _request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
|
||||
let (_tx, rx) = tokio::sync::mpsc::unbounded_channel();
|
||||
Ok(rx)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ModelListingProvider {
|
||||
kind: ProviderId,
|
||||
models: Vec<ModelInfo>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for ModelListingProvider {
|
||||
fn descriptor(&self) -> ProviderDescriptor {
|
||||
ProviderDescriptor::new(self.kind.clone())
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
|
||||
Ok(self.models.clone())
|
||||
}
|
||||
|
||||
async fn stream(&self, _request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
|
||||
let (_tx, rx) = tokio::sync::mpsc::unbounded_channel();
|
||||
Ok(rx)
|
||||
}
|
||||
}
|
||||
|
||||
/// The builder was `pub` inside a private `mod builder`, re-exported nowhere:
|
||||
/// `Runtime::builder()` worked on inference, but a downstream helper taking or
|
||||
/// returning a half-built runtime could not write its signature at all. This
|
||||
/// test is that helper, written from outside the crate — it pins the re-export
|
||||
/// at both public paths, because compiling is the whole claim.
|
||||
#[test]
|
||||
fn a_runtime_builder_is_a_type_downstream_code_can_name() {
|
||||
fn with_volatile_store(builder: mentra::RuntimeBuilder) -> mentra::runtime::RuntimeBuilder {
|
||||
builder.with_store(VolatileRuntimeStore::new())
|
||||
}
|
||||
|
||||
let _ = with_volatile_store(Runtime::builder());
|
||||
}
|
||||
|
||||
/// The two paths a downstream crate re-exports these from, written from
|
||||
/// outside mentra so a rename is a failing test rather than a broken host.
|
||||
///
|
||||
/// basis re-exports `ProviderRetry` and `ResponsesTransport` from its own
|
||||
/// `basis::runtime` so a host of *its* never names mentra, which makes both
|
||||
/// paths part of this crate's contract rather than an implementation detail
|
||||
/// that happens to be reachable. `ResponsesTransport` is visible at the crate
|
||||
/// root too; `mentra::provider::` is the one to depend on, since it sits with
|
||||
/// `ResponsesRequestOptions` and the rest of the wire vocabulary.
|
||||
#[test]
|
||||
fn the_retry_schedule_and_transport_are_types_downstream_code_can_name() {
|
||||
fn patient(retry: mentra::runtime::ProviderRetry) -> mentra::runtime::RunOptions {
|
||||
mentra::runtime::RunOptions::default().with_provider_retry(retry)
|
||||
}
|
||||
|
||||
fn over(transport: mentra::provider::ResponsesTransport) -> mentra::RuntimeBuilder {
|
||||
Runtime::empty_builder().with_responses_transport(transport)
|
||||
}
|
||||
|
||||
let options = patient(mentra::runtime::ProviderRetry {
|
||||
base_delay: std::time::Duration::from_secs(1),
|
||||
max_delay: std::time::Duration::from_secs(30),
|
||||
..Default::default()
|
||||
});
|
||||
let _ = over(mentra::provider::ResponsesTransport::HttpSse);
|
||||
|
||||
// The delegation contract basis depends on, asserted from outside: a
|
||||
// subagent meets the same provider with the same patience, so a host
|
||||
// states its schedule once rather than at every boundary.
|
||||
let child = options.child();
|
||||
assert_eq!(child.provider_retry, options.provider_retry);
|
||||
assert_eq!(child.retry_budget, options.retry_budget);
|
||||
}
|
||||
|
||||
/// The downstream shape this exists for, written from outside the crate: a
|
||||
/// host registers an executor that serves named targets and a tool that names
|
||||
/// one. Every guard around a shell command still applies — only the executor
|
||||
/// reads the name. Compiling is half the claim; the other half is that the
|
||||
/// name survives the trip.
|
||||
#[tokio::test]
|
||||
async fn a_tool_can_name_the_executor_a_command_runs_on() {
|
||||
#[derive(Clone, Default)]
|
||||
struct TargetLog(Arc<Mutex<Vec<Option<String>>>>);
|
||||
|
||||
#[async_trait]
|
||||
impl RuntimeExecutor for TargetLog {
|
||||
async fn run(&self, request: CommandRequest) -> Result<CommandOutput, String> {
|
||||
self.0
|
||||
.lock()
|
||||
.expect("target log poisoned")
|
||||
.push(request.target.clone());
|
||||
Ok(CommandOutput {
|
||||
stdout: format!("ran on {}", request.target.unwrap_or("local".to_string())),
|
||||
stderr: String::new(),
|
||||
success: true,
|
||||
status_code: Some(0),
|
||||
timed_out: false,
|
||||
stdout_truncated: false,
|
||||
stderr_truncated: false,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
struct TargetedShellTool;
|
||||
|
||||
#[async_trait]
|
||||
impl ToolDefinition for TargetedShellTool {
|
||||
fn descriptor(&self) -> ToolSpec {
|
||||
ToolSpec::builder("targeted_shell")
|
||||
.description("Run a command on a named host")
|
||||
.input_schema(json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"target": { "type": "string" },
|
||||
"command": { "type": "string" }
|
||||
},
|
||||
"required": ["command"]
|
||||
}))
|
||||
.build()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ToolExecutor for TargetedShellTool {
|
||||
async fn execute(&self, ctx: ParallelToolContext, input: Value) -> ToolResult {
|
||||
let command = input
|
||||
.get("command")
|
||||
.and_then(Value::as_str)
|
||||
.ok_or_else(|| "command is required".to_string())?
|
||||
.to_string();
|
||||
let target = input
|
||||
.get("target")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::to_string);
|
||||
let cwd = ctx.resolve_working_directory(None)?;
|
||||
let output = ctx
|
||||
.execute_shell_command_on(target, command, None, None, cwd)
|
||||
.await?;
|
||||
Ok(output.stdout)
|
||||
}
|
||||
}
|
||||
|
||||
let log = TargetLog::default();
|
||||
let model = ModelInfo::new("mock-model", BuiltinProvider::OpenAI);
|
||||
let provider = ScriptedProvider::new(model.provider.clone(), vec![model.clone()]);
|
||||
provider.push_turns(vec![
|
||||
Turn::ToolCalls(vec![ScriptedToolCall::new(
|
||||
"targeted_shell",
|
||||
json!({ "target": "mac", "command": "xcodebuild -version" }),
|
||||
)]),
|
||||
Turn::Text("done".to_string()),
|
||||
]);
|
||||
|
||||
let runtime = Runtime::builder()
|
||||
.with_runtime_identifier(format!("public-api-{}", now_nanos()))
|
||||
.with_store(VolatileRuntimeStore::new())
|
||||
.with_provider_instance(provider)
|
||||
.with_policy(RuntimePolicy::permissive())
|
||||
.with_executor(log.clone())
|
||||
.build()
|
||||
.expect("build runtime");
|
||||
runtime.register_tool(TargetedShellTool);
|
||||
let mut agent = runtime.spawn("target-agent", model).expect("spawn agent");
|
||||
|
||||
agent
|
||||
.send(vec![ContentBlock::text("build it on the mac")])
|
||||
.await
|
||||
.expect("run completes");
|
||||
|
||||
assert_eq!(
|
||||
log.0.lock().expect("target log poisoned").as_slice(),
|
||||
[Some("mac".to_string())]
|
||||
);
|
||||
}
|
||||
|
||||
fn now_nanos() -> u128 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_nanos()
|
||||
}
|
||||
|
||||
fn spawn_models_server(response_body: &str) -> (String, thread::JoinHandle<String>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server");
|
||||
let addr = listener.local_addr().expect("read listener addr");
|
||||
let response_body = response_body.to_string();
|
||||
|
||||
let handle = thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().expect("accept request");
|
||||
let mut request = Vec::new();
|
||||
let mut temp = [0_u8; 1024];
|
||||
|
||||
loop {
|
||||
let read = stream.read(&mut temp).expect("read request");
|
||||
if read == 0 {
|
||||
break;
|
||||
}
|
||||
request.extend_from_slice(&temp[..read]);
|
||||
if request.windows(4).any(|window| window == b"\r\n\r\n") {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
let response = format!(
|
||||
concat!(
|
||||
"HTTP/1.1 200 OK\r\n",
|
||||
"content-type: application/json\r\n",
|
||||
"content-length: {}\r\n\r\n",
|
||||
"{}"
|
||||
),
|
||||
response_body.len(),
|
||||
response_body
|
||||
);
|
||||
stream
|
||||
.write_all(response.as_bytes())
|
||||
.expect("write response");
|
||||
|
||||
String::from_utf8(request).expect("request should be valid utf8")
|
||||
});
|
||||
|
||||
(format!("http://{addr}/"), handle)
|
||||
}
|
||||
|
||||
fn model_with_created_at(id: &str, provider: BuiltinProvider, unix_timestamp: i64) -> ModelInfo {
|
||||
let mut model = ModelInfo::new(id, provider);
|
||||
model.created_at = Some(
|
||||
time::OffsetDateTime::from_unix_timestamp(unix_timestamp)
|
||||
.expect("timestamp should be valid"),
|
||||
);
|
||||
model
|
||||
}
|
||||
254
vendor/mentra/tests/responses_transport.rs
vendored
Normal file
254
vendor/mentra/tests/responses_transport.rs
vendored
Normal file
@@ -0,0 +1,254 @@
|
||||
//! Which transport a runtime's Responses requests go out on, and what happens
|
||||
//! when a provider cannot serve the one it was handed.
|
||||
|
||||
use std::{
|
||||
fs,
|
||||
sync::{
|
||||
Arc, Mutex,
|
||||
atomic::{AtomicU64, Ordering},
|
||||
},
|
||||
time::{SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use mentra::{
|
||||
AgentConfig, BuiltinProvider, ContentBlock, ProviderCapabilities, Role, Runtime,
|
||||
provider::{
|
||||
ContentBlockDelta, ContentBlockStart, ModelInfo, Provider, ProviderDescriptor,
|
||||
ProviderError, ProviderEvent, ProviderEventStream, ProviderId, ProviderRequestOptions,
|
||||
Request, ResponsesTransport,
|
||||
},
|
||||
runtime::SqliteRuntimeStore,
|
||||
};
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
static NEXT_TEMP_ID: AtomicU64 = AtomicU64::new(1);
|
||||
|
||||
/// Answers one short text turn and remembers the options it was handed.
|
||||
///
|
||||
/// The options are what this file is about: the transport is settled before the
|
||||
/// request leaves the runtime, so what arrives here is the only evidence of
|
||||
/// which one was chosen.
|
||||
#[derive(Clone)]
|
||||
struct RecordingProvider {
|
||||
id: ProviderId,
|
||||
display_name: Option<String>,
|
||||
capabilities: ProviderCapabilities,
|
||||
seen: Arc<Mutex<Vec<ProviderRequestOptions>>>,
|
||||
}
|
||||
|
||||
impl RecordingProvider {
|
||||
fn new(id: impl Into<ProviderId>, supports_websockets: bool) -> Self {
|
||||
Self {
|
||||
id: id.into(),
|
||||
display_name: None,
|
||||
capabilities: ProviderCapabilities {
|
||||
supports_streaming: true,
|
||||
supports_websockets,
|
||||
..ProviderCapabilities::default()
|
||||
},
|
||||
seen: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn named(mut self, display_name: &str) -> Self {
|
||||
self.display_name = Some(display_name.to_string());
|
||||
self
|
||||
}
|
||||
|
||||
fn transports(&self) -> Vec<ResponsesTransport> {
|
||||
self.seen
|
||||
.lock()
|
||||
.expect("recorded options")
|
||||
.iter()
|
||||
.map(|options| options.responses.transport)
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for RecordingProvider {
|
||||
fn descriptor(&self) -> ProviderDescriptor {
|
||||
ProviderDescriptor {
|
||||
id: self.id.clone(),
|
||||
display_name: self.display_name.clone(),
|
||||
description: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
self.capabilities
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
|
||||
Ok(vec![ModelInfo::new("model", self.id.clone())])
|
||||
}
|
||||
|
||||
async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
|
||||
self.seen
|
||||
.lock()
|
||||
.expect("recorded options")
|
||||
.push(request.provider_request_options.clone());
|
||||
|
||||
let (tx, rx) = mpsc::unbounded_channel();
|
||||
for event in [
|
||||
ProviderEvent::MessageStarted {
|
||||
id: "msg-1".to_string(),
|
||||
model: "model".to_string(),
|
||||
role: Role::Assistant,
|
||||
},
|
||||
ProviderEvent::ContentBlockStarted {
|
||||
index: 0,
|
||||
kind: ContentBlockStart::Text,
|
||||
},
|
||||
ProviderEvent::ContentBlockDelta {
|
||||
index: 0,
|
||||
delta: ContentBlockDelta::Text("ok".to_string()),
|
||||
},
|
||||
ProviderEvent::ContentBlockStopped { index: 0 },
|
||||
ProviderEvent::MessageStopped,
|
||||
] {
|
||||
tx.send(Ok(event)).expect("test receiver alive");
|
||||
}
|
||||
Ok(rx)
|
||||
}
|
||||
}
|
||||
|
||||
fn temp_store() -> SqliteRuntimeStore {
|
||||
let unique = NEXT_TEMP_ID.fetch_add(1, Ordering::Relaxed);
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.expect("system time")
|
||||
.as_nanos();
|
||||
let path = std::env::temp_dir().join(format!(
|
||||
"mentra-responses-transport-{timestamp}-{unique}.sqlite"
|
||||
));
|
||||
if let Some(parent) = path.parent() {
|
||||
fs::create_dir_all(parent).expect("create temp dir");
|
||||
}
|
||||
SqliteRuntimeStore::new(path)
|
||||
}
|
||||
|
||||
/// A runtime around `provider`, optionally told which transport to use.
|
||||
fn runtime_for(provider: RecordingProvider, transport: Option<ResponsesTransport>) -> Runtime {
|
||||
let mut builder = Runtime::empty_builder()
|
||||
.with_provider_instance(provider)
|
||||
.with_store(temp_store());
|
||||
if let Some(transport) = transport {
|
||||
builder = builder.with_responses_transport(transport);
|
||||
}
|
||||
builder.build().expect("build runtime")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_runtime_that_chooses_nothing_still_streams_over_http_sse() {
|
||||
let provider = RecordingProvider::new(BuiltinProvider::OpenAI, true);
|
||||
let recorder = provider.clone();
|
||||
let runtime = runtime_for(provider, None);
|
||||
|
||||
let mut agent = runtime
|
||||
.spawn("agent", ModelInfo::new("model", BuiltinProvider::OpenAI))
|
||||
.expect("spawn agent");
|
||||
agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("the turn runs");
|
||||
|
||||
assert_eq!(recorder.transports(), vec![ResponsesTransport::HttpSse]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_runtime_reports_the_transport_it_was_given() {
|
||||
// Without this reader the choice is write-only: the only evidence a host's
|
||||
// selection reached the runtime is a turn run against a provider that
|
||||
// records what it was handed, so anything downstream can test its own field
|
||||
// and stop at the seam — the shape of test that passes while the wiring
|
||||
// between the two is broken.
|
||||
let unset = runtime_for(RecordingProvider::new(BuiltinProvider::OpenAI, true), None);
|
||||
assert_eq!(unset.responses_transport(), None);
|
||||
|
||||
let chosen = runtime_for(
|
||||
RecordingProvider::new(BuiltinProvider::OpenAI, true),
|
||||
Some(ResponsesTransport::WebSocket),
|
||||
);
|
||||
assert_eq!(
|
||||
chosen.responses_transport(),
|
||||
Some(ResponsesTransport::WebSocket)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_chosen_websocket_transport_reaches_the_request() {
|
||||
// The gap this closes: `ResponsesRequestOptions.transport` existed and the
|
||||
// websocket path was compiled in, but nothing in the runtime ever set the
|
||||
// field, so every request went out over HTTP+SSE whatever the host wanted.
|
||||
let provider = RecordingProvider::new(BuiltinProvider::OpenAI, true);
|
||||
let recorder = provider.clone();
|
||||
let runtime = runtime_for(provider, Some(ResponsesTransport::WebSocket));
|
||||
|
||||
let mut agent = runtime
|
||||
.spawn("agent", ModelInfo::new("model", BuiltinProvider::OpenAI))
|
||||
.expect("spawn agent");
|
||||
agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("the turn runs");
|
||||
|
||||
assert_eq!(recorder.transports(), vec![ResponsesTransport::WebSocket]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_runtime_choice_settles_a_disagreeing_agent_config() {
|
||||
// Two live opinions about one socket is not a state the runtime keeps: the
|
||||
// connection-level answer is the one that holds.
|
||||
let provider = RecordingProvider::new(BuiltinProvider::OpenAI, true);
|
||||
let recorder = provider.clone();
|
||||
let runtime = runtime_for(provider, Some(ResponsesTransport::HttpSse));
|
||||
|
||||
let mut config = AgentConfig::default();
|
||||
config.provider_request_options.responses.transport = ResponsesTransport::WebSocket;
|
||||
let mut agent = runtime
|
||||
.spawn_with_config(
|
||||
"agent",
|
||||
ModelInfo::new("model", BuiltinProvider::OpenAI),
|
||||
config,
|
||||
)
|
||||
.expect("spawn agent");
|
||||
agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect("the turn runs");
|
||||
|
||||
assert_eq!(recorder.transports(), vec![ResponsesTransport::HttpSse]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_provider_without_websockets_refuses_rather_than_pretending() {
|
||||
// anthropic and gemini report `supports_websockets: false`. Answering over
|
||||
// HTTP+SSE would look like success and be a transport nobody asked for.
|
||||
let provider = RecordingProvider::new(BuiltinProvider::Anthropic, false).named("Anthropic");
|
||||
let recorder = provider.clone();
|
||||
let runtime = runtime_for(provider, Some(ResponsesTransport::WebSocket));
|
||||
|
||||
let mut agent = runtime
|
||||
.spawn("agent", ModelInfo::new("model", BuiltinProvider::Anthropic))
|
||||
.expect("spawn agent");
|
||||
let error = agent
|
||||
.send(vec![ContentBlock::text("hello")])
|
||||
.await
|
||||
.expect_err("a transport the provider cannot serve is refused");
|
||||
|
||||
let message = error.to_string();
|
||||
assert!(
|
||||
message.contains("Anthropic"),
|
||||
"the refusal must name the provider: {message}"
|
||||
);
|
||||
assert!(
|
||||
message.contains("websocket"),
|
||||
"and say what it could not do: {message}"
|
||||
);
|
||||
assert!(
|
||||
recorder.transports().is_empty(),
|
||||
"the request must never have been sent"
|
||||
);
|
||||
}
|
||||
145
vendor/mentra/tests/skills_api.rs
vendored
Normal file
145
vendor/mentra/tests/skills_api.rs
vendored
Normal file
@@ -0,0 +1,145 @@
|
||||
//! Public-API tests for skill registration and enumeration.
|
||||
//!
|
||||
//! These exercise what a host can actually reach: registering roots through
|
||||
//! `Runtime`, listing what loaded, and naming the error type in its own
|
||||
//! signatures.
|
||||
|
||||
use std::{
|
||||
fs,
|
||||
path::{Path, PathBuf},
|
||||
sync::atomic::{AtomicU64, Ordering},
|
||||
time::{SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use mentra::{BuiltinProvider, Runtime, SkillInfo, SkillLoadError};
|
||||
|
||||
static NEXT_ID: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
fn temp_dir(label: &str) -> PathBuf {
|
||||
let unique = NEXT_ID.fetch_add(1, Ordering::Relaxed);
|
||||
let stamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.expect("system time")
|
||||
.as_nanos();
|
||||
let path = std::env::temp_dir().join(format!("mentra-skills-api-{label}-{stamp}-{unique}"));
|
||||
fs::create_dir_all(&path).expect("create temp dir");
|
||||
path
|
||||
}
|
||||
|
||||
fn write_skill(root: &Path, dir: &str, name: &str, description: &str, body: &str) {
|
||||
let skill_dir = root.join(dir);
|
||||
fs::create_dir_all(&skill_dir).expect("create skill dir");
|
||||
fs::write(
|
||||
skill_dir.join("SKILL.md"),
|
||||
format!("---\nname: {name}\ndescription: {description}\n---\n{body}\n"),
|
||||
)
|
||||
.expect("write skill");
|
||||
}
|
||||
|
||||
fn runtime() -> Runtime {
|
||||
Runtime::builder()
|
||||
.with_provider(BuiltinProvider::OpenAI, "test-key")
|
||||
.build()
|
||||
.expect("runtime builds")
|
||||
}
|
||||
|
||||
/// The error type must be nameable by a caller — this signature is the test.
|
||||
fn register(runtime: &Runtime, path: &Path) -> Result<(), SkillLoadError> {
|
||||
runtime.register_skills_dir(path)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_host_can_name_the_error_type() {
|
||||
let runtime = runtime();
|
||||
let missing = temp_dir("nameable").join("does-not-exist");
|
||||
|
||||
let error = register(&runtime, &missing).expect_err("an unreadable root is an error");
|
||||
|
||||
// And match on it, which is the point of it being an enum.
|
||||
assert!(matches!(error, SkillLoadError::ReadDir { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registering_a_second_root_keeps_the_first() {
|
||||
let workspace = temp_dir("layer-workspace");
|
||||
write_skill(&workspace, "review", "review", "project review", "PROJECT");
|
||||
let global = temp_dir("layer-global");
|
||||
write_skill(&global, "review", "review", "personal review", "PERSONAL");
|
||||
write_skill(&global, "deploy", "deploy", "personal deploy", "DEPLOY");
|
||||
|
||||
let runtime = runtime();
|
||||
runtime
|
||||
.register_skills_dirs([workspace.as_path(), global.as_path()])
|
||||
.expect("both roots register");
|
||||
|
||||
let skills = runtime.skills();
|
||||
let names: Vec<&str> = skills.iter().map(|skill| skill.name.as_str()).collect();
|
||||
assert_eq!(names, vec!["deploy", "review"]);
|
||||
|
||||
let review = skills
|
||||
.iter()
|
||||
.find(|skill| skill.name == "review")
|
||||
.expect("review present");
|
||||
assert_eq!(
|
||||
review.description, "project review",
|
||||
"the earlier root must win the collision"
|
||||
);
|
||||
assert!(review.path.starts_with(&workspace));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn enumeration_reports_name_description_and_source() {
|
||||
let root = temp_dir("enumerate");
|
||||
write_skill(&root, "haiku", "haiku", "writes haiku", "BODY");
|
||||
|
||||
let runtime = runtime();
|
||||
runtime.register_skills_dir(&root).expect("registers");
|
||||
|
||||
let skills = runtime.skills();
|
||||
assert_eq!(skills.len(), 1);
|
||||
let SkillInfo {
|
||||
name,
|
||||
description,
|
||||
path,
|
||||
} = &skills[0];
|
||||
assert_eq!(name, "haiku");
|
||||
assert_eq!(description, "writes haiku");
|
||||
assert_eq!(path, &root.join("haiku").join("SKILL.md"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_runtime_without_skills_lists_none() {
|
||||
assert!(runtime().skills().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_duplicate_name_inside_one_root_is_still_an_error() {
|
||||
let root = temp_dir("duplicate");
|
||||
write_skill(&root, "first", "shared", "one", "A");
|
||||
write_skill(&root, "second", "shared", "two", "B");
|
||||
|
||||
let error = runtime()
|
||||
.register_skills_dir(&root)
|
||||
.expect_err("a repeated name in one root is a mistake");
|
||||
|
||||
assert!(matches!(error, SkillLoadError::DuplicateSkillName { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn roots_before_a_failing_one_stay_registered() {
|
||||
let good = temp_dir("partial-good");
|
||||
write_skill(&good, "keep", "keep", "kept", "BODY");
|
||||
let missing = temp_dir("partial-missing").join("absent");
|
||||
|
||||
let runtime = runtime();
|
||||
let error = runtime
|
||||
.register_skills_dirs([good.as_path(), missing.as_path()])
|
||||
.expect_err("the second root fails");
|
||||
|
||||
assert!(matches!(error, SkillLoadError::ReadDir { .. }));
|
||||
assert_eq!(
|
||||
runtime.skills().len(),
|
||||
1,
|
||||
"the root that loaded before the failure stays registered"
|
||||
);
|
||||
}
|
||||
Reference in New Issue
Block a user