feat(grid-agent): add generic LLM tool loop (#119)
This commit is contained in:
950
crates/metacrate-grid-agent/tests/llm_transport.rs
Normal file
950
crates/metacrate-grid-agent/tests/llm_transport.rs
Normal file
@@ -0,0 +1,950 @@
|
||||
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"));
|
||||
}
|
||||
Reference in New Issue
Block a user