Skip to main content

codex_exec_server_protocol/
rpc.rs

1//! JSON-RPC wire envelopes used by exec-server.
2//!
3//! Exec-server uses the Codex JSON-RPC dialect, which omits the
4//! `"jsonrpc": "2.0"` field on the wire.
5
6use std::fmt;
7
8use codex_protocol::protocol::W3cTraceContext;
9use serde::Deserialize;
10use serde::Deserializer;
11use serde::Serialize;
12use serde::de;
13use serde::de::DeserializeSeed;
14use serde::de::MapAccess;
15use serde::de::SeqAccess;
16use serde::de::Visitor;
17use serde_json::Map;
18use serde_json::Number;
19use serde_json::Value;
20
21pub const JSONRPC_VERSION: &str = "2.0";
22
23// A maximum-size fs/walk response has at most 50,000 entries and needs roughly
24// 150,000 JSON values. Keep ample headroom for legitimate protocol messages
25// while preventing compact arrays from expanding into millions of heap values.
26const MAX_JSONRPC_VALUE_NODES: usize = 256 * 1024;
27
28// With `arbitrary_precision` enabled, serde_json presents decimal, exponent,
29// and out-of-range integer values to visitors as a synthetic one-entry map.
30const SERDE_JSON_NUMBER_TOKEN: &str = "$serde_json::private::Number";
31const SERDE_JSON_RAW_VALUE_TOKEN: &str = "$serde_json::private::RawValue";
32
33#[derive(Debug, Clone, PartialEq, PartialOrd, Ord, Deserialize, Serialize, Hash, Eq)]
34#[serde(untagged)]
35pub enum RequestId {
36    String(String),
37    Integer(i64),
38}
39
40impl fmt::Display for RequestId {
41    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
42        match self {
43            Self::String(value) => f.write_str(value),
44            Self::Integer(value) => write!(f, "{value}"),
45        }
46    }
47}
48
49pub type Result = serde_json::Value;
50
51/// Any valid exec-server JSON-RPC object that can be decoded from or encoded onto the wire.
52#[derive(Debug, Clone, PartialEq, Serialize)]
53#[serde(untagged)]
54pub enum JSONRPCMessage {
55    Request(JSONRPCRequest),
56    Notification(JSONRPCNotification),
57    Response(JSONRPCResponse),
58    Error(JSONRPCError),
59}
60
61impl<'de> Deserialize<'de> for JSONRPCMessage {
62    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
63    where
64        D: Deserializer<'de>,
65    {
66        let mut remaining = MAX_JSONRPC_VALUE_NODES;
67        let value = BoundedValueSeed {
68            remaining: &mut remaining,
69        }
70        .deserialize(deserializer)?;
71        let object = value
72            .as_object()
73            .ok_or_else(|| de::Error::custom("expected a JSON-RPC object"))?;
74
75        if object.contains_key("method") {
76            if object.contains_key("id") {
77                JSONRPCRequest::deserialize(value)
78                    .map(Self::Request)
79                    .map_err(de::Error::custom)
80            } else {
81                JSONRPCNotification::deserialize(value)
82                    .map(Self::Notification)
83                    .map_err(de::Error::custom)
84            }
85        } else if object.contains_key("result") {
86            JSONRPCResponse::deserialize(value)
87                .map(Self::Response)
88                .map_err(de::Error::custom)
89        } else {
90            JSONRPCError::deserialize(value)
91                .map(Self::Error)
92                .map_err(de::Error::custom)
93        }
94    }
95}
96
97struct BoundedValueSeed<'a> {
98    remaining: &'a mut usize,
99}
100
101impl<'de> DeserializeSeed<'de> for BoundedValueSeed<'_> {
102    type Value = Value;
103
104    fn deserialize<D>(self, deserializer: D) -> std::result::Result<Self::Value, D::Error>
105    where
106        D: Deserializer<'de>,
107    {
108        let Some(remaining) = self.remaining.checked_sub(1) else {
109            return Err(de::Error::custom(format!(
110                "JSON-RPC message exceeds the limit of {MAX_JSONRPC_VALUE_NODES} JSON values"
111            )));
112        };
113        *self.remaining = remaining;
114        deserializer.deserialize_any(BoundedValueVisitor {
115            remaining: self.remaining,
116        })
117    }
118}
119
120struct BoundedValueVisitor<'a> {
121    remaining: &'a mut usize,
122}
123
124impl<'de> Visitor<'de> for BoundedValueVisitor<'_> {
125    type Value = Value;
126
127    fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
128        formatter.write_str("a JSON value within the exec-server complexity limit")
129    }
130
131    fn visit_bool<E>(self, value: bool) -> std::result::Result<Self::Value, E> {
132        Ok(Value::Bool(value))
133    }
134
135    fn visit_i64<E>(self, value: i64) -> std::result::Result<Self::Value, E> {
136        Ok(Value::Number(value.into()))
137    }
138
139    fn visit_u64<E>(self, value: u64) -> std::result::Result<Self::Value, E> {
140        Ok(Value::Number(value.into()))
141    }
142
143    fn visit_f64<E>(self, value: f64) -> std::result::Result<Self::Value, E> {
144        Ok(Number::from_f64(value).map_or(Value::Null, Value::Number))
145    }
146
147    fn visit_str<E>(self, value: &str) -> std::result::Result<Self::Value, E> {
148        Ok(Value::String(value.to_string()))
149    }
150
151    fn visit_string<E>(self, value: String) -> std::result::Result<Self::Value, E> {
152        Ok(Value::String(value))
153    }
154
155    fn visit_none<E>(self) -> std::result::Result<Self::Value, E> {
156        Ok(Value::Null)
157    }
158
159    fn visit_some<D>(self, deserializer: D) -> std::result::Result<Self::Value, D::Error>
160    where
161        D: Deserializer<'de>,
162    {
163        BoundedValueSeed {
164            remaining: self.remaining,
165        }
166        .deserialize(deserializer)
167    }
168
169    fn visit_unit<E>(self) -> std::result::Result<Self::Value, E> {
170        Ok(Value::Null)
171    }
172
173    fn visit_seq<A>(self, mut sequence: A) -> std::result::Result<Self::Value, A::Error>
174    where
175        A: SeqAccess<'de>,
176    {
177        let mut values = Vec::new();
178        while let Some(value) = sequence.next_element_seed(BoundedValueSeed {
179            remaining: &mut *self.remaining,
180        })? {
181            values.push(value);
182        }
183        Ok(Value::Array(values))
184    }
185
186    fn visit_map<A>(self, mut object: A) -> std::result::Result<Self::Value, A::Error>
187    where
188        A: MapAccess<'de>,
189    {
190        let Some(first_key) = object.next_key::<String>()? else {
191            return Ok(Value::Object(Map::new()));
192        };
193
194        if first_key == SERDE_JSON_NUMBER_TOKEN {
195            let encoded = object.next_value::<String>()?;
196            let number = encoded.parse::<Number>().map_err(de::Error::custom)?;
197            return Ok(Value::Number(number));
198        }
199
200        if first_key == SERDE_JSON_RAW_VALUE_TOKEN {
201            let encoded = object.next_value::<String>()?;
202            let mut deserializer = serde_json::Deserializer::from_str(&encoded);
203            // The raw wrapper already consumed one value from the budget. Reuse
204            // that slot for the decoded root while charging all of its children.
205            let value = (&mut deserializer)
206                .deserialize_any(BoundedValueVisitor {
207                    remaining: &mut *self.remaining,
208                })
209                .map_err(de::Error::custom)?;
210            deserializer.end().map_err(de::Error::custom)?;
211            return Ok(value);
212        }
213
214        let mut values = Map::new();
215        let first_value = object.next_value_seed(BoundedValueSeed {
216            remaining: &mut *self.remaining,
217        })?;
218        values.insert(first_key, first_value);
219
220        while let Some(key) = object.next_key::<String>()? {
221            if values.contains_key(&key) {
222                return Err(de::Error::custom(format!(
223                    "duplicate JSON object key `{key}`"
224                )));
225            }
226            let value = object.next_value_seed(BoundedValueSeed {
227                remaining: &mut *self.remaining,
228            })?;
229            values.insert(key, value);
230        }
231        Ok(Value::Object(values))
232    }
233}
234
235/// A request that expects a response.
236#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
237pub struct JSONRPCRequest {
238    pub id: RequestId,
239    pub method: String,
240    #[serde(default, skip_serializing_if = "Option::is_none")]
241    pub params: Option<serde_json::Value>,
242    #[serde(default, skip_serializing_if = "Option::is_none")]
243    pub trace: Option<W3cTraceContext>,
244}
245
246/// A notification that does not expect a response.
247#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
248pub struct JSONRPCNotification {
249    pub method: String,
250    #[serde(default, skip_serializing_if = "Option::is_none")]
251    pub params: Option<serde_json::Value>,
252}
253
254/// A successful response to a request.
255#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
256pub struct JSONRPCResponse {
257    pub id: RequestId,
258    pub result: Result,
259}
260
261/// A response indicating that a request failed.
262#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
263pub struct JSONRPCError {
264    pub error: JSONRPCErrorError,
265    pub id: RequestId,
266}
267
268#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
269pub struct JSONRPCErrorError {
270    pub code: i64,
271    #[serde(default, skip_serializing_if = "Option::is_none")]
272    pub data: Option<serde_json::Value>,
273    pub message: String,
274}
275
276#[cfg(test)]
277#[path = "rpc_tests.rs"]
278mod tests;