Add operator Mentra memory browser
Some checks failed
CI / rust-skia (Rust only) (push) Has been cancelled
CI / required (push) Has been cancelled

This commit is contained in:
2026-08-23 19:18:41 +02:00
parent 95cca3e777
commit 360203d647
190 changed files with 70246 additions and 45 deletions

562
vendor/mentra/tests/agent_runtime.rs vendored Normal file
View 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
View 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
View 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
View 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
}

View 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
View 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"
);
}