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
165 changes: 164 additions & 1 deletion codex-rs/exec-server-protocol/src/rpc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,24 @@ use std::fmt;

use codex_protocol::protocol::W3cTraceContext;
use serde::Deserialize;
use serde::Deserializer;
use serde::Serialize;
use serde::de;
use serde::de::DeserializeSeed;
use serde::de::MapAccess;
use serde::de::SeqAccess;
use serde::de::Visitor;
use serde_json::Map;
use serde_json::Number;
use serde_json::Value;

pub const JSONRPC_VERSION: &str = "2.0";

// A maximum-size fs/walk response has at most 50,000 entries and needs roughly
// 150,000 JSON values. Keep ample headroom for legitimate protocol messages
// while preventing compact arrays from expanding into millions of heap values.
const MAX_JSONRPC_VALUE_NODES: usize = 256 * 1024;

#[derive(Debug, Clone, PartialEq, PartialOrd, Ord, Deserialize, Serialize, Hash, Eq)]
#[serde(untagged)]
pub enum RequestId {
Expand All @@ -30,7 +44,7 @@ impl fmt::Display for RequestId {
pub type Result = serde_json::Value;

/// Any valid exec-server JSON-RPC object that can be decoded from or encoded onto the wire.
#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
#[derive(Debug, Clone, PartialEq, Serialize)]
#[serde(untagged)]
pub enum JSONRPCMessage {
Request(JSONRPCRequest),
Expand All @@ -39,6 +53,151 @@ pub enum JSONRPCMessage {
Error(JSONRPCError),
}

#[derive(Deserialize)]
#[serde(untagged)]
enum JSONRPCMessageRepr {
Request(JSONRPCRequest),
Notification(JSONRPCNotification),
Response(JSONRPCResponse),
Error(JSONRPCError),
}

impl From<JSONRPCMessageRepr> for JSONRPCMessage {
fn from(value: JSONRPCMessageRepr) -> Self {
match value {
JSONRPCMessageRepr::Request(request) => Self::Request(request),
JSONRPCMessageRepr::Notification(notification) => Self::Notification(notification),
JSONRPCMessageRepr::Response(response) => Self::Response(response),
JSONRPCMessageRepr::Error(error) => Self::Error(error),
}
}
}

impl<'de> Deserialize<'de> for JSONRPCMessage {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let mut remaining = MAX_JSONRPC_VALUE_NODES;
let value = BoundedValueSeed {
remaining: &mut remaining,
}
.deserialize(deserializer)?;
JSONRPCMessageRepr::deserialize(value)
.map(Self::from)
.map_err(de::Error::custom)
}
}

struct BoundedValueSeed<'a> {
remaining: &'a mut usize,
}

impl<'de> DeserializeSeed<'de> for BoundedValueSeed<'_> {
type Value = Value;

fn deserialize<D>(self, deserializer: D) -> std::result::Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
let Some(remaining) = self.remaining.checked_sub(1) else {
return Err(de::Error::custom(format!(
"JSON-RPC message exceeds the limit of {MAX_JSONRPC_VALUE_NODES} JSON values"
)));
};
*self.remaining = remaining;
deserializer.deserialize_any(BoundedValueVisitor {
remaining: self.remaining,
})
}
}

struct BoundedValueVisitor<'a> {
remaining: &'a mut usize,
}

impl<'de> Visitor<'de> for BoundedValueVisitor<'_> {
type Value = Value;

fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a JSON value within the exec-server complexity limit")
}

fn visit_bool<E>(self, value: bool) -> std::result::Result<Self::Value, E> {
Ok(Value::Bool(value))
}

fn visit_i64<E>(self, value: i64) -> std::result::Result<Self::Value, E> {
Ok(Value::Number(value.into()))
}

fn visit_u64<E>(self, value: u64) -> std::result::Result<Self::Value, E> {
Ok(Value::Number(value.into()))
}

fn visit_f64<E>(self, value: f64) -> std::result::Result<Self::Value, E> {
Ok(Number::from_f64(value).map_or(Value::Null, Value::Number))
}

fn visit_str<E>(self, value: &str) -> std::result::Result<Self::Value, E> {
Ok(Value::String(value.to_string()))
}

fn visit_string<E>(self, value: String) -> std::result::Result<Self::Value, E> {
Ok(Value::String(value))
}

fn visit_none<E>(self) -> std::result::Result<Self::Value, E> {
Ok(Value::Null)
}

fn visit_some<D>(self, deserializer: D) -> std::result::Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
BoundedValueSeed {
remaining: self.remaining,
}
.deserialize(deserializer)
}

