Skip to main content

vv_agent/app_server/protocol/
jsonrpc.rs

1use schemars::JsonSchema;
2use serde::de::Error as _;
3use serde::ser::SerializeMap;
4use serde::{Deserialize, Deserializer, Serialize, Serializer};
5use serde_json::Value;
6use ts_rs::TS;
7
8use super::errors::AppServerErrorCode;
9
10pub const JSON_RPC_VERSION: &str = "2.0";
11
12#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, JsonSchema, TS)]
13#[serde(untagged)]
14pub enum RequestId {
15    String(String),
16    Integer(i64),
17    Null,
18}
19
20#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema, TS)]
21#[serde(untagged)]
22pub enum JsonRpcMessage {
23    Request(JsonRpcRequest),
24    Notification(JsonRpcNotification),
25    Response(JsonRpcResponse),
26    Error(JsonRpcError),
27}
28
29#[derive(Debug, Clone, PartialEq, JsonSchema, TS)]
30pub struct JsonRpcRequest {
31    pub id: RequestId,
32    pub method: String,
33    #[serde(default, skip_serializing_if = "Option::is_none")]
34    pub params: Option<Value>,
35}
36
37#[derive(Debug, Clone, PartialEq, JsonSchema, TS)]
38pub struct JsonRpcNotification {
39    pub method: String,
40    #[serde(default, skip_serializing_if = "Option::is_none")]
41    pub params: Option<Value>,
42}
43
44#[derive(Debug, Clone, PartialEq, JsonSchema, TS)]
45pub struct JsonRpcResponse {
46    pub id: RequestId,
47    pub result: Value,
48}
49
50#[derive(Debug, Clone, PartialEq, JsonSchema, TS)]
51pub struct JsonRpcError {
52    pub id: RequestId,
53    pub error: JsonRpcErrorBody,
54}
55
56#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema, TS)]
57pub struct JsonRpcErrorBody {
58    pub code: i64,
59    pub message: String,
60    #[serde(default, skip_serializing_if = "Option::is_none")]
61    pub data: Option<Value>,
62}
63
64#[derive(Deserialize)]
65#[serde(deny_unknown_fields)]
66struct JsonRpcRequestWire {
67    jsonrpc: String,
68    id: RequestId,
69    method: String,
70    #[serde(default)]
71    params: Option<Value>,
72}
73
74#[derive(Deserialize)]
75#[serde(deny_unknown_fields)]
76struct JsonRpcNotificationWire {
77    jsonrpc: String,
78    method: String,
79    #[serde(default)]
80    params: Option<Value>,
81}
82
83#[derive(Deserialize)]
84#[serde(deny_unknown_fields)]
85struct JsonRpcResponseWire {
86    jsonrpc: String,
87    id: RequestId,
88    result: Value,
89}
90
91#[derive(Deserialize)]
92#[serde(deny_unknown_fields)]
93struct JsonRpcErrorWire {
94    jsonrpc: String,
95    id: RequestId,
96    error: JsonRpcErrorBody,
97}
98
99fn validate_version<E: serde::de::Error>(version: &str) -> Result<(), E> {
100    if version == JSON_RPC_VERSION {
101        Ok(())
102    } else {
103        Err(E::custom("jsonrpc must be exactly 2.0"))
104    }
105}
106
107fn non_null_id_error(id: &RequestId) -> Option<&'static str> {
108    if matches!(id, RequestId::Null) {
109        Some("JSON-RPC request id cannot be null")
110    } else {
111        None
112    }
113}
114
115fn validate_method<E: serde::de::Error>(method: &str) -> Result<(), E> {
116    if method.is_empty() {
117        Err(E::custom("JSON-RPC method cannot be empty"))
118    } else {
119        Ok(())
120    }
121}
122
123impl<'de> Deserialize<'de> for JsonRpcRequest {
124    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
125    where
126        D: Deserializer<'de>,
127    {
128        let wire = JsonRpcRequestWire::deserialize(deserializer)?;
129        validate_version::<D::Error>(&wire.jsonrpc)?;
130        if let Some(message) = non_null_id_error(&wire.id) {
131            return Err(D::Error::custom(message));
132        }
133        validate_method::<D::Error>(&wire.method)?;
134        Ok(Self {
135            id: wire.id,
136            method: wire.method,
137            params: wire.params,
138        })
139    }
140}
141
142impl<'de> Deserialize<'de> for JsonRpcNotification {
143    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
144    where
145        D: Deserializer<'de>,
146    {
147        let wire = JsonRpcNotificationWire::deserialize(deserializer)?;
148        validate_version::<D::Error>(&wire.jsonrpc)?;
149        validate_method::<D::Error>(&wire.method)?;
150        Ok(Self {
151            method: wire.method,
152            params: wire.params,
153        })
154    }
155}
156
157impl<'de> Deserialize<'de> for JsonRpcResponse {
158    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
159    where
160        D: Deserializer<'de>,
161    {
162        let wire = JsonRpcResponseWire::deserialize(deserializer)?;
163        validate_version::<D::Error>(&wire.jsonrpc)?;
164        if let Some(message) = non_null_id_error(&wire.id) {
165            return Err(D::Error::custom(message));
166        }
167        Ok(Self {
168            id: wire.id,
169            result: wire.result,
170        })
171    }
172}
173
174impl<'de> Deserialize<'de> for JsonRpcError {
175    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
176    where
177        D: Deserializer<'de>,
178    {
179        let wire = JsonRpcErrorWire::deserialize(deserializer)?;
180        validate_version::<D::Error>(&wire.jsonrpc)?;
181        Ok(Self {
182            id: wire.id,
183            error: wire.error,
184        })
185    }
186}
187
188impl<'de> Deserialize<'de> for JsonRpcMessage {
189    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
190    where
191        D: Deserializer<'de>,
192    {
193        let value = Value::deserialize(deserializer)?;
194        let object = value
195            .as_object()
196            .ok_or_else(|| D::Error::custom("JSON-RPC message must be an object"))?;
197        let has_method = object.contains_key("method");
198        let has_id = object.contains_key("id");
199        let has_result = object.contains_key("result");
200        let has_error = object.contains_key("error");
201        if has_method && !has_result && !has_error {
202            return if has_id {
203                serde_json::from_value::<JsonRpcRequest>(value)
204                    .map(Self::Request)
205                    .map_err(D::Error::custom)
206            } else {
207                serde_json::from_value::<JsonRpcNotification>(value)
208                    .map(Self::Notification)
209                    .map_err(D::Error::custom)
210            };
211        }
212        if has_id && has_result && !has_error && !has_method {
213            return serde_json::from_value::<JsonRpcResponse>(value)
214                .map(Self::Response)
215                .map_err(D::Error::custom);
216        }
217        if has_id && has_error && !has_result && !has_method {
218            return serde_json::from_value::<JsonRpcError>(value)
219                .map(Self::Error)
220                .map_err(D::Error::custom);
221        }
222        Err(D::Error::custom("Invalid JSON-RPC message"))
223    }
224}
225
226impl Serialize for JsonRpcRequest {
227    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
228    where
229        S: Serializer,
230    {
231        if let Some(message) = non_null_id_error(&self.id) {
232            return Err(<S::Error as serde::ser::Error>::custom(message));
233        }
234        let mut map = serializer.serialize_map(Some(3 + usize::from(self.params.is_some())))?;
235        map.serialize_entry("jsonrpc", JSON_RPC_VERSION)?;
236        map.serialize_entry("id", &self.id)?;
237        map.serialize_entry("method", &self.method)?;
238        if let Some(params) = &self.params {
239            map.serialize_entry("params", params)?;
240        }
241        map.end()
242    }
243}
244
245impl Serialize for JsonRpcNotification {
246    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
247    where
248        S: Serializer,
249    {
250        let mut map = serializer.serialize_map(Some(2 + usize::from(self.params.is_some())))?;
251        map.serialize_entry("jsonrpc", JSON_RPC_VERSION)?;
252        map.serialize_entry("method", &self.method)?;
253        if let Some(params) = &self.params {
254            map.serialize_entry("params", params)?;
255        }
256        map.end()
257    }
258}
259
260impl Serialize for JsonRpcResponse {
261    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
262    where
263        S: Serializer,
264    {
265        if let Some(message) = non_null_id_error(&self.id) {
266            return Err(<S::Error as serde::ser::Error>::custom(message));
267        }
268        let mut map = serializer.serialize_map(Some(3))?;
269        map.serialize_entry("jsonrpc", JSON_RPC_VERSION)?;
270        map.serialize_entry("id", &self.id)?;
271        map.serialize_entry("result", &self.result)?;
272        map.end()
273    }
274}
275
276impl Serialize for JsonRpcError {
277    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
278    where
279        S: Serializer,
280    {
281        let mut map = serializer.serialize_map(Some(3))?;
282        map.serialize_entry("jsonrpc", JSON_RPC_VERSION)?;
283        map.serialize_entry("id", &self.id)?;
284        map.serialize_entry("error", &self.error)?;
285        map.end()
286    }
287}
288
289impl JsonRpcError {
290    pub fn new(id: RequestId, code: AppServerErrorCode, message: impl Into<String>) -> Self {
291        Self {
292            id,
293            error: JsonRpcErrorBody {
294                code: code.code(),
295                message: message.into(),
296                data: None,
297            },
298        }
299    }
300
301    pub fn with_data(mut self, data: Value) -> Self {
302        self.error.data = Some(data);
303        self
304    }
305}