Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 82 additions & 2 deletions rust/src/jsonrpc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,68 @@ impl JsonRpcResponse {

const CONTENT_LENGTH_HEADER: &str = "Content-Length: ";

/// Rewrites unpaired UTF-16 surrogate escapes to `\uFFFD`.
///
/// Returns `None` when the body contains no unpaired surrogate, so valid
/// frames do not incur a repair allocation.
fn repair_lone_surrogates(body: &[u8]) -> Option<Vec<u8>> {
fn hex_escape_at(body: &[u8], index: usize) -> Option<u16> {
let digits = body.get(index + 2..index + 6)?;
let text = std::str::from_utf8(digits).ok()?;
u16::from_str_radix(text, 16).ok()
}

let mut repaired = None;
let mut in_string = false;
let mut index = 0;

while index < body.len() {
let byte = body[index];

if !in_string {
in_string = byte == b'"';
index += 1;
continue;
}

match byte {
b'"' => {
in_string = false;
index += 1;
}
// Consume non-Unicode escapes whole so an escaped backslash cannot
// be mistaken for the start of a surrogate escape.
b'\\' if body.get(index + 1) != Some(&b'u') => index += 2,
b'\\' => {
let Some(unit) = hex_escape_at(body, index) else {
index += 2;
continue;
};

let is_pair = (0xD800..0xDC00).contains(&unit)
&& body.get(index + 6) == Some(&b'\\')
&& body.get(index + 7) == Some(&b'u')
&& hex_escape_at(body, index + 6)
.is_some_and(|low| (0xDC00..0xE000).contains(&low));

if is_pair {
index += 12;
continue;
}

if (0xD800..0xE000).contains(&unit) {
let output = repaired.get_or_insert_with(|| body.to_vec());
output[index..index + 6].copy_from_slice(br"\ufffd");
}
index += 6;
}
_ => index += 1,
}
}

repaired
}

/// One framed JSON-RPC message handed to the writer actor.
///
/// `frame` is the fully serialized bytes (header + body); the caller pays
Expand Down Expand Up @@ -428,8 +490,26 @@ impl JsonRpcClient {
let mut body = vec![0u8; length];
reader.read_exact(&mut body).await?;

let message: JsonRpcMessage = serde_json::from_slice(&body)?;
Ok(Some(message))
match serde_json::from_slice::<JsonRpcMessage>(&body) {
Ok(message) => Ok(Some(message)),
Err(error) => {
// Dropping an undecodable frame could leave its pending
// request waiting forever because this layer has no timeout.
match repair_lone_surrogates(&body)
.and_then(|repaired| serde_json::from_slice::<JsonRpcMessage>(&repaired).ok())
{
Some(message) => {
warn!(
error = %error,
length,
"recovered JSON-RPC frame containing unpaired UTF-16 surrogates"
);
Ok(Some(message))
}
None => Err(error.into()),
}
}
}
}

/// Send a JSON-RPC request and wait for the matching response.
Expand Down
137 changes: 136 additions & 1 deletion rust/tests/jsonrpc_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
#![allow(clippy::unwrap_used)]

use github_copilot_sdk::test_support::{JsonRpcClient, JsonRpcNotification, JsonRpcRequest};
use tokio::io::{AsyncWrite, AsyncWriteExt, duplex};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, duplex};
use tokio::sync::{broadcast, mpsc};

/// Write a Content-Length framed JSON-RPC message to a writer.
Expand All @@ -13,6 +13,28 @@ async fn write_framed(writer: &mut (impl AsyncWrite + Unpin), body: &[u8]) {
writer.flush().await.unwrap();
}

async fn read_framed(reader: &mut (impl AsyncRead + Unpin)) -> Vec<u8> {
let mut header = String::new();
loop {
let mut byte = [0u8; 1];
reader.read_exact(&mut byte).await.unwrap();
header.push(byte[0] as char);
if header.ends_with("\r\n\r\n") {
break;
}
}

let length = header
.trim()
.strip_prefix("Content-Length: ")
.unwrap()
.parse()
.unwrap();
let mut body = vec![0u8; length];
reader.read_exact(&mut body).await.unwrap();
body
}