fn visit_unit<E>(self) -> std::result::Result<Self::Value, E> {
Ok(Value::Null)
}

fn visit_seq<A>(self, mut sequence: A) -> std::result::Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let mut values = Vec::new();
while let Some(value) = sequence.next_element_seed(BoundedValueSeed {
remaining: &mut *self.remaining,
})? {
values.push(value);
}
Ok(Value::Array(values))
}

fn visit_map<A>(self, mut object: A) -> std::result::Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let mut values = Map::new();
while let Some(key) = object.next_key::<String>()? {
if values.contains_key(&key) {
return Err(de::Error::custom(format!(
"duplicate JSON object key `{key}`"
)));
}
let value = object.next_value_seed(BoundedValueSeed {
remaining: &mut *self.remaining,
})?;
values.insert(key, value);
}
Ok(Value::Object(values))
}
}

/// A request that expects a response.
#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
pub struct JSONRPCRequest {
Expand Down Expand Up @@ -79,3 +238,7 @@ pub struct JSONRPCErrorError {
pub data: Option<serde_json::Value>,
pub message: String,
}

#[cfg(test)]
#[path = "rpc_tests.rs"]
mod tests;
99 changes: 99 additions & 0 deletions codex-rs/exec-server-protocol/src/rpc_tests.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
use pretty_assertions::assert_eq;
use serde_json::json;

use super::JSONRPCError;
use super::JSONRPCErrorError;
use super::JSONRPCMessage;
use super::JSONRPCNotification;
use super::JSONRPCRequest;
use super::JSONRPCResponse;
use super::MAX_JSONRPC_VALUE_NODES;
use super::RequestId;

#[test]
fn round_trips_every_jsonrpc_message_variant() -> serde_json::Result<()> {
let messages = [
JSONRPCMessage::Request(JSONRPCRequest {
id: RequestId::Integer(1),
method: "request".to_string(),
params: Some(json!({"items": [1, 2, 3]})),
trace: None,
}),
JSONRPCMessage::Notification(JSONRPCNotification {
method: "notification".to_string(),
params: Some(json!({"enabled": true})),
}),
JSONRPCMessage::Response(JSONRPCResponse {
id: RequestId::String("response".to_string()),
result: json!({"value": "ok"}),
}),
JSONRPCMessage::Error(JSONRPCError {
error: JSONRPCErrorError {
code: -32603,
data: Some(json!({"retryable": false})),
message: "failed".to_string(),
},
id: RequestId::Integer(2),
}),
];

for expected in messages {
let encoded = serde_json::to_string(&expected)?;
let actual = serde_json::from_str::<JSONRPCMessage>(&encoded)?;
assert_eq!(actual, expected);
}

Ok(())
}

#[test]
fn accepts_large_scalar_payload() -> serde_json::Result<()> {
let expected = JSONRPCMessage::Notification(JSONRPCNotification {
method: "large".to_string(),
params: Some(json!({"data": "x".repeat(MAX_JSONRPC_VALUE_NODES + 1)})),
});

let encoded = serde_json::to_string(&expected)?;
let actual = serde_json::from_str::<JSONRPCMessage>(&encoded)?;

assert_eq!(actual, expected);
Ok(())
}

#[test]
fn rejects_duplicate_object_keys() {
let error = serde_json::from_str::<JSONRPCMessage>(r#"{"method":"safe","method":"dangerous"}"#)
.expect_err("duplicate JSON object keys should be rejected");

assert!(
error
.to_string()
.contains("duplicate JSON object key `method`"),
"unexpected error: {error}"
);
}

#[test]
fn rejects_compact_array_heap_amplification() {
const REPRO_VALUE_COUNT: usize = 2_097_137;
const REPRO_MESSAGE_BYTES: usize = 4_194_303;

let mut encoded = String::with_capacity(REPRO_MESSAGE_BYTES);
encoded.push_str(r#"{"method":"probe","params":["#);
for index in 0..REPRO_VALUE_COUNT {
if index != 0 {
encoded.push(',');
}
encoded.push('0');
}
encoded.push_str("]}");
assert_eq!(encoded.len(), REPRO_MESSAGE_BYTES);

let error = serde_json::from_str::<JSONRPCMessage>(&encoded)
.expect_err("amplification payload should exceed the JSON value limit");
let expected_error = format!("exceeds the limit of {MAX_JSONRPC_VALUE_NODES} JSON values");
assert!(
error.to_string().contains(&expected_error),
"unexpected error: {error}"
);
}
Loading
Loading