387 lines
15 KiB
Rust
387 lines
15 KiB
Rust
use async_trait::async_trait;
|
|
use mentra::memory::MemoryStore as _;
|
|
use mentra::runtime::HybridRuntimeStore;
|
|
use mentra::tool::{ParallelToolContext, RuntimeToolDescriptor, ToolResult};
|
|
use mentra::{BuiltinProvider, ContentBlock, ModelInfo, Runtime};
|
|
use metacrate_grid_agent::{
|
|
AgentConfig, BoundedText, LlmClient, LlmError, ToolDefinition, ToolSchema,
|
|
};
|
|
use serde_json::{Value, json};
|
|
use std::collections::{BTreeMap, BTreeSet};
|
|
use std::fmt::Write as _;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
use tokio::net::TcpListener;
|
|
use tokio::sync::mpsc;
|
|
|
|
static NEXT_TEST: AtomicU64 = AtomicU64::new(1);
|
|
|
|
#[derive(Clone)]
|
|
struct EchoTool;
|
|
|
|
impl mentra::tool::ToolDefinition for EchoTool {
|
|
fn descriptor(&self) -> RuntimeToolDescriptor {
|
|
RuntimeToolDescriptor::builder("echo")
|
|
.description("Echo a target.")
|
|
.input_schema(json!({
|
|
"type":"object",
|
|
"properties":{"target":{"type":"string"}},
|
|
"required":["target"],
|
|
"additionalProperties":false
|
|
}))
|
|
.non_strict()
|
|
.build()
|
|
}
|
|
}
|
|
|
|
#[async_trait]
|
|
impl mentra::tool::ToolExecutor for EchoTool {
|
|
async fn execute(&self, _: ParallelToolContext, input: Value) -> ToolResult {
|
|
Ok(format!(
|
|
"echoed {}",
|
|
input["target"].as_str().ok_or("target is required")?
|
|
))
|
|
}
|
|
}
|
|
|
|
struct CapturedRequest {
|
|
head: String,
|
|
body: Value,
|
|
}
|
|
|
|
async fn fake_responses(
|
|
responses: Vec<String>,
|
|
) -> (
|
|
String,
|
|
mpsc::Receiver<CapturedRequest>,
|
|
tokio::task::JoinHandle<()>,
|
|
) {
|
|
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
|
|
.await
|
|
.expect("bind fake provider");
|
|
let address = listener.local_addr().expect("listener address");
|
|
let (sender, receiver) = mpsc::channel(responses.len());
|
|
let task = tokio::spawn(async move {
|
|
for body in responses {
|
|
let (mut stream, _) = listener.accept().await.expect("provider connection");
|
|
sender
|
|
.send(read_request(&mut stream).await)
|
|
.await
|
|
.expect("capture request");
|
|
let head = format!(
|
|
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
|
|
body.len()
|
|
);
|
|
stream.write_all(head.as_bytes()).await.expect("head");
|
|
stream.write_all(body.as_bytes()).await.expect("body");
|
|
stream.shutdown().await.expect("shutdown");
|
|
}
|
|
});
|
|
(format!("http://{address}/v1"), receiver, task)
|
|
}
|
|
|
|
async fn read_request(stream: &mut tokio::net::TcpStream) -> CapturedRequest {
|
|
let mut bytes = Vec::new();
|
|
let header_end = loop {
|
|
let mut chunk = [0_u8; 4096];
|
|
let count = stream.read(&mut chunk).await.expect("request read");
|
|
assert!(count > 0 && bytes.len() < 8 * 1024 * 1024);
|
|
bytes.extend_from_slice(&chunk[..count]);
|
|
if let Some(position) = bytes.windows(4).position(|window| window == b"\r\n\r\n") {
|
|
break position + 4;
|
|
}
|
|
};
|
|
let head = String::from_utf8(bytes[..header_end].to_vec()).expect("header UTF-8");
|
|
let length = head
|
|
.lines()
|
|
.find_map(|line| {
|
|
let (name, value) = line.split_once(':')?;
|
|
name.eq_ignore_ascii_case("content-length")
|
|
.then(|| value.trim().parse::<usize>().expect("content length"))
|
|
})
|
|
.expect("content length");
|
|
while bytes.len() - header_end < length {
|
|
let mut chunk = [0_u8; 4096];
|
|
let count = stream.read(&mut chunk).await.expect("body read");
|
|
assert!(count > 0);
|
|
bytes.extend_from_slice(&chunk[..count]);
|
|
}
|
|
CapturedRequest {
|
|
head,
|
|
body: serde_json::from_slice(&bytes[header_end..header_end + length]).expect("JSON body"),
|
|
}
|
|
}
|
|
|
|
fn events(events: impl IntoIterator<Item = Value>) -> String {
|
|
events.into_iter().fold(String::new(), |mut body, event| {
|
|
write!(body, "data: {event}\n\n").expect("writing to String cannot fail");
|
|
body
|
|
})
|
|
}
|
|
|
|
fn tool_response() -> String {
|
|
events([
|
|
json!({"type":"response.created","response":{"id":"resp-tool","model":"test-model","status":"in_progress"}}),
|
|
json!({"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"call-1","name":"echo","arguments":""}}),
|
|
json!({"type":"response.function_call_arguments.delta","output_index":0,"delta":"{\"target\":\"grid\"}"}),
|
|
json!({"type":"response.output_item.done","output_index":0,"item":{"type":"function_call","call_id":"call-1","name":"echo","arguments":"{\"target\":\"grid\"}"}}),
|
|
json!({"type":"response.completed","response":{"id":"resp-tool","model":"test-model","status":"completed","usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}),
|
|
])
|
|
}
|
|
|
|
fn memory_pin_response() -> String {
|
|
events([
|
|
json!({"type":"response.created","response":{"id":"resp-memory","model":"test-model","status":"in_progress"}}),
|
|
json!({"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"memory-1","name":"memory_pin","arguments":""}}),
|
|
json!({"type":"response.function_call_arguments.delta","output_index":0,"delta":"{\"content\":\"The welcome area fountain is north of the landing point.\"}"}),
|
|
json!({"type":"response.output_item.done","output_index":0,"item":{"type":"function_call","call_id":"memory-1","name":"memory_pin","arguments":"{\"content\":\"The welcome area fountain is north of the landing point.\"}"}}),
|
|
json!({"type":"response.completed","response":{"id":"resp-memory","model":"test-model","status":"completed","usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}),
|
|
])
|
|
}
|
|
|
|
fn text_response(id: &str, text: &str) -> String {
|
|
events([
|
|
json!({"type":"response.created","response":{"id":id,"model":"test-model","status":"in_progress"}}),
|
|
json!({"type":"response.output_item.added","output_index":0,"item":{"type":"message","content":[]}}),
|
|
json!({"type":"response.output_text.delta","output_index":0,"content_index":0,"delta":text}),
|
|
json!({"type":"response.output_item.done","output_index":0,"item":{"type":"message","content":[{"type":"output_text","text":text}]}}),
|
|
json!({"type":"response.completed","response":{"id":id,"model":"test-model","status":"completed","usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}),
|
|
])
|
|
}
|
|
|
|
fn agent_config() -> mentra::AgentConfig {
|
|
let root = std::env::temp_dir().join(format!(
|
|
"metacrate-mentra-test-{}-{}",
|
|
std::process::id(),
|
|
NEXT_TEST.fetch_add(1, Ordering::Relaxed)
|
|
));
|
|
let mut config = mentra::AgentConfig {
|
|
system: Some("Use the registered grid tools.".to_owned()),
|
|
tool_profile: mentra::agent::ToolProfile::only(["echo"]),
|
|
memory: mentra::agent::MemoryConfig {
|
|
auto_recall_enabled: false,
|
|
write_tools_enabled: false,
|
|
..Default::default()
|
|
},
|
|
..Default::default()
|
|
};
|
|
config.compaction.transcript_dir = root.join("transcripts");
|
|
config.task.tasks_dir = root.join("tasks");
|
|
config.team.team_dir = root.join("teams");
|
|
config.workspace.base_dir = root;
|
|
config
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn mentra_owns_responses_sse_images_and_tool_rounds() {
|
|
let (endpoint, mut requests, server) =
|
|
fake_responses(vec![tool_response(), text_response("resp-final", "done")]).await;
|
|
let connection = AgentConfig::offline(&endpoint, "super-secret-api-key")
|
|
.expect("config")
|
|
.llm;
|
|
let client = LlmClient::new(connection);
|
|
assert_eq!(
|
|
client.mentra_provider().definition().base_url.as_deref(),
|
|
Some(endpoint.as_str())
|
|
);
|
|
let runtime = Runtime::empty_builder()
|
|
.with_store(mentra::runtime::VolatileRuntimeStore::default())
|
|
.with_registered_provider(client.mentra_provider())
|
|
.with_tool(EchoTool)
|
|
.build()
|
|
.expect("Mentra runtime");
|
|
let mut agent = runtime
|
|
.spawn_with_config(
|
|
"grid-agent",
|
|
ModelInfo::new(client.configured_model(), BuiltinProvider::OpenAI),
|
|
agent_config(),
|
|
)
|
|
.expect("Mentra agent");
|
|
let message = agent
|
|
.send(vec![
|
|
ContentBlock::text("Inspect this and use echo."),
|
|
ContentBlock::image_url("data:image/jpeg;base64,/9j/2Q=="),
|
|
])
|
|
.await
|
|
.expect("tool-using turn");
|
|
assert!(
|
|
message
|
|
.content
|
|
.iter()
|
|
.any(|part| matches!(part, ContentBlock::Text { text } if text == "done"))
|
|
);
|
|
|
|
let first = requests.recv().await.expect("first request");
|
|
assert!(first.head.starts_with("POST /v1/responses HTTP/1.1"));
|
|
assert!(
|
|
first
|
|
.head
|
|
.to_ascii_lowercase()
|
|
.contains("authorization: bearer super-secret-api-key")
|
|
);
|
|
assert_eq!(first.body["tools"][0]["name"], "echo");
|
|
assert_eq!(first.body["input"][0]["content"][1]["type"], "input_image");
|
|
assert_eq!(
|
|
first.body["input"][0]["content"][1]["image_url"],
|
|
"data:image/jpeg;base64,/9j/2Q=="
|
|
);
|
|
let second = requests.recv().await.expect("tool result request");
|
|
assert!(
|
|
serde_json::to_string(&second.body)
|
|
.expect("request JSON")
|
|
.contains("echoed grid")
|
|
);
|
|
server.await.expect("server");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn mentra_persists_turns_and_auto_compacts_long_conversations() {
|
|
let (endpoint, mut requests, server) = fake_responses(vec![
|
|
text_response("resp-first", "first done"),
|
|
text_response("resp-summary", "important facts retained"),
|
|
text_response("resp-second", "second done"),
|
|
])
|
|
.await;
|
|
let connection = AgentConfig::offline(&endpoint, "test-key")
|
|
.expect("config")
|
|
.llm;
|
|
let client = LlmClient::new(connection);
|
|
let runtime = Runtime::empty_builder()
|
|
.with_store(mentra::runtime::VolatileRuntimeStore::default())
|
|
.with_registered_provider(client.mentra_provider())
|
|
.with_tool(EchoTool)
|
|
.build()
|
|
.expect("Mentra runtime");
|
|
let mut config = agent_config();
|
|
config.compaction.auto_compact_threshold_tokens = Some(1);
|
|
let mut agent = runtime
|
|
.spawn_with_config(
|
|
"conversation-agent",
|
|
ModelInfo::new(client.configured_model(), BuiltinProvider::OpenAI),
|
|
config,
|
|
)
|
|
.expect("Mentra agent");
|
|
let mut events = agent.subscribe_events();
|
|
agent
|
|
.send(vec![ContentBlock::text("first")])
|
|
.await
|
|
.expect("first turn");
|
|
agent
|
|
.send(vec![ContentBlock::text("second")])
|
|
.await
|
|
.expect("second turn");
|
|
assert!(agent.history().iter().any(|message| {
|
|
message.content.iter().any(
|
|
|part| matches!(part, ContentBlock::Text { text } if text.contains("[Compaction summary]"))
|
|
)
|
|
}));
|
|
let mut compacted = false;
|
|
while let Ok(event) = events.try_recv() {
|
|
compacted |= matches!(event, mentra::agent::AgentEvent::ContextCompacted { .. });
|
|
}
|
|
assert!(compacted, "Mentra must emit the compaction event");
|
|
for _ in 0..3 {
|
|
assert!(
|
|
requests
|
|
.recv()
|
|
.await
|
|
.expect("compaction request sequence")
|
|
.body
|
|
.get("previous_response_id")
|
|
.is_none(),
|
|
"MetaCrate uses Mentra transcript replay for proxy compatibility"
|
|
);
|
|
}
|
|
server.await.expect("server");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn mentra_persists_agents_and_searchable_environment_memory_across_restart() {
|
|
let (endpoint, mut requests, server) = fake_responses(vec![
|
|
memory_pin_response(),
|
|
text_response("resp-memory-done", "remembered"),
|
|
])
|
|
.await;
|
|
let client = LlmClient::new(
|
|
AgentConfig::offline(&endpoint, "test-key")
|
|
.expect("config")
|
|
.llm,
|
|
);
|
|
let root = std::env::temp_dir().join(format!(
|
|
"metacrate-mentra-persistence-{}-{}",
|
|
std::process::id(),
|
|
NEXT_TEST.fetch_add(1, Ordering::Relaxed)
|
|
));
|
|
let store = HybridRuntimeStore::with_memory_path(
|
|
root.join("runtime.sqlite"),
|
|
root.join("memory.sqlite"),
|
|
);
|
|
let runtime = Runtime::builder()
|
|
.with_runtime_identifier("metacrate-test")
|
|
.with_store(store.clone())
|
|
.with_registered_provider(client.mentra_provider())
|
|
.build()
|
|
.expect("Mentra runtime");
|
|
let mut config = agent_config();
|
|
config.tool_profile =
|
|
mentra::agent::ToolProfile::only(["memory_search", "memory_pin", "memory_forget"]);
|
|
config.memory = mentra::agent::MemoryConfig::default();
|
|
let mut agent = runtime
|
|
.spawn_with_config(
|
|
"grid-conversation-authorized-test",
|
|
ModelInfo::new(client.configured_model(), BuiltinProvider::OpenAI),
|
|
config,
|
|
)
|
|
.expect("Mentra agent");
|
|
let agent_id = agent.id().to_owned();
|
|
agent
|
|
.send(vec![ContentBlock::text("Remember what you learned here.")])
|
|
.await
|
|
.expect("memory turn");
|
|
for _ in 0..2 {
|
|
requests.recv().await.expect("memory request sequence");
|
|
}
|
|
server.await.expect("server");
|
|
assert!(
|
|
store
|
|
.search_records(&agent_id, "welcome fountain", 10)
|
|
.expect("search persisted memory")
|
|
.iter()
|
|
.any(|record| record.content.contains("north of the landing point"))
|
|
);
|
|
drop(agent);
|
|
drop(runtime);
|
|
|
|
let rebooted = Runtime::builder()
|
|
.with_runtime_identifier("metacrate-test")
|
|
.with_store(store)
|
|
.with_registered_provider(client.mentra_provider())
|
|
.build()
|
|
.expect("rebooted Mentra runtime");
|
|
let resumed = rebooted.resume("metacrate-test").expect("resume agents");
|
|
assert_eq!(resumed.len(), 1);
|
|
assert_eq!(resumed[0].id(), agent_id);
|
|
assert_eq!(resumed[0].name(), "grid-conversation-authorized-test");
|
|
assert!(!resumed[0].history().is_empty());
|
|
}
|
|
|
|
#[test]
|
|
fn metacrate_still_validates_grid_tool_schemas_before_mentra_registration() {
|
|
let mut properties = BTreeMap::new();
|
|
properties.insert("target".to_owned(), ToolSchema::String);
|
|
let valid = ToolDefinition {
|
|
name: BoundedText::new("name", "face_avatar").expect("name"),
|
|
description: BoundedText::new("description", "Face an avatar.").expect("description"),
|
|
schema: ToolSchema::Object {
|
|
properties,
|
|
required: BTreeSet::from(["target".to_owned()]),
|
|
additional_properties: false,
|
|
},
|
|
mutating: true,
|
|
};
|
|
assert_eq!(valid.validate(), Ok(()));
|
|
let mut invalid = valid;
|
|
invalid.name = BoundedText::new("name", "invalid tool name").expect("bounded");
|
|
assert_eq!(invalid.validate(), Err(LlmError::InvalidToolSchema));
|
|
}
|