codex_exec_server_protocol/
rpc.rs1use 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
23const MAX_JSONRPC_VALUE_NODES: usize = 256 * 1024;
27
28const 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#[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 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#[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#[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#[derive(Debug, Clone, PartialEq, Deserialize, Serialize)]
256pub struct JSONRPCResponse {
257 pub id: RequestId,
258 pub result: Result,
259}
260
261#[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;