#![allow(clippy::unwrap_used)]
use std::marker::PhantomData;
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use futures_util::FutureExt;
use github_copilot_sdk::handler::{
McpAuthHandler, McpAuthRequest, McpAuthResult, PermissionHandler, PermissionResult,
};
use github_copilot_sdk::session::PreparedSession;
use github_copilot_sdk::subscription::{EventSubscription, RecvErrorKind};
use github_copilot_sdk::types::{
CloudSessionOptions, CloudSessionRepository, MessageOptions, PermissionRequestData, RequestId,
ResumeSessionConfig, SessionConfig, SessionId,
};
use github_copilot_sdk::{Client, ErrorKind, SessionErrorKind};
use serde_json::{Value, json};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, duplex};
use tokio::time::timeout;
const TIMEOUT: Duration = Duration::from_secs(5);
const QUIET: Duration = Duration::from_millis(150);
const BURST: usize = 600;
async fn write_framed(writer: &mut (impl AsyncWrite + Unpin), body: &[u8]) {
let header = format!("Content-Length: {}\r\n\r\n", body.len());
writer.write_all(header.as_bytes()).await.unwrap();
writer.write_all(body).await.unwrap();
writer.flush().await.unwrap();
}
async fn read_framed(reader: &mut (impl AsyncRead + Unpin)) -> Value {
let mut header = String::new();
loop {
let mut byte = [0u8; 1];
tokio::io::AsyncReadExt::read_exact(reader, &mut byte)
.await
.unwrap();
header.push(byte[0] as char);
if header.ends_with("\r\n\r\n") {
break;
}
}
let length: usize = header
.trim()
.strip_prefix("Content-Length: ")
.unwrap()
.parse()
.unwrap();
let mut buf = vec![0u8; length];
tokio::io::AsyncReadExt::read_exact(reader, &mut buf)
.await
.unwrap();
serde_json::from_slice(&buf).unwrap()
}
struct FakeServer {
read: tokio::io::DuplexStream,
write: tokio::io::DuplexStream,
}
impl FakeServer {
async fn read_request(&mut self) -> Value {
timeout(TIMEOUT, read_framed(&mut self.read)).await.unwrap()
}
async fn expect_quiet(&mut self) {
assert!(
timeout(QUIET, read_framed(&mut self.read)).await.is_err(),
"expected no wire traffic"
);
}
async fn respond(&mut self, request: &Value, result: Value) {
let id = request["id"].as_u64().unwrap();
let response = json!({ "jsonrpc": "2.0", "id": id, "result": result });
write_framed(&mut self.write, &serde_json::to_vec(&response).unwrap()).await;
}
async fn respond_error(&mut self, request: &Value, code: i64, message: &str) {
self.respond_error_with_data(request, code, message, None)
.await;
}
async fn respond_error_with_data(
&mut self,
request: &Value,
code: i64,
message: &str,
data: Option<Value>,
) {
let id = request["id"].as_u64().unwrap();
let mut response = json!({
"jsonrpc": "2.0",
"id": id,
"error": { "code": code, "message": message },
});
if let Some(data) = data {
response["error"]["data"] = data;
}
write_framed(&mut self.write, &serde_json::to_vec(&response).unwrap()).await;
}
async fn send_event(&mut self, session_id: &str, id: &str, event_type: &str, ephemeral: bool) {
let notification = json!({
"jsonrpc": "2.0",
"method": "session.event",
"params": {
"sessionId": session_id,
"event": {
"id": id,
"timestamp": "2025-01-01T00:00:00Z",
"ephemeral": ephemeral,
"type": event_type,
"data": {},
},
},
});
write_framed(&mut self.write, &serde_json::to_vec(¬ification).unwrap()).await;
}
async fn send_startup_burst(&mut self, session_id: &str) {
for i in 0..BURST {
self.send_event(
session_id,
&format!("evt-{i}"),
"assistant.message_delta",
false,
)
.await;
}
self.send_event(session_id, "evt-idle", "session.idle", true)
.await;
}
async fn answer_skills_reload(&mut self) {
let request = self.read_request().await;
assert_eq!(request["method"], "session.skills.reload");
self.respond(&request, json!({})).await;
}
async fn await_publication(&mut self, session_id: &str, published: &tokio::sync::Notify) {
let notification = json!({
"jsonrpc": "2.0",
"method": "session.event",
"params": {
"sessionId": session_id,
"event": {
"id": "publication-fence",
"timestamp": "2025-01-01T00:00:00Z",
"type": "permission.requested",
"data": { "requestId": "publication-fence", "kind": "read" },
},
},
});
write_framed(&mut self.write, &serde_json::to_vec(¬ification).unwrap()).await;
timeout(TIMEOUT, published.notified()).await.unwrap();
}
}
struct PublicationFence(Arc<tokio::sync::Notify>);
#[async_trait]
impl PermissionHandler for PublicationFence {
async fn handle(
&self,
_session_id: SessionId,
_request_id: RequestId,
_request: PermissionRequestData,
) -> PermissionResult {
self.0.notify_one();
PermissionResult::no_result()
}
}
struct CancelMcpAuthHandler;
#[async_trait]
impl McpAuthHandler for CancelMcpAuthHandler {
async fn handle(
&self,
_session_id: SessionId,
_request_id: RequestId,
_request: McpAuthRequest,
) -> McpAuthResult {
McpAuthResult::Cancelled
}
}
fn make_client() -> (Client, FakeServer) {
let (client_write, server_read) = duplex(1 << 20);
let (server_write, client_read) = duplex(1 << 20);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
(
client,
FakeServer {
read: server_read,
write: server_write,
},
)
}
fn cloud_options() -> CloudSessionOptions {
CloudSessionOptions::with_repository(CloudSessionRepository::new("octocat", "hello-world"))
}
#[tokio::test]
async fn client_call_preserves_structured_rpc_error_data() {
let (client, mut server) = make_client();
let call = tokio::spawn({
let client = client.clone();
async move { client.call("session.raw", None).await }
});
let request = server.read_request().await;
let data = json!({
"code": "managed_policy_blocked",
"setting": "extensions",
"message": "Extensions are disabled by policy",
});
server
.respond_error_with_data(
&request,
-32001,
"managed policy blocked",
Some(data.clone()),
)
.await;
let error = timeout(TIMEOUT, call).await.unwrap().unwrap().unwrap_err();
assert_eq!(error.kind(), &ErrorKind::Rpc { code: -32001 });
assert_eq!(error.rpc_code(), Some(-32001));
assert_eq!(error.message(), Some("managed policy blocked"));
assert_eq!(error.rpc_data(), Some(&data));
assert!(std::error::Error::source(&error).is_none());
assert_eq!(
error.to_string(),
"RPC error -32001: managed policy blocked"
);
for debug in [format!("{error:?}"), format!("{error:#?}")] {
assert!(debug.contains("context"));
assert!(debug.contains("Rpc"));
assert!(debug.contains("managed policy blocked"));
assert!(!debug.contains("rpc_data"));
assert!(
!debug.contains("managed_policy_blocked"),
"RPC payloads must be accessed explicitly, not included in Debug"
);
}
}
#[tokio::test]
async fn client_call_preserves_non_object_rpc_error_data() {
let (client, mut server) = make_client();
for data in [
json!([{"detail": "example"}, null, false, 42]),
json!("detail"),
json!(0),
json!(false),
] {
let call = tokio::spawn({
let client = client.clone();
async move { client.call("session.raw", None).await }
});
let request = server.read_request().await;
server
.respond_error_with_data(&request, -32001, "request failed", Some(data.clone()))
.await;
let error = timeout(TIMEOUT, call).await.unwrap().unwrap().unwrap_err();
assert_eq!(error.rpc_data(), Some(&data));
}
}
#[tokio::test]
async fn client_call_handles_rpc_error_without_data() {
let (client, mut server) = make_client();
for data in [None, Some(Value::Null)] {
let call = tokio::spawn({
let client = client.clone();
async move { client.call("session.raw", None).await }
});
let request = server.read_request().await;
server
.respond_error_with_data(&request, -32002, "request failed", data)
.await;
let error = timeout(TIMEOUT, call).await.unwrap().unwrap().unwrap_err();
assert_eq!(error.kind(), &ErrorKind::Rpc { code: -32002 });
assert_eq!(error.rpc_code(), Some(-32002));
assert_eq!(error.message(), Some("request failed"));
assert_eq!(error.rpc_data(), None);
assert!(std::error::Error::source(&error).is_none());
assert_eq!(error.to_string(), "RPC error -32002: request failed");
}
}
fn create_result(session_id: &str) -> Value {
json!({ "sessionId": session_id, "workspacePath": "/tmp/workspace" })
}
async fn expect_startup_burst(events: &mut EventSubscription) {
for i in 0..BURST {
let event = timeout(TIMEOUT, events.recv())
.await
.unwrap_or_else(|_| panic!("timed out waiting for event {i}"))
.unwrap_or_else(|error| panic!("event {i} not delivered: {error}"));
assert_eq!(event.id.as_str(), format!("evt-{i}"), "out-of-order event");
}
let idle = timeout(TIMEOUT, events.recv()).await.unwrap().unwrap();
assert_eq!(idle.id.as_str(), "evt-idle");
assert_eq!(idle.event_type, "session.idle");
assert_eq!(idle.ephemeral, Some(true));
}
async fn expect_event_id(events: &mut EventSubscription, id: &str) {
let event = timeout(TIMEOUT, events.recv())
.await
.unwrap_or_else(|_| panic!("timed out waiting for {id}"))
.unwrap_or_else(|error| panic!("expected {id}, got {error}"));
assert_eq!(event.id, id);
}
fn expect_error<T>(result: Result<T, github_copilot_sdk::Error>) -> github_copilot_sdk::Error {
match result {
Ok(_) => panic!("expected an error"),
Err(error) => error,
}
}
fn assert_only_registration(client: &Client, session_id: &SessionId, context: &str) {
let registered = client.registered_session_ids_for_test();
let matches_expected = registered.len() == 1 && registered[0] == *session_id;
assert!(matches_expected, "{context}");
}
async fn await_no_registrations(client: &Client) {
let deadline = tokio::time::Instant::now() + TIMEOUT;
loop {
let outstanding = client.registered_session_count_for_test();
if outstanding == 0 {
return;
}
assert!(
tokio::time::Instant::now() < deadline,
"{outstanding} session registration(s) were never cleaned up"
);
tokio::task::yield_now().await;
}
}
async fn expect_closed(events: &mut EventSubscription) {
loop {
match timeout(TIMEOUT, events.recv()).await.unwrap() {
Ok(_) => continue,
Err(error) => {
assert!(
matches!(error.kind(), RecvErrorKind::Closed),
"expected Closed, got {:?}",
error.kind()
);
return;
}
}
}
}
#[tokio::test]
async fn prepared_create_delivers_pre_response_burst_to_concurrent_drain() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-create-concurrent");
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_event_buffer_capacity(2048),
)
.unwrap();
let mut events = prepared.subscribe();
let drain = tokio::spawn(async move {
expect_startup_burst(&mut events).await;
});
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
server.send_startup_burst(session_id.as_str()).await;
server
.respond(&create_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
timeout(TIMEOUT, drain).await.unwrap().unwrap();
drop(session);
}
#[tokio::test]
async fn prepared_create_retains_burst_for_deferred_consumer() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-create-deferred");
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_event_buffer_capacity(2048),
)
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
server.send_startup_burst(session_id.as_str()).await;
server
.respond(&create_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
expect_startup_burst(&mut events).await;
drop(session);
}
#[tokio::test]
async fn prepared_resume_delivers_pre_response_burst() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-resume");
let prepared = client
.prepare_resume_session(
ResumeSessionConfig::new(session_id.clone())
.with_continue_pending_work(true)
.with_event_buffer_capacity(2048),
)
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let resume_req = server.read_request().await;
assert_eq!(resume_req["method"], "session.resume");
assert_eq!(resume_req["params"]["continuePendingWork"], true);
server.send_startup_burst(session_id.as_str()).await;
server
.respond(&resume_req, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
expect_startup_burst(&mut events).await;
let mut late = session.subscribe();
assert!(late.recv().now_or_never().is_none());
server
.send_event(session_id.as_str(), "live", "assistant.message", false)
.await;
for subscription in [&mut events, &mut late] {
expect_event_id(subscription, "live").await;
}
drop(session);
}
#[tokio::test]
async fn undersized_buffer_reports_lag_and_keeps_live_tail() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-lag");
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_event_buffer_capacity(8),
)
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
server.send_startup_burst(session_id.as_str()).await;
server
.respond(&create_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
let mut lagged = None;
let mut last_index: Option<usize> = None;
while lagged.is_none() {
match timeout(TIMEOUT, events.recv()).await.unwrap() {
Ok(event) => {
if let Some(index) = event.id.as_str().strip_prefix("evt-")
&& let Ok(index) = index.parse::<usize>()
{
if let Some(previous) = last_index {
assert!(index > previous, "delivered events must stay ordered");
}
last_index = Some(index);
}
}
Err(error) => match error.kind() {
RecvErrorKind::Lagged(lag) => lagged = Some(lag.skipped()),
other => panic!("expected lag, got {other:?}"),
},
}
}
assert!(lagged.unwrap() > 0, "lag must report the skipped count");
server
.send_event(session_id.as_str(), "evt-live", "assistant.message", false)
.await;
let live = loop {
match timeout(TIMEOUT, events.recv()).await.unwrap() {
Ok(event) if event.id.as_str() == "evt-live" => break event,
Ok(_) => continue,
Err(error) => match error.kind() {
RecvErrorKind::Lagged(_) => continue,
other => panic!("subscription ended before the live tail: {other:?}"),
},
}
};
assert_eq!(live.event_type, "assistant.message");
drop(session);
}
#[tokio::test]
async fn prepare_is_inert_until_start_is_polled() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-inert");
let tasks_before = tokio::runtime::Handle::current()
.metrics()
.num_alive_tasks();
let prepared = client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap();
let _events = prepared.subscribe();
server.expect_quiet().await;
assert!(client.registered_session_ids_for_test().is_empty());
assert_eq!(
tokio::runtime::Handle::current()
.metrics()
.num_alive_tasks(),
tasks_before,
"prepare must not spawn a task"
);
let start = prepared.start();
server.expect_quiet().await;
assert!(client.registered_session_ids_for_test().is_empty());
let start = tokio::spawn(start);
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
server
.respond(&create_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
drop(session);
}
#[tokio::test]
async fn dropping_unstarted_prepared_session_leaves_no_state() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-dropped");
let prepared = client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap();
let mut events = prepared.subscribe();
drop(prepared);
assert!(matches!(
timeout(TIMEOUT, events.recv())
.await
.unwrap()
.unwrap_err()
.kind(),
RecvErrorKind::Closed
));
assert!(client.registered_session_ids_for_test().is_empty());
server.expect_quiet().await;
}
#[tokio::test]
async fn cancelled_prepared_create_cleans_up_and_allows_retry() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-create-cancel");
let prepared = client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
start.abort();
let _ = start.await;
await_no_registrations(&client).await;
expect_closed(&mut events).await;
let retry = tokio::spawn(
client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap()
.start(),
);
let retry_req = server.read_request().await;
assert_eq!(retry_req["method"], "session.create");
server
.respond(&retry_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, retry).await.unwrap().unwrap().unwrap();
assert_eq!(session.id(), &session_id);
drop(session);
}
#[tokio::test]
async fn cancelled_prepared_resume_cleans_up_and_allows_retry() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-resume-cancel");
let prepared = client
.prepare_resume_session(ResumeSessionConfig::new(session_id.clone()))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let resume_req = server.read_request().await;
assert_eq!(resume_req["method"], "session.resume");
start.abort();
let _ = start.await;
await_no_registrations(&client).await;
expect_closed(&mut events).await;
let retry = tokio::spawn(
client
.prepare_resume_session(ResumeSessionConfig::new(session_id.clone()))
.unwrap()
.start(),
);
let retry_req = server.read_request().await;
assert_eq!(retry_req["method"], "session.resume");
server
.respond(&retry_req, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, retry).await.unwrap().unwrap().unwrap();
assert_eq!(session.id(), &session_id);
drop(session);
}
#[tokio::test]
async fn create_rpc_error_preserves_kind_and_cleans_up() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-create-rpc-error");
let prepared = client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
server
.respond_error(&create_req, -32000, "session create failed")
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
assert!(
matches!(error.kind(), ErrorKind::Rpc { code: -32000 }),
"unexpected error kind: {:?}",
error.kind()
);
await_no_registrations(&client).await;
expect_closed(&mut events).await;
}
#[tokio::test]
async fn create_result_parse_error_preserves_kind_and_cleans_up() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-create-parse-error");
let prepared = client
.prepare_session(SessionConfig::default().with_session_id(session_id))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
server
.respond(&create_req, json!({ "sessionId": 42 }))
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
assert!(matches!(error.kind(), ErrorKind::Json), "{error}");
await_no_registrations(&client).await;
expect_closed(&mut events).await;
}
#[tokio::test]
async fn create_session_id_mismatch_preserves_kind_and_cleans_up() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-mismatch");
let prepared = client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
server
.respond(&create_req, create_result("some-other-id"))
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
let ErrorKind::Session(SessionErrorKind::SessionIdMismatch {
requested,
returned,
}) = error.kind()
else {
panic!("unexpected error kind: {:?}", error.kind());
};
assert_eq!(requested, &session_id);
assert_eq!(returned.as_str(), "some-other-id");
await_no_registrations(&client).await;
expect_closed(&mut events).await;
}
#[tokio::test]
async fn resume_session_id_mismatch_preserves_kind_and_cleans_up() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-resume-mismatch");
let prepared = client
.prepare_resume_session(ResumeSessionConfig::new(session_id.clone()))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let resume_req = server.read_request().await;
server
.respond(&resume_req, json!({ "sessionId": "another-session" }))
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
assert!(
matches!(
error.kind(),
ErrorKind::Session(SessionErrorKind::SessionIdMismatch { .. })
),
"unexpected error kind: {:?}",
error.kind()
);
await_no_registrations(&client).await;
expect_closed(&mut events).await;
}
#[tokio::test]
async fn create_mcp_auth_interest_error_preserves_kind_and_cleans_up() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-create-interest-error");
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_mcp_auth_handler(Arc::new(CancelMcpAuthHandler)),
)
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
server
.respond(&create_req, create_result(session_id.as_str()))
.await;
let interest_req = server.read_request().await;
assert_eq!(interest_req["method"], "session.eventLog.registerInterest");
assert_eq!(interest_req["params"]["eventType"], "mcp.oauth_required");
server
.respond_error(&interest_req, -32003, "interest registration failed")
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
assert!(
matches!(error.kind(), ErrorKind::Rpc { code: -32003 }),
"unexpected error kind: {:?}",
error.kind()
);
expect_closed(&mut events).await;
await_no_registrations(&client).await;
}
#[tokio::test]
async fn resume_mcp_auth_interest_error_preserves_kind_and_cleans_up() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-resume-interest-error");
let prepared = client
.prepare_resume_session(
ResumeSessionConfig::new(session_id.clone())
.with_mcp_auth_handler(Arc::new(CancelMcpAuthHandler)),
)
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let resume_req = server.read_request().await;
assert_eq!(resume_req["method"], "session.resume");
server
.respond(&resume_req, json!({ "sessionId": session_id.as_str() }))
.await;
let interest_req = server.read_request().await;
assert_eq!(interest_req["method"], "session.eventLog.registerInterest");
server
.respond_error(&interest_req, -32004, "interest registration failed")
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
assert!(
matches!(error.kind(), ErrorKind::Rpc { code: -32004 }),
"unexpected error kind: {:?}",
error.kind()
);
server.expect_quiet().await;
expect_closed(&mut events).await;
await_no_registrations(&client).await;
}
#[tokio::test]
async fn zero_event_buffer_capacity_is_invalid_config() {
let (client, _server) = make_client();
let error = expect_error(
client.prepare_session(SessionConfig::default().with_event_buffer_capacity(0)),
);
assert!(matches!(error.kind(), ErrorKind::InvalidConfig));
let error = expect_error(client.prepare_resume_session(
ResumeSessionConfig::new(SessionId::new("zero")).with_event_buffer_capacity(0),
));
assert!(matches!(error.kind(), ErrorKind::InvalidConfig));
let error = expect_error(
client
.create_session(SessionConfig::default().with_event_buffer_capacity(0))
.await,
);
assert!(matches!(error.kind(), ErrorKind::InvalidConfig));
}
#[tokio::test]
async fn early_and_late_subscribers_share_one_event_loop() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-two-subscribers");
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_event_buffer_capacity(2048),
)
.unwrap();
let mut early = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
server
.send_event(session_id.as_str(), "evt-early", "assistant.message", false)
.await;
server
.respond(&create_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
let mut late = session.subscribe();
server
.send_event(session_id.as_str(), "evt-late", "assistant.message", false)
.await;
assert_eq!(
timeout(TIMEOUT, early.recv())
.await
.unwrap()
.unwrap()
.id
.as_str(),
"evt-early"
);
assert_eq!(
timeout(TIMEOUT, early.recv())
.await
.unwrap()
.unwrap()
.id
.as_str(),
"evt-late"
);
assert_eq!(
timeout(TIMEOUT, late.recv())
.await
.unwrap()
.unwrap()
.id
.as_str(),
"evt-late"
);
assert!(
timeout(QUIET, late.recv()).await.is_err(),
"duplicate delivery implies more than one event loop"
);
assert!(
timeout(QUIET, early.recv()).await.is_err(),
"duplicate delivery implies more than one event loop"
);
drop(session);
}
#[tokio::test]
async fn create_session_wrapper_keeps_rpc_sequence() {
let (client, mut server) = make_client();
let start = tokio::spawn({
let client = client.clone();
async move { client.create_session(SessionConfig::default()).await }
});
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
let session_id = create_req["params"]["sessionId"]
.as_str()
.unwrap()
.to_string();
server
.respond(&create_req, create_result(&session_id))
.await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
assert_eq!(session.id().as_str(), session_id);
server.expect_quiet().await;
drop(session);
}
#[tokio::test]
async fn resume_session_wrapper_keeps_rpc_sequence() {
let (client, mut server) = make_client();
let session_id = SessionId::new("wrapper-resume");
let start = tokio::spawn({
let client = client.clone();
let session_id = session_id.clone();
async move {
client
.resume_session(ResumeSessionConfig::new(session_id))
.await
}
});
let resume_req = server.read_request().await;
assert_eq!(resume_req["method"], "session.resume");
assert_eq!(resume_req["params"]["sessionId"], session_id.as_str());
server
.respond(&resume_req, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
assert_eq!(session.id(), &session_id);
server.expect_quiet().await;
drop(session);
}
#[tokio::test]
async fn resume_bootstrap_retains_all_startup_phases_and_keeps_later_subscribers_live() {
let (client, mut server) = make_client();
let session_id = SessionId::new("resume-bootstrap-phases");
let published = Arc::new(tokio::sync::Notify::new());
let start = tokio::spawn({
let client = client.clone();
let session_id = session_id.clone();
let published = published.clone();
async move {
client
.resume_session(
ResumeSessionConfig::new(session_id)
.with_event_buffer_capacity(1)
.with_permission_handler(Arc::new(PublicationFence(published)))
.with_mcp_auth_handler(Arc::new(CancelMcpAuthHandler)),
)
.await
}
});
let resume = server.read_request().await;
assert_eq!(resume["method"], "session.resume");
server
.send_event(
session_id.as_str(),
"pre-durable",
"session.model_change",
false,
)
.await;
server
.send_event(session_id.as_str(), "pre-ephemeral", "session.idle", true)
.await;
server
.await_publication(session_id.as_str(), &published)
.await;
server
.respond(&resume, json!({ "sessionId": session_id.as_str() }))
.await;
let interest = server.read_request().await;
assert_eq!(interest["method"], "session.eventLog.registerInterest");
server.send_startup_burst(session_id.as_str()).await;
server
.await_publication(session_id.as_str(), &published)
.await;
server.respond(&interest, json!({})).await;
let reload = server.read_request().await;
assert_eq!(reload["method"], "session.skills.reload");
server
.send_event(session_id.as_str(), "during-reload", "session.idle", true)
.await;
server
.await_publication(session_id.as_str(), &published)
.await;
server.respond(&reload, json!({})).await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
server
.send_event(
session_id.as_str(),
"post-setup",
"assistant.message",
false,
)
.await;
server
.await_publication(session_id.as_str(), &published)
.await;
let mut first = session.subscribe();
let mut second = session.subscribe();
server
.send_event(session_id.as_str(), "during-catchup", "session.idle", true)
.await;
expect_event_id(&mut second, "during-catchup").await;
for (id, ephemeral) in [
("pre-durable", Some(false)),
("pre-ephemeral", Some(true)),
("publication-fence", None),
] {
let event = timeout(TIMEOUT, first.recv()).await.unwrap().unwrap();
assert_eq!(event.id, id);
assert_eq!(event.ephemeral, ephemeral);
}
expect_startup_burst(&mut first).await;
for id in [
"publication-fence",
"during-reload",
"publication-fence",
"post-setup",
"publication-fence",
"during-catchup",
] {
expect_event_id(&mut first, id).await;
}
assert!(first.recv().now_or_never().is_none());
assert!(second.recv().now_or_never().is_none());
server
.send_event(session_id.as_str(), "live", "assistant.message", false)
.await;
for events in [&mut first, &mut second] {
expect_event_id(events, "live").await;
assert!(events.recv().now_or_never().is_none());
}
timeout(TIMEOUT, session.stop_event_loop()).await.unwrap();
drop(session);
expect_closed(&mut first).await;
expect_closed(&mut second).await;
}
#[tokio::test]
async fn structured_output_completion_preserves_resume_bootstrap_for_first_observer() {
check_structured_output_preserves_bootstrap(false).await;
}
#[tokio::test]
async fn structured_output_cancellation_preserves_resume_bootstrap_for_first_observer() {
check_structured_output_preserves_bootstrap(true).await;
}
async fn check_structured_output_preserves_bootstrap(cancel: bool) {
let (client, mut server) = make_client();
let session_id = SessionId::new("resume-structured-output");
let published = Arc::new(tokio::sync::Notify::new());
let start = tokio::spawn({
let client = client.clone();
let session_id = session_id.clone();
let published = published.clone();
async move {
client
.resume_session(
ResumeSessionConfig::new(session_id)
.with_permission_handler(Arc::new(PublicationFence(published))),
)
.await
}
});
let resume = server.read_request().await;
server
.send_event(session_id.as_str(), "startup-idle", "session.idle", true)
.await;
server
.await_publication(session_id.as_str(), &published)
.await;
server
.respond(&resume, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = Arc::new(timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap());
let waiting = tokio::spawn({
let session = session.clone();
async move {
session
.send_and_wait(
MessageOptions::new("structured")
.with_response_schema(json!({ "type": "object" })),
)
.await
}
});
let request = server.read_request().await;
assert_eq!(request["method"], "session.send");
if cancel {
waiting.abort();
let error = timeout(TIMEOUT, waiting).await.unwrap().unwrap_err();
assert!(error.is_cancelled(), "expected cancellation, got {error:?}");
} else {
server
.respond(&request, json!({ "messageId": "structured-user" }))
.await;
let notification = json!({
"jsonrpc": "2.0", "method": "session.event",
"params": { "sessionId": session_id.as_str(), "event": {
"id": "structured-answer", "timestamp": "2025-01-01T00:00:00Z",
"type": "assistant.message", "data": {
"messageId": "structured-answer", "originatingMessageId": "structured-user",
"content": "{\"ok\":true}"
}
} }
});
write_framed(
&mut server.write,
&serde_json::to_vec(¬ification).unwrap(),
)
.await;
server
.send_event(session_id.as_str(), "structured-idle", "session.idle", true)
.await;
let result = timeout(TIMEOUT, waiting)
.await
.unwrap()
.unwrap()
.unwrap()
.unwrap();
assert_eq!(result.id, "structured-answer");
}
let mut observer = session.subscribe();
expect_event_id(&mut observer, "startup-idle").await;
expect_event_id(&mut observer, "publication-fence").await;
if !cancel {
expect_event_id(&mut observer, "structured-answer").await;
expect_event_id(&mut observer, "structured-idle").await;
}
assert!(observer.recv().now_or_never().is_none());
timeout(TIMEOUT, session.stop_event_loop()).await.unwrap();
drop(session);
expect_closed(&mut observer).await;
}
#[tokio::test]
async fn populated_resume_bootstrap_cleans_up_after_setup_failure() {
check_populated_resume_cleanup(false).await;
}
#[tokio::test]
async fn populated_resume_bootstrap_cleans_up_after_setup_cancellation() {
check_populated_resume_cleanup(true).await;
}
async fn check_populated_resume_cleanup(cancel: bool) {
let (client, mut server) = make_client();
let session_id = SessionId::new("resume-bootstrap-cleanup");
let published = Arc::new(tokio::sync::Notify::new());
let start = tokio::spawn({
let client = client.clone();
let session_id = session_id.clone();
let published = published.clone();
async move {
client
.resume_session(
ResumeSessionConfig::new(session_id)
.with_permission_handler(Arc::new(PublicationFence(published)))
.with_mcp_auth_handler(Arc::new(CancelMcpAuthHandler)),
)
.await
}
});
let resume = server.read_request().await;
assert_eq!(resume["method"], "session.resume");
server
.respond(&resume, json!({ "sessionId": session_id.as_str() }))
.await;
let interest = server.read_request().await;
assert_eq!(interest["method"], "session.eventLog.registerInterest");
server.send_startup_burst(session_id.as_str()).await;
server
.await_publication(session_id.as_str(), &published)
.await;
if cancel {
start.abort();
let error = timeout(TIMEOUT, start)
.await
.unwrap()
.err()
.expect("cancelled resume task unexpectedly completed");
assert!(error.is_cancelled(), "expected cancellation, got {error:?}");
} else {
server
.respond_error(&interest, -32004, "interest registration failed")
.await;
let error = expect_error(timeout(TIMEOUT, start).await.unwrap().unwrap());
assert!(matches!(error.kind(), ErrorKind::Rpc { code: -32004 }));
}
await_no_registrations(&client).await;
server.expect_quiet().await;
let retry = tokio::spawn({
let client = client.clone();
let session_id = session_id.clone();
async move {
client
.resume_session(ResumeSessionConfig::new(session_id))
.await
}
});
let request = server.read_request().await;
assert_eq!(request["method"], "session.resume");
server
.respond(&request, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, retry).await.unwrap().unwrap().unwrap();
let mut events = session.subscribe();
server
.send_event(session_id.as_str(), "after-retry", "session.idle", true)
.await;
expect_event_id(&mut events, "after-retry").await;
timeout(TIMEOUT, session.stop_event_loop()).await.unwrap();
drop(session);
expect_closed(&mut events).await;
await_no_registrations(&client).await;
}
#[tokio::test]
async fn stopping_resume_releases_unclaimed_bootstrap() {
check_stopping_resume_bootstrap(false).await;
}
#[tokio::test]
async fn claimed_bootstrap_drains_after_stopping_and_dropping_session() {
check_stopping_resume_bootstrap(true).await;
}
async fn check_stopping_resume_bootstrap(claim: bool) {
let (client, mut server) = make_client();
let session_id = SessionId::new("resume-bootstrap-stop");
let published = Arc::new(tokio::sync::Notify::new());
let start = tokio::spawn({
let client = client.clone();
let session_id = session_id.clone();
let published = published.clone();
async move {
let prepared = client
.prepare_resume_session(
ResumeSessionConfig::new(session_id)
.with_event_buffer_capacity(1)
.with_permission_handler(Arc::new(PublicationFence(published))),
)
.unwrap();
drop(prepared.subscribe());
prepared.start().await
}
});
let resume = server.read_request().await;
server.send_startup_burst(session_id.as_str()).await;
server
.await_publication(session_id.as_str(), &published)
.await;
server
.respond(&resume, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
let claimed = claim.then(|| session.subscribe());
timeout(TIMEOUT, session.stop_event_loop()).await.unwrap();
let mut events = claimed.unwrap_or_else(|| session.subscribe());
if claim {
drop(session);
expect_startup_burst(&mut events).await;
expect_event_id(&mut events, "publication-fence").await;
} else {
assert!(
events.recv().now_or_never().is_none(),
"unclaimed backlog was replayed"
);
drop(session);
}
assert!(matches!(
timeout(TIMEOUT, events.recv())
.await
.unwrap()
.unwrap_err()
.kind(),
RecvErrorKind::Closed
));
await_no_registrations(&client).await;
}
#[tokio::test]
async fn active_prepared_resume_subscription_remains_bounded() {
let (client, mut server) = make_client();
let session_id = SessionId::new("prepared-resume-bounded");
let published = Arc::new(tokio::sync::Notify::new());
let prepared = client
.prepare_resume_session(
ResumeSessionConfig::new(session_id.clone())
.with_event_buffer_capacity(1)
.with_permission_handler(Arc::new(PublicationFence(published.clone()))),
)
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let resume = server.read_request().await;
server.send_startup_burst(session_id.as_str()).await;
server
.await_publication(session_id.as_str(), &published)
.await;
server
.respond(&resume, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, start).await.unwrap().unwrap().unwrap();
let error = timeout(TIMEOUT, events.recv()).await.unwrap().unwrap_err();
let RecvErrorKind::Lagged(lag) = error.kind() else {
panic!("expected bounded prepared delivery to lag, got {error:?}");
};
assert_eq!(lag.skipped(), (BURST + 1) as u64);
expect_event_id(&mut events, "publication-fence").await;
let mut late = session.subscribe();
server
.send_event(session_id.as_str(), "live", "session.idle", true)
.await;
for subscription in [&mut events, &mut late] {
expect_event_id(subscription, "live").await;
}
timeout(TIMEOUT, session.stop_event_loop()).await.unwrap();
drop(session);
expect_closed(&mut events).await;
expect_closed(&mut late).await;
}
struct CloneProbe<T>(PhantomData<T>);
impl<T: Clone> CloneProbe<T> {
fn is_clone(&self) -> bool {
true
}
}
trait MaybeClone {
fn is_clone(&self) -> bool {
false
}
}
impl<T> MaybeClone for CloneProbe<T> {}
#[test]
fn prepared_session_is_send_static_and_not_clone() {
fn assert_send_static<T: Send + 'static>() {}
assert_send_static::<PreparedSession>();
assert!(CloneProbe::<String>(PhantomData).is_clone());
assert!(!CloneProbe::<PreparedSession>(PhantomData).is_clone());
}
const DRIVE: Duration = Duration::from_millis(50);
#[tokio::test]
async fn stale_create_guard_does_not_unregister_same_id_retry() {
let (client, mut server) = make_client();
let session_id = SessionId::new("stale-create-guard");
let mut first = Box::pin(
client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap()
.start(),
);
let _ = timeout(DRIVE, &mut first).await;
let first_req = server.read_request().await;
assert_eq!(first_req["method"], "session.create");
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_event_buffer_capacity(64),
)
.unwrap();
let mut events = prepared.subscribe();
let mut second = Box::pin(prepared.start());
let _ = timeout(DRIVE, &mut second).await;
let second_req = server.read_request().await;
assert_eq!(second_req["method"], "session.create");
drop(first);
assert_only_registration(
&client,
&session_id,
"a stale startup guard unregistered the live retry",
);
server
.respond(&second_req, create_result(session_id.as_str()))
.await;
let session = timeout(TIMEOUT, &mut second).await.unwrap().unwrap();
server
.send_event(
session_id.as_str(),
"evt-after-stale",
"assistant.message",
false,
)
.await;
let event = timeout(TIMEOUT, events.recv()).await.unwrap().unwrap();
assert_eq!(event.id.as_str(), "evt-after-stale");
drop(session);
}
#[tokio::test]
async fn stale_resume_guard_does_not_unregister_same_id_retry() {
let (client, mut server) = make_client();
let session_id = SessionId::new("stale-resume-guard");
let mut first = Box::pin(
client
.prepare_resume_session(ResumeSessionConfig::new(session_id.clone()))
.unwrap()
.start(),
);
let _ = timeout(DRIVE, &mut first).await;
let first_req = server.read_request().await;
assert_eq!(first_req["method"], "session.resume");
let prepared = client
.prepare_resume_session(
ResumeSessionConfig::new(session_id.clone()).with_event_buffer_capacity(64),
)
.unwrap();
let mut events = prepared.subscribe();
let mut second = Box::pin(prepared.start());
let _ = timeout(DRIVE, &mut second).await;
let second_req = server.read_request().await;
assert_eq!(second_req["method"], "session.resume");
drop(first);
assert_only_registration(
&client,
&session_id,
"a stale startup guard unregistered the live retry",
);
let second = tokio::spawn(second);
server
.respond(&second_req, json!({ "sessionId": session_id.as_str() }))
.await;
server.answer_skills_reload().await;
let session = timeout(TIMEOUT, second).await.unwrap().unwrap().unwrap();
server
.send_event(
session_id.as_str(),
"evt-after-stale",
"assistant.message",
false,
)
.await;
let event = timeout(TIMEOUT, events.recv()).await.unwrap().unwrap();
assert_eq!(event.id.as_str(), "evt-after-stale");
drop(session);
}
#[tokio::test]
async fn dropping_superseded_session_does_not_unregister_its_replacement() {
let (client, mut server) = make_client();
let session_id = SessionId::new("superseded-session");
let first = tokio::spawn(
client
.prepare_session(SessionConfig::default().with_session_id(session_id.clone()))
.unwrap()
.start(),
);
let first_req = server.read_request().await;
server
.respond(&first_req, create_result(session_id.as_str()))
.await;
let first_session = timeout(TIMEOUT, first).await.unwrap().unwrap().unwrap();
let prepared = client
.prepare_session(
SessionConfig::default()
.with_session_id(session_id.clone())
.with_event_buffer_capacity(64),
)
.unwrap();
let mut events = prepared.subscribe();
let second = tokio::spawn(prepared.start());
let second_req = server.read_request().await;
server
.respond(&second_req, create_result(session_id.as_str()))
.await;
let second_session = timeout(TIMEOUT, second).await.unwrap().unwrap().unwrap();
drop(first_session);
assert_only_registration(
&client,
&session_id,
"dropping a superseded Session unregistered its replacement",
);
server
.send_event(
session_id.as_str(),
"evt-survivor",
"assistant.message",
false,
)
.await;
let event = timeout(TIMEOUT, events.recv()).await.unwrap().unwrap();
assert_eq!(event.id.as_str(), "evt-survivor");
drop(second_session);
}
#[tokio::test]
async fn cancelled_deferred_create_leaves_no_registration() {
let (client, mut server) = make_client();
let prepared = client
.prepare_session(SessionConfig::default().with_cloud(cloud_options()))
.unwrap();
let mut events = prepared.subscribe();
let start = tokio::spawn(prepared.start());
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
assert!(create_req["params"]["sessionId"].is_null());
start.abort();
let _ = start.await;
server
.respond(&create_req, create_result("server-assigned-id"))
.await;
expect_closed(&mut events).await;
await_no_registrations(&client).await;
server.expect_quiet().await;
assert!(
client.registered_session_ids_for_test().is_empty(),
"a cancelled deferred create left a registration behind"
);
let retry = tokio::spawn(
client
.prepare_session(SessionConfig::default().with_cloud(cloud_options()))
.unwrap()
.start(),
);
let retry_req = server.read_request().await;
server
.respond(&retry_req, create_result("server-assigned-retry"))
.await;
let session = timeout(TIMEOUT, retry).await.unwrap().unwrap().unwrap();
assert_eq!(session.id().as_str(), "server-assigned-retry");
drop(session);
}
async fn await_registered(client: &Client, session_id: &str) {
let deadline = tokio::time::Instant::now() + TIMEOUT;
while !client
.registered_session_ids_for_test()
.iter()
.any(|id| id.as_str() == session_id)
{
assert!(
tokio::time::Instant::now() < deadline,
"inline callback never registered the expected session"
);
tokio::task::yield_now().await;
}
}
#[tokio::test]
async fn deferred_create_cancelled_after_callback_registered_is_cleaned_up() {
let (client, mut server) = make_client();
let prepared = client
.prepare_session(SessionConfig::default().with_cloud(cloud_options()))
.unwrap();
let mut events = prepared.subscribe();
let mut start = Box::pin(prepared.start());
let _ = timeout(DRIVE, &mut start).await;
let create_req = server.read_request().await;
assert_eq!(create_req["method"], "session.create");
server
.respond(&create_req, create_result("registered-then-cancelled"))
.await;
await_registered(&client, "registered-then-cancelled").await;
drop(start);
expect_closed(&mut events).await;
await_no_registrations(&client).await;
server.expect_quiet().await;
let retry = tokio::spawn(
client
.prepare_session(SessionConfig::default().with_cloud(cloud_options()))
.unwrap()
.start(),
);
let retry_req = server.read_request().await;
server
.respond(&retry_req, create_result("registered-then-cancelled"))
.await;
let session = timeout(TIMEOUT, retry).await.unwrap().unwrap().unwrap();
assert_eq!(session.id().as_str(), "registered-then-cancelled");
assert_eq!(
client.registered_session_count_for_test(),
1,
"retry must hold exactly one registration"
);
drop(session);
}