#![allow(clippy::unwrap_used)]
use github_copilot_sdk::Client;
use tokio::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt, duplex};
#[tokio::test]
async fn dropping_external_client_closes_its_streams() {
let (client_write, mut server_read) = duplex(8192);
let (_server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
client
.register_request_handler("host.shutdown", |_, client| async move {
client.call("ping", None).await
})
.unwrap();
drop(client);
let mut byte = [0];
let count = tokio::time::timeout(
std::time::Duration::from_secs(2),
server_read.read(&mut byte),
)
.await
.unwrap()
.unwrap();
assert_eq!(count, 0);
}
#[tokio::test]
async fn connection_handler_can_await_sdk_replies_before_responding() {
let (client_write, mut server_read) = duplex(8192);
let (mut server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
client
.register_request_handler("host.shutdown", |_, client| async move {
client.call("ping", None).await
})
.unwrap();
let request = serde_json::json!({"jsonrpc":"2.0","id":41,"method":"host.shutdown","params":{}});
write_framed(&mut server_write, &serde_json::to_vec(&request).unwrap()).await;
let ping = tokio::time::timeout(
std::time::Duration::from_secs(2),
read_framed(&mut server_read),
)
.await
.unwrap();
assert_eq!(ping["method"], "ping");
let response = serde_json::json!({"jsonrpc":"2.0","id":ping["id"],"result":{"drained":true}});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
let shutdown = tokio::time::timeout(
std::time::Duration::from_secs(2),
read_framed(&mut server_read),
)
.await
.unwrap();
assert_eq!(shutdown["id"], 41);
assert_eq!(shutdown["result"], serde_json::json!({"drained":true}));
drop(client);
let mut byte = [0];
assert_eq!(
tokio::time::timeout(
std::time::Duration::from_secs(2),
server_read.read(&mut byte),
)
.await
.unwrap()
.unwrap(),
0
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn synchronous_handler_setup_does_not_block_the_reader() {
let (client_write, mut server_read) = duplex(8192);
let (mut server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, ".".into()).unwrap();
let (started_tx, mut started) = tokio::sync::mpsc::unbounded_channel();
let (unblock, gate) = std::sync::mpsc::channel();
let gate = std::sync::Mutex::new(gate);
client
.register_request_handler("blocking.setup", move |_, _| {
started_tx.send(()).unwrap();
gate.lock()
.unwrap()
.recv_timeout(std::time::Duration::from_secs(10))
.unwrap();
async { Ok(serde_json::json!({})) }
})
.unwrap();
let request = serde_json::json!({"jsonrpc":"2.0","id":44,"method":"blocking.setup"});
write_framed(&mut server_write, &serde_json::to_vec(&request).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), started.recv())
.await
.unwrap()
.unwrap();
let ping = tokio::spawn({
let client = client.clone();
async move { client.call("ping", None).await }
});
let request = read_framed(&mut server_read).await;
let response = serde_json::json!({"jsonrpc":"2.0","id":request["id"],"result":{}});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), ping)
.await
.unwrap()
.unwrap()
.unwrap();
unblock.send(()).unwrap();
assert_eq!(read_framed(&mut server_read).await["id"], 44);
client.force_stop();
}
#[tokio::test]
async fn connection_handlers_reject_duplicate_registration_and_report_errors() {
let (client_write, mut server_read) = duplex(8192);
let (mut server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
client
.register_request_handler("host.shutdown", |_, _| async {
Err(github_copilot_sdk::Error::with_message(
github_copilot_sdk::ErrorKind::InvalidConfig,
"drain failed",
))
})
.unwrap();
assert!(
client
.register_request_handler("host.shutdown", |_, _| async { Ok(serde_json::json!({})) })
.is_err()
);
assert!(
client
.register_request_handler("", |_, _| async { Ok(serde_json::json!({})) })
.is_err()
);
let request = serde_json::json!({"jsonrpc":"2.0","id":42,"method":"host.shutdown","params":{}});
write_framed(&mut server_write, &serde_json::to_vec(&request).unwrap()).await;
let response = tokio::time::timeout(
std::time::Duration::from_secs(2),
read_framed(&mut server_read),
)
.await
.unwrap();
assert_eq!(response["error"]["code"], -32603);
assert_eq!(response["error"]["message"], "drain failed");
client.force_stop();
}
#[tokio::test]
async fn connection_eof_cancels_inbound_handlers() {
use std::sync::Arc;
use tokio::sync::Notify;
struct OnDrop(Arc<Notify>);
impl Drop for OnDrop {
fn drop(&mut self) {
self.0.notify_one();
}
}
let (client_write, _server_read) = duplex(8192);
let (mut server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
let started = Arc::new(Notify::new());
let cancelled = Arc::new(Notify::new());
let callback_started = started.clone();
let callback_cancelled = cancelled.clone();
client
.register_request_handler("host.shutdown", move |_, client| {
let started = callback_started.clone();
let cancelled = callback_cancelled.clone();
async move {
let _guard = OnDrop(cancelled);
let _client = client;
started.notify_one();
std::future::pending().await
}
})
.unwrap();
let request = serde_json::json!({"jsonrpc":"2.0","id":43,"method":"host.shutdown","params":{}});
write_framed(&mut server_write, &serde_json::to_vec(&request).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), started.notified())
.await
.unwrap();
drop(client);
server_write.shutdown().await.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(2), cancelled.notified())
.await
.unwrap();
}
#[cfg(not(feature = "runtime"))]
#[tokio::test]
async fn external_stream_build_cannot_launch_or_discover_a_runtime() {
let error = Client::start(github_copilot_sdk::ClientOptions::default())
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("requires the `runtime` Cargo feature")
);
}
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 tokio::io::AsyncRead + Unpin)) -> serde_json::Value {
let mut header = String::new();
loop {
let mut byte = [0u8; 1];
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];
AsyncReadExt::read_exact(reader, &mut buf).await.unwrap();
serde_json::from_slice(&buf).unwrap()
}
async fn verify_with_result(
result: serde_json::Value,
) -> (Result<(), github_copilot_sdk::Error>, Option<u32>) {
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let connect_req = read_framed(&mut server_read).await;
assert_eq!(connect_req["method"], "connect");
let not_found = serde_json::json!({
"jsonrpc": "2.0",
"id": connect_req["id"],
"error": { "code": -32601, "message": "Method not found" },
});
write_framed(&mut server_write, &serde_json::to_vec(¬_found).unwrap()).await;
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "ping");
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": result,
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
let res = tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap();
let version = client.protocol_version();
(res, version)
}
#[tokio::test]
async fn accepted_when_version_in_range() {
let (res, version) = verify_with_result(serde_json::json!({ "protocolVersion": 3 })).await;
assert!(res.is_ok());
assert_eq!(version, Some(3));
}
#[tokio::test]
async fn rejected_when_version_out_of_range() {
let (res, version) = verify_with_result(serde_json::json!({ "protocolVersion": 1 })).await;
let err = res.unwrap_err();
assert!(matches!(
err.kind(),
github_copilot_sdk::ErrorKind::Protocol(
github_copilot_sdk::ProtocolErrorKind::VersionMismatch { server: 1, .. }
)
));
assert_eq!(version, None);
}
#[tokio::test]
async fn succeeds_when_version_missing() {
let (res, version) = verify_with_result(serde_json::json!({ "message": "pong" })).await;
assert!(res.is_ok());
assert_eq!(version, None);
}
#[tokio::test]
async fn connect_handshake_supplies_protocol_version() {
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams(client_read, client_write, std::env::temp_dir()).unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "connect");
assert!(req["params"].get("token").is_none());
assert!(req["params"].get("clientInfo").is_none());
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": { "ok": true, "protocolVersion": 3, "version": "test-1.0.0" },
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
let res = tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap();
assert!(res.is_ok());
assert_eq!(client.protocol_version(), Some(3));
}
#[tokio::test]
async fn connect_handshake_forwards_explicit_token() {
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams_with_connection_token(
client_read,
client_write,
std::env::temp_dir(),
Some("explicit-token-abc".to_string()),
)
.unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "connect");
assert_eq!(req["params"]["token"], "explicit-token-abc");
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": { "ok": true, "protocolVersion": 3, "version": "test-1.0.0" },
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap()
.unwrap();
}
#[tokio::test]
async fn connect_handshake_forwards_auto_generated_token() {
let token = Client::generate_connection_token_for_test();
assert_eq!(token.len(), 32, "expected 32-char hex, got {token:?}");
assert!(
token
.chars()
.all(|c| c.is_ascii_hexdigit() && !c.is_uppercase()),
"expected lowercase hex, got {token:?}",
);
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams_with_connection_token(
client_read,
client_write,
std::env::temp_dir(),
Some(token.clone()),
)
.unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "connect");
assert_eq!(req["params"]["token"], token);
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": { "ok": true, "protocolVersion": 3, "version": "test-1.0.0" },
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap()
.unwrap();
}
#[tokio::test]
async fn connect_handshake_forwards_client_info() {
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams_with_client_info(
client_read,
client_write,
std::env::temp_dir(),
Some(
github_copilot_sdk::ClientInfo::new()
.with_application_name("acme-developer-portal")
.with_application_version("2.4.0")
.with_integration_name("copilot-assistant")
.with_integration_version("1.5.0"),
),
)
.unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "connect");
let client_info = &req["params"]["clientInfo"];
assert_eq!(client_info["editorName"], "acme-developer-portal");
assert_eq!(client_info["editorVersion"], "2.4.0");
assert_eq!(client_info["extensionName"], "copilot-assistant");
assert_eq!(client_info["extensionVersion"], "1.5.0");
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": { "ok": true, "protocolVersion": 3, "version": "test-1.0.0" },
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap()
.unwrap();
}
#[tokio::test]
async fn connect_handshake_omits_empty_client_info_fields() {
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams_with_client_info(
client_read,
client_write,
std::env::temp_dir(),
Some(
github_copilot_sdk::ClientInfo::new()
.with_application_name("example-app")
.with_application_version(""),
),
)
.unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "connect");
let client_info = &req["params"]["clientInfo"];
assert_eq!(client_info["editorName"], "example-app");
assert!(client_info.get("editorVersion").is_none());
assert!(client_info.get("extensionName").is_none());
assert!(client_info.get("extensionVersion").is_none());
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": { "ok": true, "protocolVersion": 3, "version": "test-1.0.0" },
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap()
.unwrap();
}
#[tokio::test]
async fn connect_handshake_omits_all_empty_client_info() {
let (client_write, server_read) = duplex(8192);
let (server_write, client_read) = duplex(8192);
let client = Client::from_streams_with_client_info(
client_read,
client_write,
std::env::temp_dir(),
Some(
github_copilot_sdk::ClientInfo::new()
.with_application_name("")
.with_application_version("")
.with_integration_name("")
.with_integration_version(""),
),
)
.unwrap();
let mut server_read = server_read;
let mut server_write = server_write;
let verify_handle = tokio::spawn({
let client = client.clone();
async move { client.verify_protocol_version().await }
});
let req = read_framed(&mut server_read).await;
assert_eq!(req["method"], "connect");
assert!(
req["params"].get("clientInfo").is_none(),
"an all-empty clientInfo must be omitted from the handshake"
);
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": req["id"],
"result": { "ok": true, "protocolVersion": 3, "version": "test-1.0.0" },
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
tokio::time::timeout(std::time::Duration::from_secs(2), verify_handle)
.await
.unwrap()
.unwrap()
.unwrap();
}