feat(grid-agent): isolate conversation memory (#122)
This commit is contained in:
149
crates/metacrate-grid-agent/tests/conversation_memory.rs
Normal file
149
crates/metacrate-grid-agent/tests/conversation_memory.rs
Normal file
@@ -0,0 +1,149 @@
|
||||
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");
|
||||
}
|
||||
Reference in New Issue
Block a user