#[tokio::test]
async fn request_response_round_trip() {
// duplex: client_write → server_read, server_write → client_read
Expand Down Expand Up @@ -410,3 +432,116 @@ async fn send_request_cancellation_does_not_leak_pending() {
assert_eq!(response.result.unwrap()["ok"], true);
server_task.await.unwrap();
}

#[test]
fn lone_surrogate_yields_unexpected_end_of_hex_escape() {
let error = serde_json::from_slice::<serde_json::Value>(br#""\ud83d""#).unwrap_err();

assert_eq!(
error.to_string(),
"unexpected end of hex escape at line 1 column 8"
);
}

#[tokio::test]
async fn lone_surrogate_frame_is_recovered_without_closing_connection() {
let (client_write, mut server_read) = duplex(4096);
let (mut server_write, client_read) = duplex(4096);
let (notification_tx, _) = broadcast::channel(16);
let (request_tx, _) = mpsc::unbounded_channel();
let client = JsonRpcClient::new(client_write, client_read, notification_tx, request_tx);

let server_task = tokio::spawn(async move {
let request: JsonRpcRequest =
serde_json::from_slice(&read_framed(&mut server_read).await).unwrap();
let response = format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"name":"invalid \ud83d value"}}}}"#,
request.id
);
write_framed(&mut server_write, response.as_bytes()).await;

let request: JsonRpcRequest =
serde_json::from_slice(&read_framed(&mut server_read).await).unwrap();
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": request.id,
"result": {"name": "still connected"}
});
write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await;
});

let response = client.send_request("models.list", None).await.unwrap();
assert_eq!(
response.result.unwrap()["name"],
serde_json::json!("invalid \u{FFFD} value")
);

let response = client.send_request("account.getQuota", None).await.unwrap();
assert_eq!(
response.result.unwrap()["name"],
serde_json::json!("still connected")
);
server_task.await.unwrap();
}

#[tokio::test]
async fn unrepairable_frame_remains_fatal() {
let (client_write, mut server_read) = duplex(4096);
let (mut server_write, client_read) = duplex(4096);
let (notification_tx, _) = broadcast::channel(16);
let (request_tx, _) = mpsc::unbounded_channel();
let client = JsonRpcClient::new(client_write, client_read, notification_tx, request_tx);

let server_task = tokio::spawn(async move {
let request: JsonRpcRequest =
serde_json::from_slice(&read_framed(&mut server_read).await).unwrap();
let response = format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"surrogate":"\ud83d","escape":"\q"}}}}"#,
request.id
);
write_framed(&mut server_write, response.as_bytes()).await;
});

let error = tokio::time::timeout(
std::time::Duration::from_secs(2),
client.send_request("models.list", None),
)
.await
.expect("unrepairable frame did not terminate the pending request")
.unwrap_err();

assert_eq!(error.to_string(), "request cancelled");
assert!(error.is_transport_failure());
server_task.await.unwrap();
}

#[tokio::test]
async fn valid_pairs_and_escaped_backslashes_are_untouched() {
let (client_write, mut server_read) = duplex(4096);
let (mut server_write, client_read) = duplex(4096);
let (notification_tx, _) = broadcast::channel(16);
let (request_tx, _) = mpsc::unbounded_channel();
let client = JsonRpcClient::new(client_write, client_read, notification_tx, request_tx);

let server_task = tokio::spawn(async move {
let request: JsonRpcRequest =
serde_json::from_slice(&read_framed(&mut server_read).await).unwrap();
let response = format!(
r#"{{"jsonrpc":"2.0","id":{},"result":{{"emoji":"\ud83d\ude00","path":"C:\\ud83d","invalid":"\ud83d"}}}}"#,
request.id
);
write_framed(&mut server_write, response.as_bytes()).await;
});

let result = client
.send_request("models.list", None)
.await
.unwrap()
.result
.unwrap();

assert_eq!(result["emoji"], serde_json::json!("😀"));
assert_eq!(result["path"], serde_json::json!(r"C:\ud83d"));
assert_eq!(result["invalid"], serde_json::json!("\u{FFFD}"));
server_task.await.unwrap();
}
Loading