Files
MetaCrate/crates/metacrate-grid-agent/tests/llm_transport.rs
Chili Palmer 11a8c37f21
Some checks failed
CI / rust-skia (Rust only) (push) Successful in 2m48s
CI / required (push) Failing after 54s
Implement pure Rust visual snapshots (#132)
2026-08-18 11:57:50 +02:00

990 lines
32 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 multimodal_capability_rejection_is_typed_and_never_retried() {
let mut server = fake_server(vec![ResponsePlan::status(415)]).await;
let mut retry_limits = limits();
retry_limits.max_retries = 2;
let image = CompletionMessage {
role: MessageRole::Avatar,
content: metacrate_grid_agent::BoundedVec::try_from_vec(
"message.content",
vec![ContentPart::Image {
url: metacrate_grid_agent::BoundedText::new(
"image.url",
"data:image/png;base64,iVBORw0KGgo=",
)
.unwrap(),
detail: ImageDetail::Low,
}],
)
.unwrap(),
tool_call_id: None,
proposed_calls: metacrate_grid_agent::BoundedVec::new(),
};
assert_eq!(
client(&server.url, retry_limits)
.complete(&[image], &[], &CancellationToken::default())
.await
.expect_err("multimodal rejection"),
LlmError::MultimodalUnsupported
);
server.requests.recv().await.expect("one image request");
assert!(
!matches!(
tokio::time::timeout(Duration::from_millis(100), server.requests.recv()).await,
Ok(Some(_))
),
"capability rejection must not resend the large image"
);
}
#[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"));
}