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
101 changes: 101 additions & 0 deletions codex-rs/code-mode-protocol/src/host/codec.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
use std::io;
use std::mem::size_of;

use serde::Serialize;
use serde::de::DeserializeOwned;
use tokio::io::AsyncRead;
use tokio::io::AsyncReadExt;
use tokio::io::AsyncWrite;
use tokio::io::AsyncWriteExt;

/// Maximum JSON payload size accepted for one IPC frame.
pub const MAX_FRAME_BYTES: usize = 64 * 1024 * 1024;

/// Decodes JSON messages prefixed by a four-byte little-endian payload length.
pub struct FramedReader<R> {
reader: R,
}

impl<R> FramedReader<R>
where
R: AsyncRead + Unpin,
{
pub fn new(reader: R) -> Self {
Self { reader }
}

/// Reads the next frame, returning `None` only for EOF at a frame boundary.
pub async fn read<T>(&mut self) -> io::Result<Option<T>>
where
T: DeserializeOwned,
{
let mut length_bytes = [0_u8; size_of::<u32>()];
if self.reader.read(&mut length_bytes[..1]).await? == 0 {
return Ok(None);
}
self.reader.read_exact(&mut length_bytes[1..]).await?;

let length = u32::from_le_bytes(length_bytes) as usize;
if length > MAX_FRAME_BYTES {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("code-mode IPC frame length {length} exceeds {MAX_FRAME_BYTES} bytes"),
));
}

let mut payload = vec![0; length];
self.reader.read_exact(&mut payload).await?;
serde_json::from_slice(&payload).map(Some).map_err(|err| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("failed to decode code-mode IPC frame: {err}"),
)
})
}
}

/// Encodes JSON messages with a four-byte little-endian payload length.
pub struct FramedWriter<W> {
writer: W,
}

impl<W> FramedWriter<W>
where
W: AsyncWrite + Unpin,
{
pub fn new(writer: W) -> Self {
Self { writer }
}

/// Writes and flushes one complete frame.
pub async fn write<T>(&mut self, message: &T) -> io::Result<()>
where
T: Serialize,
{
let payload = serde_json::to_vec(message).map_err(|err| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("failed to encode code-mode IPC frame: {err}"),
)
})?;
if payload.len() > MAX_FRAME_BYTES {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"code-mode IPC frame length {} exceeds {MAX_FRAME_BYTES} bytes",
payload.len()
),
));
}
let length = u32::try_from(payload.len()).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
"code-mode IPC frame length exceeds u32",
)
})?;

self.writer.write_all(&length.to_le_bytes()).await?;
self.writer.write_all(&payload).await?;
self.writer.flush().await
}
}
104 changes: 104 additions & 0 deletions codex-rs/code-mode-protocol/src/host/codec_tests.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
use pretty_assertions::assert_eq;
use serde_json::json;
use tokio::io::AsyncReadExt;
use tokio::io::AsyncWriteExt;

use super::FramedReader;
use super::FramedWriter;
use super::MAX_FRAME_BYTES;

#[tokio::test]
async fn frame_wire_format_is_little_endian_length_prefixed_json() {
let (writer, mut reader) = tokio::io::duplex(/*max_buf_size*/ 128);
let write = tokio::spawn(async move {
FramedWriter::new(writer)
.write(&json!({"value": 1}))
.await
.expect("write frame");
});

let mut bytes = Vec::new();
reader.read_to_end(&mut bytes).await.expect("read bytes");
write.await.expect("writer task");

let payload = br#"{"value":1}"#;
let mut expected = (payload.len() as u32).to_le_bytes().to_vec();
expected.extend_from_slice(payload);
assert_eq!(bytes, expected);
}

#[tokio::test]
async fn fragmented_frame_round_trips() {
let value = json!({"type": "session/open", "sessionId": "session-1"});
let payload = serde_json::to_vec(&value).expect("serialize");
let mut bytes = (payload.len() as u32).to_le_bytes().to_vec();
bytes.extend(payload);

let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 128);
let write = tokio::spawn(async move {
for byte in bytes {
writer.write_all(&[byte]).await.expect("write byte");
tokio::task::yield_now().await;
}
});

assert_eq!(
FramedReader::new(reader)
.read::<serde_json::Value>()
.await
.expect("read frame"),
Some(value)
);
write.await.expect("writer task");
}

#[tokio::test]
async fn eof_is_clean_only_at_a_frame_boundary() {
let (writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
drop(writer);
assert_eq!(
FramedReader::new(reader)
.read::<serde_json::Value>()
.await
.expect("clean eof"),
None
);

let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
writer
.write_all(&[1, 0])
.await
.expect("write partial header");
drop(writer);
let err = FramedReader::new(reader)
.read::<serde_json::Value>()
.await
.expect_err("truncated header");
assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof);
}

#[tokio::test]
async fn oversized_and_malformed_frames_are_rejected() {
let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
writer
.write_all(&((MAX_FRAME_BYTES as u32) + 1).to_le_bytes())
.await
.expect("write oversized header");
let err = FramedReader::new(reader)
.read::<serde_json::Value>()
.await
.expect_err("oversized frame");
assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);

let (mut writer, reader) = tokio::io::duplex(/*max_buf_size*/ 16);
writer
.write_all(&(1_u32).to_le_bytes())
.await
.expect("write length");
writer.write_all(b"{").await.expect("write malformed json");
let err = FramedReader::new(reader)
.read::<serde_json::Value>()
.await
.expect_err("malformed frame");
assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
}
Loading
Loading