150 lines
6.0 KiB
Rust
150 lines
6.0 KiB
Rust
use libremetaverse_types::UUID;
|
|
use libremetaverse_types::compat::CancellationTokenSource;
|
|
use metacrate_grid_agent::{
|
|
AgentConfig, ConversationChannel, ConversationKey, ConversationLimits, ConversationStore,
|
|
LlmClient, LlmTransportLimits, MemoryRecord,
|
|
};
|
|
use serde_json::{Value, json};
|
|
use std::sync::Arc;
|
|
use std::time::Duration;
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
use tokio::net::TcpListener;
|
|
use tokio::sync::mpsc;
|
|
|
|
const MAX_REQUEST_BYTES: usize = 1024 * 1024;
|
|
|
|
fn avatar(number: u64) -> UUID {
|
|
UUID::new_with_u_int64(number).expect("fixture UUID")
|
|
}
|
|
|
|
fn key(number: u64, channel: ConversationChannel) -> ConversationKey {
|
|
ConversationKey::new(avatar(number), channel).expect("nonzero fixture")
|
|
}
|
|
|
|
async fn capture_server(
|
|
request_count: usize,
|
|
) -> (String, mpsc::Receiver<Value>, tokio::task::JoinHandle<()>) {
|
|
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
|
|
.await
|
|
.expect("bind fake LLM");
|
|
let address = listener.local_addr().expect("listener address");
|
|
let (sender, receiver) = mpsc::channel(request_count);
|
|
let task = tokio::spawn(async move {
|
|
for _ in 0..request_count {
|
|
let (mut stream, _) = listener.accept().await.expect("LLM connection");
|
|
let request = read_request(&mut stream).await;
|
|
sender.send(request).await.expect("capture receiver");
|
|
let response = serde_json::to_vec(&json!({
|
|
"choices": [{"message": {"content": "ok", "tool_calls": []}}],
|
|
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
|
|
}))
|
|
.expect("response JSON");
|
|
let head = format!(
|
|
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
|
|
response.len()
|
|
);
|
|
stream
|
|
.write_all(head.as_bytes())
|
|
.await
|
|
.expect("response head");
|
|
stream.write_all(&response).await.expect("response body");
|
|
stream.shutdown().await.expect("response shutdown");
|
|
}
|
|
});
|
|
(format!("http://{address}/exact/chat"), receiver, task)
|
|
}
|
|
|
|
async fn read_request(stream: &mut tokio::net::TcpStream) -> Value {
|
|
let mut bytes = Vec::new();
|
|
let header_end = loop {
|
|
assert!(bytes.len() < MAX_REQUEST_BYTES, "bounded request headers");
|
|
let mut chunk = [0_u8; 2048];
|
|
let count = stream.read(&mut chunk).await.expect("request read");
|
|
assert!(count > 0, "complete request headers");
|
|
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 headers = std::str::from_utf8(&bytes[..header_end]).expect("UTF-8 headers");
|
|
let content_length = headers
|
|
.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 header");
|
|
assert!(content_length <= MAX_REQUEST_BYTES, "bounded request body");
|
|
while bytes.len() - header_end < content_length {
|
|
let mut chunk = [0_u8; 4096];
|
|
let count = stream.read(&mut chunk).await.expect("request body read");
|
|
assert!(count > 0, "complete request body");
|
|
bytes.extend_from_slice(&chunk[..count]);
|
|
}
|
|
serde_json::from_slice(&bytes[header_end..header_end + content_length]).expect("request JSON")
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn actual_llm_requests_never_cross_avatar_or_channel_boundaries() {
|
|
let (endpoint, mut requests, server) = capture_server(3).await;
|
|
let config = AgentConfig::offline(&endpoint, "test-key").expect("offline config");
|
|
let limits = 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(3),
|
|
max_prompt_bytes: 64 * 1024,
|
|
max_response_bytes: 64 * 1024,
|
|
max_concurrent_requests: 2,
|
|
max_retries: 0,
|
|
max_retry_delay: Duration::from_millis(10),
|
|
};
|
|
let client = Arc::new(LlmClient::new(config.llm, limits).expect("LLM client"));
|
|
let store = ConversationStore::open(ConversationLimits::default(), None).expect("store");
|
|
let alice_public = key(1, ConversationChannel::PublicChat);
|
|
let alice_direct = key(1, ConversationChannel::DirectIm);
|
|
let bob_public = key(2, ConversationChannel::PublicChat);
|
|
store
|
|
.append(
|
|
alice_public,
|
|
MemoryRecord::avatar_message("alice-public-only"),
|
|
)
|
|
.expect("alice public");
|
|
store
|
|
.append(
|
|
alice_direct,
|
|
MemoryRecord::avatar_message("alice-direct-only"),
|
|
)
|
|
.expect("alice direct");
|
|
store
|
|
.append(bob_public, MemoryRecord::avatar_message("bob-public-only"))
|
|
.expect("bob public");
|
|
|
|
for session_key in [alice_public, alice_direct, bob_public] {
|
|
let messages = store
|
|
.context(session_key)
|
|
.expect("isolated context")
|
|
.llm_messages()
|
|
.expect("LLM messages");
|
|
client
|
|
.complete(&messages, &[], &CancellationTokenSource::new().token())
|
|
.await
|
|
.expect("fake completion");
|
|
}
|
|
|
|
let expected = ["alice-public-only", "alice-direct-only", "bob-public-only"];
|
|
for own_text in expected {
|
|
let request = requests.recv().await.expect("captured request");
|
|
let serialized = serde_json::to_string(&request).expect("request string");
|
|
assert!(serialized.contains(own_text));
|
|
for other_text in expected {
|
|
if other_text != own_text {
|
|
assert!(!serialized.contains(other_text));
|
|
}
|
|
}
|
|
}
|
|
server.await.expect("capture server");
|
|
}
|