951 lines
30 KiB
Rust
951 lines
30 KiB
Rust
use libremetaverse_types::compat::{CancellationToken, CancellationTokenSource};
|
|
use metacrate_grid_agent::{
|
|
AgentConfig, BoundedText, BoundedVec, CompletionMessage, ContentPart, HistorySummarizer,
|
|
ImageDetail, LlmClient, LlmError, LlmTransportLimits, MessageRole, SessionGeneration,
|
|
ToolDefinition, ToolExecution, ToolExecutor, ToolFuture, ToolLoop, ToolLoopError,
|
|
ToolLoopLimits, ToolSchema,
|
|
};
|
|
use serde_json::{Value, json};
|
|
use std::collections::{BTreeMap, BTreeSet};
|
|
use std::fmt::Write as _;
|
|
use std::sync::Arc;
|
|
use std::sync::atomic::{AtomicUsize, Ordering};
|
|
use std::time::Duration;
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
use tokio::net::TcpListener;
|
|
use tokio::sync::{Notify, mpsc};
|
|
|
|
const MAX_REQUEST_BYTES: usize = 2 * 1024 * 1024;
|
|
|
|
#[derive(Clone)]
|
|
struct ResponsePlan {
|
|
status: u16,
|
|
headers: Vec<(String, String)>,
|
|
chunks: Vec<Vec<u8>>,
|
|
delay: Duration,
|
|
}
|
|
|
|
impl ResponsePlan {
|
|
#[allow(clippy::needless_pass_by_value)] // Keeps nested JSON fixtures readable at call sites.
|
|
fn json(value: Value) -> Self {
|
|
Self {
|
|
status: 200,
|
|
headers: vec![("content-type".into(), "application/json".into())],
|
|
chunks: vec![serde_json::to_vec(&value).expect("fixture JSON")],
|
|
delay: Duration::ZERO,
|
|
}
|
|
}
|
|
|
|
fn status(status: u16) -> Self {
|
|
Self {
|
|
status,
|
|
headers: Vec::new(),
|
|
chunks: vec![b"<html>bounded error</html>".to_vec()],
|
|
delay: Duration::ZERO,
|
|
}
|
|
}
|
|
|
|
fn delayed(mut self, delay: Duration) -> Self {
|
|
self.delay = delay;
|
|
self
|
|
}
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
struct CapturedRequest {
|
|
head: String,
|
|
body: Value,
|
|
}
|
|
|
|
struct FakeServer {
|
|
url: String,
|
|
requests: mpsc::Receiver<CapturedRequest>,
|
|
task: tokio::task::JoinHandle<()>,
|
|
max_active: Arc<AtomicUsize>,
|
|
}
|
|
|
|
impl Drop for FakeServer {
|
|
fn drop(&mut self) {
|
|
self.task.abort();
|
|
}
|
|
}
|
|
|
|
async fn fake_server(plans: Vec<ResponsePlan>) -> FakeServer {
|
|
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
|
|
.await
|
|
.expect("bind fake LLM");
|
|
let address = listener.local_addr().expect("fake address");
|
|
let (request_sender, requests) = mpsc::channel(plans.len().max(1));
|
|
let active = Arc::new(AtomicUsize::new(0));
|
|
let max_active = Arc::new(AtomicUsize::new(0));
|
|
let task_active = Arc::clone(&active);
|
|
let task_max = Arc::clone(&max_active);
|
|
let task = tokio::spawn(async move {
|
|
let mut handlers = Vec::with_capacity(plans.len());
|
|
for plan in plans {
|
|
let Ok((stream, _)) = listener.accept().await else {
|
|
break;
|
|
};
|
|
let sender = request_sender.clone();
|
|
let active = Arc::clone(&task_active);
|
|
let max_active = Arc::clone(&task_max);
|
|
handlers.push(tokio::spawn(async move {
|
|
handle_connection(stream, plan, sender, active, max_active).await;
|
|
}));
|
|
}
|
|
for handler in handlers {
|
|
let _ = handler.await;
|
|
}
|
|
});
|
|
FakeServer {
|
|
url: format!("http://{address}/exact/chat?route=operator"),
|
|
requests,
|
|
task,
|
|
max_active,
|
|
}
|
|
}
|
|
|
|
async fn handle_connection(
|
|
mut stream: tokio::net::TcpStream,
|
|
plan: ResponsePlan,
|
|
sender: mpsc::Sender<CapturedRequest>,
|
|
active: Arc<AtomicUsize>,
|
|
max_active: Arc<AtomicUsize>,
|
|
) {
|
|
let now = active.fetch_add(1, Ordering::AcqRel) + 1;
|
|
max_active.fetch_max(now, Ordering::AcqRel);
|
|
let request = read_request(&mut stream).await.expect("bounded request");
|
|
sender.send(request).await.expect("capture owner");
|
|
tokio::time::sleep(plan.delay).await;
|
|
let body_len = plan.chunks.iter().map(Vec::len).sum::<usize>();
|
|
let reason = match plan.status {
|
|
200 => "OK",
|
|
302 => "Found",
|
|
400 => "Bad Request",
|
|
429 => "Too Many Requests",
|
|
500 => "Internal Server Error",
|
|
503 => "Service Unavailable",
|
|
_ => "Response",
|
|
};
|
|
let mut head = format!(
|
|
"HTTP/1.1 {} {}\r\nContent-Length: {}\r\nConnection: close\r\n",
|
|
plan.status, reason, body_len
|
|
);
|
|
for (name, value) in plan.headers {
|
|
write!(head, "{name}: {value}\r\n").expect("write response header");
|
|
}
|
|
head.push_str("\r\n");
|
|
stream
|
|
.write_all(head.as_bytes())
|
|
.await
|
|
.expect("response head");
|
|
for chunk in plan.chunks {
|
|
stream.write_all(&chunk).await.expect("response chunk");
|
|
tokio::task::yield_now().await;
|
|
}
|
|
let _ = stream.shutdown().await;
|
|
active.fetch_sub(1, Ordering::AcqRel);
|
|
}
|
|
|
|
async fn read_request(stream: &mut tokio::net::TcpStream) -> Result<CapturedRequest, ()> {
|
|
let mut bytes = Vec::with_capacity(4096);
|
|
let header_end = loop {
|
|
if bytes.len() >= MAX_REQUEST_BYTES {
|
|
return Err(());
|
|
}
|
|
let mut chunk = [0_u8; 1024];
|
|
let read = stream.read(&mut chunk).await.map_err(|_| ())?;
|
|
if read == 0 {
|
|
return Err(());
|
|
}
|
|
bytes.extend_from_slice(&chunk[..read]);
|
|
if let Some(offset) = bytes.windows(4).position(|window| window == b"\r\n\r\n") {
|
|
break offset + 4;
|
|
}
|
|
};
|
|
let head = String::from_utf8(bytes[..header_end].to_vec()).map_err(|_| ())?;
|
|
let content_length = head
|
|
.lines()
|
|
.find_map(|line| {
|
|
line.split_once(':').and_then(|(name, value)| {
|
|
name.eq_ignore_ascii_case("content-length")
|
|
.then(|| value.trim().parse::<usize>().ok())
|
|
.flatten()
|
|
})
|
|
})
|
|
.ok_or(())?;
|
|
if content_length > MAX_REQUEST_BYTES {
|
|
return Err(());
|
|
}
|
|
while bytes.len() - header_end < content_length {
|
|
let mut chunk = [0_u8; 1024];
|
|
let read = stream.read(&mut chunk).await.map_err(|_| ())?;
|
|
if read == 0 {
|
|
return Err(());
|
|
}
|
|
bytes.extend_from_slice(&chunk[..read]);
|
|
}
|
|
let body =
|
|
serde_json::from_slice(&bytes[header_end..header_end + content_length]).map_err(|_| ())?;
|
|
Ok(CapturedRequest { head, body })
|
|
}
|
|
|
|
fn response_text(text: &str) -> Value {
|
|
json!({
|
|
"id":"ignored-provider-field",
|
|
"choices":[{"message":{"role":"assistant","content":text},"unknown":"kept-compatible"}],
|
|
"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5},
|
|
"unknown_root":{"harmless":true}
|
|
})
|
|
}
|
|
|
|
#[allow(clippy::needless_pass_by_value)] // Keeps nested JSON fixtures readable at call sites.
|
|
fn response_calls(calls: Value) -> Value {
|
|
json!({"choices":[{"message":{"role":"assistant","content":null,"tool_calls":calls}}]})
|
|
}
|
|
|
|
fn call(id: &str, name: &str, arguments: &str) -> Value {
|
|
json!({"id":id,"type":"function","function":{"name":name,"arguments":arguments}})
|
|
}
|
|
|
|
fn limits() -> LlmTransportLimits {
|
|
LlmTransportLimits {
|
|
connect_timeout: Duration::from_secs(1),
|
|
request_timeout: Duration::from_secs(2),
|
|
read_idle_timeout: Duration::from_secs(1),
|
|
pool_idle_timeout: Duration::from_secs(1),
|
|
total_timeout: Duration::from_secs(4),
|
|
max_prompt_bytes: 64 * 1024,
|
|
max_response_bytes: 64 * 1024,
|
|
max_concurrent_requests: 2,
|
|
max_retries: 0,
|
|
max_retry_delay: Duration::from_millis(20),
|
|
}
|
|
}
|
|
|
|
fn client(url: &str, limits: LlmTransportLimits) -> Arc<LlmClient> {
|
|
let connection = AgentConfig::offline(url, "super-secret-api-key")
|
|
.expect("valid fake endpoint")
|
|
.llm;
|
|
Arc::new(LlmClient::new(connection, limits).expect("valid client"))
|
|
}
|
|
|
|
fn message() -> CompletionMessage {
|
|
CompletionMessage::text(MessageRole::Avatar, "hello").expect("bounded message")
|
|
}
|
|
|
|
fn system_message() -> CompletionMessage {
|
|
CompletionMessage::text(MessageRole::System, "use registered tools only")
|
|
.expect("bounded system message")
|
|
}
|
|
|
|
fn image_message() -> CompletionMessage {
|
|
CompletionMessage {
|
|
role: MessageRole::Avatar,
|
|
content: BoundedVec::try_from_vec(
|
|
"image.content",
|
|
vec![ContentPart::Image {
|
|
url: BoundedText::new("image.url", "https://assets.invalid/object.png")
|
|
.expect("image URL"),
|
|
detail: ImageDetail::Low,
|
|
}],
|
|
)
|
|
.expect("image content"),
|
|
tool_call_id: None,
|
|
proposed_calls: BoundedVec::new(),
|
|
}
|
|
}
|
|
|
|
fn tool(name: &str, mutating: bool) -> ToolDefinition {
|
|
let mut properties = BTreeMap::new();
|
|
properties.insert("target".into(), ToolSchema::String);
|
|
ToolDefinition {
|
|
name: BoundedText::new("tool.name", name).expect("tool name"),
|
|
description: BoundedText::new("tool.description", "test tool").expect("description"),
|
|
schema: ToolSchema::Object {
|
|
properties,
|
|
required: BTreeSet::from(["target".into()]),
|
|
additional_properties: false,
|
|
},
|
|
mutating,
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn exact_url_auth_schema_fragmentation_and_metadata_are_verified() {
|
|
let body = serde_json::to_vec(&response_text("done")).expect("response");
|
|
let midpoint = body.len() / 2;
|
|
let mut server = fake_server(vec![ResponsePlan {
|
|
status: 200,
|
|
headers: vec![("content-type".into(), "application/json".into())],
|
|
chunks: vec![body[..midpoint].to_vec(), body[midpoint..].to_vec()],
|
|
delay: Duration::ZERO,
|
|
}])
|
|
.await;
|
|
let client = client(&server.url, limits());
|
|
let source = CancellationTokenSource::new();
|
|
let messages = [system_message(), message(), image_message()];
|
|
let completion = client
|
|
.complete(&messages, &[tool("look", false)], &source.token())
|
|
.await
|
|
.expect("completion");
|
|
assert_eq!(completion.request_id, 1);
|
|
assert_eq!(completion.attempts, 1);
|
|
assert_eq!(completion.usage.expect("usage").total_tokens, Some(5));
|
|
let request = server.requests.recv().await.expect("captured request");
|
|
assert!(
|
|
request
|
|
.head
|
|
.starts_with("POST /exact/chat?route=operator HTTP/1.1")
|
|
);
|
|
assert!(
|
|
request
|
|
.head
|
|
.to_ascii_lowercase()
|
|
.contains("authorization: bearer super-secret-api-key")
|
|
);
|
|
assert!(
|
|
request
|
|
.head
|
|
.to_ascii_lowercase()
|
|
.contains("x-correlation-id: agent-request-0000000000000001")
|
|
);
|
|
assert!(request.body.get("model").is_none());
|
|
assert!(request.body.get("provider").is_none());
|
|
assert_eq!(request.body["tools"][0]["type"], "function");
|
|
assert_eq!(request.body["messages"][0]["role"], "system");
|
|
assert_eq!(
|
|
request.body["messages"][2]["content"][0]["type"],
|
|
"image_url"
|
|
);
|
|
assert_eq!(
|
|
request.body["messages"][2]["content"][0]["image_url"]["detail"],
|
|
"low"
|
|
);
|
|
let diagnostic = format!("{client:?}");
|
|
assert!(!diagnostic.contains("super-secret-api-key"));
|
|
assert!(!diagnostic.contains("route=operator"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn prompt_bytes_and_schema_complexity_fail_before_network_access() {
|
|
let mut tiny = limits();
|
|
tiny.max_prompt_bytes = 32;
|
|
let client = client("http://127.0.0.1:9/exact", tiny);
|
|
assert_eq!(
|
|
client
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect_err("prompt cap"),
|
|
LlmError::PromptTooLarge
|
|
);
|
|
|
|
let mut schema = ToolSchema::String;
|
|
for _ in 0..18 {
|
|
schema = ToolSchema::Array {
|
|
items: Box::new(schema),
|
|
max_items: 1,
|
|
};
|
|
}
|
|
let definition = ToolDefinition {
|
|
name: BoundedText::new("tool.name", "deep").expect("name"),
|
|
description: BoundedText::new("tool.description", "too deep").expect("description"),
|
|
schema,
|
|
mutating: false,
|
|
};
|
|
assert_eq!(definition.validate(), Err(LlmError::InvalidToolSchema));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn malformed_html_oversize_and_duplicate_calls_are_typed() {
|
|
let cases = [
|
|
(
|
|
ResponsePlan::json(json!({"not_choices":[]})),
|
|
LlmError::MalformedJson,
|
|
),
|
|
(
|
|
ResponsePlan {
|
|
status: 200,
|
|
headers: vec![("content-type".into(), "text/html".into())],
|
|
chunks: vec![b"<html>not JSON</html>".to_vec()],
|
|
delay: Duration::ZERO,
|
|
},
|
|
LlmError::MalformedJson,
|
|
),
|
|
(ResponsePlan::status(500), LlmError::HttpStatus(500)),
|
|
(
|
|
ResponsePlan::json(response_calls(json!([{
|
|
"id":"unsupported-1",
|
|
"type":"not-a-function",
|
|
"function":{"name":"look","arguments":"{}"}
|
|
}]))),
|
|
LlmError::UnsupportedResponse,
|
|
),
|
|
(
|
|
ResponsePlan::json(response_calls(json!([
|
|
call("same", "look", "{}"),
|
|
call("same", "look", "{}")
|
|
]))),
|
|
LlmError::DuplicateToolCallId,
|
|
),
|
|
];
|
|
for (plan, expected) in cases {
|
|
let server = fake_server(vec![plan]).await;
|
|
let error = client(&server.url, limits())
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect_err("typed failure");
|
|
assert_eq!(error, expected);
|
|
}
|
|
|
|
let mut small = limits();
|
|
small.max_response_bytes = 32;
|
|
let server = fake_server(vec![ResponsePlan::json(response_text(
|
|
"this response is deliberately larger than thirty-two bytes",
|
|
))])
|
|
.await;
|
|
assert_eq!(
|
|
client(&server.url, small)
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect_err("oversize rejected"),
|
|
LlmError::ResponseTooLarge
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn redirects_never_forward_bearer_credentials() {
|
|
let mut target = fake_server(vec![ResponsePlan::json(response_text("leaked"))]).await;
|
|
let redirect = ResponsePlan {
|
|
status: 302,
|
|
headers: vec![("location".into(), target.url.clone())],
|
|
chunks: Vec::new(),
|
|
delay: Duration::ZERO,
|
|
};
|
|
let mut source = fake_server(vec![redirect]).await;
|
|
let error = client(&source.url, limits())
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect_err("redirect refused");
|
|
assert_eq!(error, LlmError::RedirectRefused);
|
|
source.requests.recv().await.expect("source request");
|
|
assert!(
|
|
tokio::time::timeout(Duration::from_millis(100), target.requests.recv())
|
|
.await
|
|
.is_err(),
|
|
"redirect target must receive no request"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn retry_classification_retry_after_timeout_and_cancellation_are_bounded() {
|
|
let retry = ResponsePlan {
|
|
status: 503,
|
|
headers: vec![("retry-after".into(), "0".into())],
|
|
chunks: Vec::new(),
|
|
delay: Duration::ZERO,
|
|
};
|
|
let mut server = fake_server(vec![retry, ResponsePlan::json(response_text("recovered"))]).await;
|
|
let mut retry_limits = limits();
|
|
retry_limits.max_retries = 1;
|
|
let completion = client(&server.url, retry_limits)
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect("transient retry");
|
|
assert_eq!(completion.attempts, 2);
|
|
server.requests.recv().await.expect("first attempt");
|
|
server.requests.recv().await.expect("second attempt");
|
|
|
|
let server = fake_server(vec![
|
|
ResponsePlan::status(400),
|
|
ResponsePlan::json(response_text("must not retry")),
|
|
])
|
|
.await;
|
|
let mut retry_limits = limits();
|
|
retry_limits.max_retries = 1;
|
|
assert_eq!(
|
|
client(&server.url, retry_limits)
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect_err("400 not retried"),
|
|
LlmError::HttpStatus(400)
|
|
);
|
|
|
|
let mut slow_limits = limits();
|
|
slow_limits.request_timeout = Duration::from_millis(50);
|
|
slow_limits.total_timeout = Duration::from_millis(100);
|
|
let server = fake_server(vec![
|
|
ResponsePlan::json(response_text("late")).delayed(Duration::from_secs(1)),
|
|
])
|
|
.await;
|
|
assert_eq!(
|
|
client(&server.url, slow_limits)
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
.expect_err("timeout"),
|
|
LlmError::Timeout
|
|
);
|
|
|
|
let server = fake_server(vec![
|
|
ResponsePlan::json(response_text("late")).delayed(Duration::from_secs(1)),
|
|
])
|
|
.await;
|
|
let client = client(&server.url, limits());
|
|
let source = CancellationTokenSource::new();
|
|
let token = source.token();
|
|
let messages = [message()];
|
|
let request = client.complete(&messages, &[], &token);
|
|
tokio::pin!(request);
|
|
tokio::select! {
|
|
result = &mut request => panic!("request unexpectedly completed: {result:?}"),
|
|
() = tokio::time::sleep(Duration::from_millis(20)) => source.cancel(),
|
|
}
|
|
assert_eq!(request.await.expect_err("cancelled"), LlmError::Cancelled);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn concurrent_slow_sessions_respect_the_shared_semaphore() {
|
|
let plans = (0..6)
|
|
.map(|_| ResponsePlan::json(response_text("done")).delayed(Duration::from_millis(80)))
|
|
.collect();
|
|
let server = fake_server(plans).await;
|
|
let client = client(&server.url, limits());
|
|
let mut tasks = Vec::with_capacity(6);
|
|
for _ in 0..6 {
|
|
let client = Arc::clone(&client);
|
|
tasks.push(tokio::spawn(async move {
|
|
client
|
|
.complete(&[message()], &[], &CancellationToken::default())
|
|
.await
|
|
}));
|
|
}
|
|
for task in tasks {
|
|
task.await
|
|
.expect("session task")
|
|
.expect("session completion");
|
|
}
|
|
assert_eq!(server.max_active.load(Ordering::Acquire), 2);
|
|
}
|
|
|
|
struct RecordingExecutor {
|
|
calls: AtomicUsize,
|
|
result: ToolExecution,
|
|
}
|
|
|
|
impl ToolExecutor for RecordingExecutor {
|
|
fn execute<'a>(
|
|
&'a self,
|
|
_definition: &'a ToolDefinition,
|
|
_call: &'a metacrate_grid_agent::ProposedToolCall,
|
|
_arguments: &'a Value,
|
|
_cancellation: &'a CancellationToken,
|
|
) -> ToolFuture<'a> {
|
|
self.calls.fetch_add(1, Ordering::AcqRel);
|
|
let result = self.result.clone();
|
|
Box::pin(async move { result })
|
|
}
|
|
}
|
|
|
|
fn loop_limits() -> ToolLoopLimits {
|
|
ToolLoopLimits {
|
|
max_turns: 3,
|
|
max_tool_calls_per_turn: 4,
|
|
max_tool_calls_per_session: 4,
|
|
max_history_messages: 8,
|
|
max_history_bytes: 4096,
|
|
wall_clock_timeout: Duration::from_secs(2),
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn valid_tool_round_trip_and_model_authored_summary_complete() {
|
|
let mut server = fake_server(vec![
|
|
ResponsePlan::json(response_calls(json!([call(
|
|
"call-1",
|
|
"look",
|
|
r#"{"target":"tree"}"#,
|
|
)]))),
|
|
ResponsePlan::json(response_text("I looked at the tree.")),
|
|
])
|
|
.await;
|
|
let executor = RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::Completed(
|
|
BoundedText::new("result", "tree is nearby").expect("result"),
|
|
),
|
|
};
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
let generation = SessionGeneration::default();
|
|
let outcome = loop_
|
|
.run(
|
|
vec![message()],
|
|
&generation,
|
|
generation.current(),
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect("tool loop");
|
|
assert_eq!(outcome.turns, 2);
|
|
assert_eq!(outcome.tool_calls, 1);
|
|
assert_eq!(executor.calls.load(Ordering::Acquire), 1);
|
|
server.requests.recv().await.expect("tool request");
|
|
let round_trip = server.requests.recv().await.expect("observation request");
|
|
let serialized = round_trip.body.to_string();
|
|
assert!(serialized.contains("tree is nearby"));
|
|
assert!(serialized.contains("call-1"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn unknown_and_malformed_calls_become_observations_without_execution() {
|
|
let mut server = fake_server(vec![
|
|
ResponsePlan::json(response_calls(json!([
|
|
call("unknown-1", "missing", "{}"),
|
|
call("bad-1", "look", "not-json"),
|
|
call("schema-1", "look", r#"{"wrong":true}"#)
|
|
]))),
|
|
ResponsePlan::json(response_text("handled safely")),
|
|
])
|
|
.await;
|
|
let executor = RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::Completed(BoundedText::new("result", "unused").expect("result")),
|
|
};
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
let generation = SessionGeneration::default();
|
|
loop_
|
|
.run(
|
|
vec![message()],
|
|
&generation,
|
|
0,
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect("safe observations");
|
|
assert_eq!(executor.calls.load(Ordering::Acquire), 0);
|
|
server.requests.recv().await.expect("first request");
|
|
let observations = server
|
|
.requests
|
|
.recv()
|
|
.await
|
|
.expect("second request")
|
|
.body
|
|
.to_string();
|
|
assert!(observations.contains("unknown tool"));
|
|
assert!(observations.contains("malformed JSON"));
|
|
assert!(observations.contains("registered schema"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn endless_loops_and_ambiguous_mutations_fail_without_reexecution() {
|
|
let server = fake_server(
|
|
(1..=3)
|
|
.map(|number| {
|
|
ResponsePlan::json(response_calls(json!([call(
|
|
&format!("call-{number}"),
|
|
"look",
|
|
r#"{"target":"tree"}"#,
|
|
)])))
|
|
})
|
|
.collect(),
|
|
)
|
|
.await;
|
|
let executor = RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::Completed(BoundedText::new("result", "again").expect("result")),
|
|
};
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
assert_eq!(
|
|
loop_
|
|
.run(
|
|
vec![message()],
|
|
&SessionGeneration::default(),
|
|
0,
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect_err("endless loop"),
|
|
ToolLoopError::EndlessToolLoop
|
|
);
|
|
|
|
let server = fake_server(vec![ResponsePlan::json(response_calls(json!([call(
|
|
"mutate-1",
|
|
"rez",
|
|
r#"{"target":"cube"}"#,
|
|
)])))])
|
|
.await;
|
|
let executor = RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::AmbiguousMutation,
|
|
};
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("rez", true)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
assert_eq!(
|
|
loop_
|
|
.run(
|
|
vec![message()],
|
|
&SessionGeneration::default(),
|
|
0,
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect_err("ambiguous mutation"),
|
|
ToolLoopError::AmbiguousMutation
|
|
);
|
|
assert_eq!(executor.calls.load(Ordering::Acquire), 1);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn repeated_call_id_across_turns_is_rejected_before_reexecution() {
|
|
let duplicate = ResponsePlan::json(response_calls(json!([call(
|
|
"same-call",
|
|
"look",
|
|
r#"{"target":"tree"}"#,
|
|
)])));
|
|
let server = fake_server(vec![duplicate.clone(), duplicate]).await;
|
|
let executor = RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::Completed(BoundedText::new("result", "first").expect("result")),
|
|
};
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
assert_eq!(
|
|
loop_
|
|
.run(
|
|
vec![message()],
|
|
&SessionGeneration::default(),
|
|
0,
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect_err("duplicate call ID"),
|
|
ToolLoopError::DuplicateToolCallId
|
|
);
|
|
assert_eq!(executor.calls.load(Ordering::Acquire), 1);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn session_call_and_wall_clock_budgets_are_enforced() {
|
|
let server = fake_server(vec![ResponsePlan::json(response_calls(json!([
|
|
call("one", "look", r#"{"target":"tree"}"#),
|
|
call("two", "look", r#"{"target":"rock"}"#)
|
|
])))])
|
|
.await;
|
|
let executor = RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::Completed(BoundedText::new("result", "unused").expect("result")),
|
|
};
|
|
let mut one_call = loop_limits();
|
|
one_call.max_tool_calls_per_turn = 1;
|
|
one_call.max_tool_calls_per_session = 1;
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
one_call,
|
|
)
|
|
.expect("loop");
|
|
assert_eq!(
|
|
loop_
|
|
.run(
|
|
vec![message()],
|
|
&SessionGeneration::default(),
|
|
0,
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect_err("call budget"),
|
|
ToolLoopError::ToolCallLimit
|
|
);
|
|
assert_eq!(executor.calls.load(Ordering::Acquire), 0);
|
|
|
|
let server = fake_server(vec![
|
|
ResponsePlan::json(response_text("too late")).delayed(Duration::from_secs(1)),
|
|
])
|
|
.await;
|
|
let mut short_loop = loop_limits();
|
|
short_loop.wall_clock_timeout = Duration::from_millis(30);
|
|
let loop_ = ToolLoop::new(client(&server.url, limits()), vec![], short_loop).expect("loop");
|
|
assert_eq!(
|
|
loop_
|
|
.run(
|
|
vec![message()],
|
|
&SessionGeneration::default(),
|
|
0,
|
|
&CancellationToken::default(),
|
|
&executor,
|
|
)
|
|
.await
|
|
.expect_err("wall-clock budget"),
|
|
ToolLoopError::WallClockTimeout
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn superseded_session_discards_late_model_result_before_execution() {
|
|
let mut server = fake_server(vec![
|
|
ResponsePlan::json(response_calls(json!([call(
|
|
"late-1",
|
|
"look",
|
|
r#"{"target":"tree"}"#,
|
|
)])))
|
|
.delayed(Duration::from_millis(100)),
|
|
])
|
|
.await;
|
|
let executor = Arc::new(RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::Completed(BoundedText::new("result", "late").expect("result")),
|
|
});
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
let generation = Arc::new(SessionGeneration::default());
|
|
let token = CancellationToken::default();
|
|
let run = loop_.run(vec![message()], &generation, 0, &token, executor.as_ref());
|
|
tokio::pin!(run);
|
|
tokio::select! {
|
|
request = server.requests.recv() => {
|
|
request.expect("in-flight request");
|
|
}
|
|
result = &mut run => panic!("session completed before supersession: {result:?}"),
|
|
}
|
|
generation.supersede();
|
|
assert_eq!(
|
|
run.await.expect_err("superseded"),
|
|
ToolLoopError::Superseded
|
|
);
|
|
assert_eq!(executor.calls.load(Ordering::Acquire), 0);
|
|
}
|
|
|
|
struct BlockingExecutor {
|
|
started: Notify,
|
|
completed: AtomicUsize,
|
|
}
|
|
|
|
impl ToolExecutor for BlockingExecutor {
|
|
fn execute<'a>(
|
|
&'a self,
|
|
_definition: &'a ToolDefinition,
|
|
_call: &'a metacrate_grid_agent::ProposedToolCall,
|
|
_arguments: &'a Value,
|
|
_cancellation: &'a CancellationToken,
|
|
) -> ToolFuture<'a> {
|
|
Box::pin(async move {
|
|
self.started.notify_one();
|
|
std::future::pending::<()>().await;
|
|
self.completed.fetch_add(1, Ordering::AcqRel);
|
|
ToolExecution::AmbiguousMutation
|
|
})
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn superseding_session_cancels_an_in_flight_tool_executor() {
|
|
let server = fake_server(vec![ResponsePlan::json(response_calls(json!([call(
|
|
"in-flight-1",
|
|
"look",
|
|
r#"{"target":"tree"}"#,
|
|
)])))])
|
|
.await;
|
|
let executor = Arc::new(BlockingExecutor {
|
|
started: Notify::new(),
|
|
completed: AtomicUsize::new(0),
|
|
});
|
|
let loop_ = ToolLoop::new(
|
|
client(&server.url, limits()),
|
|
vec![tool("look", false)],
|
|
loop_limits(),
|
|
)
|
|
.expect("loop");
|
|
let generation = SessionGeneration::default();
|
|
let token = CancellationToken::default();
|
|
let run = loop_.run(vec![message()], &generation, 0, &token, executor.as_ref());
|
|
tokio::pin!(run);
|
|
tokio::select! {
|
|
() = executor.started.notified() => {}
|
|
result = &mut run => panic!("tool loop completed before supersession: {result:?}"),
|
|
}
|
|
generation.supersede();
|
|
assert_eq!(
|
|
run.await.expect_err("superseded executor"),
|
|
ToolLoopError::Superseded
|
|
);
|
|
assert_eq!(executor.completed.load(Ordering::Acquire), 0);
|
|
}
|
|
|
|
struct FailingSummarizer;
|
|
|
|
impl HistorySummarizer for FailingSummarizer {
|
|
fn summarize(&self, _omitted: &[CompletionMessage]) -> Option<CompletionMessage> {
|
|
None
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn failed_history_summary_degrades_to_bounded_safe_truncation() {
|
|
let mut server = fake_server(vec![ResponsePlan::json(response_text("done"))]).await;
|
|
let mut compact = loop_limits();
|
|
compact.max_history_messages = 2;
|
|
let loop_ = ToolLoop::new(client(&server.url, limits()), vec![], compact)
|
|
.expect("loop")
|
|
.with_summarizer(Arc::new(FailingSummarizer));
|
|
let history = vec![
|
|
CompletionMessage::text(MessageRole::System, "old system").expect("message"),
|
|
CompletionMessage::text(MessageRole::Avatar, "old question").expect("message"),
|
|
CompletionMessage::text(MessageRole::Agent, "old answer").expect("message"),
|
|
message(),
|
|
];
|
|
loop_
|
|
.run(
|
|
history,
|
|
&SessionGeneration::default(),
|
|
0,
|
|
&CancellationToken::default(),
|
|
&RecordingExecutor {
|
|
calls: AtomicUsize::new(0),
|
|
result: ToolExecution::AmbiguousMutation,
|
|
},
|
|
)
|
|
.await
|
|
.expect("compacted completion");
|
|
let request = server
|
|
.requests
|
|
.recv()
|
|
.await
|
|
.expect("request")
|
|
.body
|
|
.to_string();
|
|
assert!(request.contains("history safely truncated"));
|
|
assert!(!request.contains("old question"));
|
|
}
|