1use serde::de::{self, DeserializeOwned};
4use serde::{Deserialize, Deserializer, Serialize};
5use serde_json::Value;
6
7pub const MAX_PROTOCOL_FRAME_BYTES: usize = 17 * 1024 * 1024;
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
10pub enum RequestFrameType {
11 #[serde(rename = "req")]
12 Req,
13}
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
16pub enum ResponseFrameType {
17 #[serde(rename = "res")]
18 Res,
19}
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
22pub enum EventFrameType {
23 #[serde(rename = "event")]
24 Event,
25}
26
27#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
28#[serde(deny_unknown_fields)]
29pub struct RequestFrame {
30 #[serde(rename = "type")]
31 pub frame_type: RequestFrameType,
32 pub id: String,
33 pub method: String,
34 #[serde(
35 default,
36 deserialize_with = "optional_json_value",
37 skip_serializing_if = "Option::is_none"
38 )]
39 pub params: Option<Value>,
40}
41
42#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
43#[serde(rename_all = "camelCase", deny_unknown_fields)]
44pub struct ResponseFrame {
45 #[serde(rename = "type")]
46 pub frame_type: ResponseFrameType,
47 pub id: String,
48 pub ok: bool,
49 #[serde(
50 default,
51 deserialize_with = "optional_json_value",
52 skip_serializing_if = "Option::is_none"
53 )]
54 pub payload: Option<Value>,
55 #[serde(
56 default,
57 deserialize_with = "optional_non_null",
58 skip_serializing_if = "Option::is_none"
59 )]
60 pub error: Option<ErrorShape>,
61}
62
63#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
64#[serde(rename_all = "camelCase", deny_unknown_fields)]
65pub struct EventFrame {
66 #[serde(rename = "type")]
67 pub frame_type: EventFrameType,
68 pub event: String,
69 #[serde(
70 default,
71 deserialize_with = "optional_json_value",
72 skip_serializing_if = "Option::is_none"
73 )]
74 pub payload: Option<Value>,
75 #[serde(
76 default,
77 deserialize_with = "optional_non_null",
78 skip_serializing_if = "Option::is_none"
79 )]
80 pub seq: Option<u64>,
81 #[serde(
82 default,
83 deserialize_with = "optional_non_null",
84 skip_serializing_if = "Option::is_none"
85 )]
86 pub state_version: Option<StateVersion>,
87}
88
89#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
90#[serde(untagged)]
91pub enum ProtocolFrame {
92 Request(RequestFrame),
93 Response(ResponseFrame),
94 Event(EventFrame),
95}
96
97#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
98#[serde(rename_all = "camelCase", deny_unknown_fields)]
99pub struct ErrorShape {
100 pub code: String,
101 pub message: String,
102 #[serde(
103 default,
104 deserialize_with = "optional_json_value",
105 skip_serializing_if = "Option::is_none"
106 )]
107 pub details: Option<Value>,
108 #[serde(
109 default,
110 deserialize_with = "optional_non_null",
111 skip_serializing_if = "Option::is_none"
112 )]
113 pub retryable: Option<bool>,
114 #[serde(
115 default,
116 deserialize_with = "optional_non_null",
117 skip_serializing_if = "Option::is_none"
118 )]
119 pub retry_after_ms: Option<u64>,
120}
121
122#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
123#[serde(rename_all = "camelCase", deny_unknown_fields)]
124pub struct StateVersion {
125 pub presence: u64,
126 pub health: u64,
127}
128
129pub fn encode_protocol_frame(frame: &ProtocolFrame) -> Result<String, String> {
130 validate_protocol_frame(frame)?;
131 let encoded = serde_json::to_string(frame).map_err(|error| error.to_string())?;
132 if encoded.len() > MAX_PROTOCOL_FRAME_BYTES {
133 return Err("protocol frame is too large".to_string());
134 }
135 Ok(encoded)
136}
137
138pub fn decode_protocol_frame(source: &str) -> Result<ProtocolFrame, String> {
139 if source.is_empty() {
140 return Err("protocol frame is empty".to_string());
141 }
142 if source.len() > MAX_PROTOCOL_FRAME_BYTES {
143 return Err("protocol frame is too large".to_string());
144 }
145 let frame = serde_json::from_str(source).map_err(|_| "invalid protocol frame".to_string())?;
146 validate_protocol_frame(&frame)?;
147 Ok(frame)
148}
149
150pub fn validate_protocol_frame(frame: &ProtocolFrame) -> Result<(), String> {
151 match frame {
152 ProtocolFrame::Request(request) => {
153 validate_non_empty("request id", &request.id)?;
154 validate_non_empty("request method", &request.method)
155 }
156 ProtocolFrame::Response(response) => {
157 validate_non_empty("response id", &response.id)?;
158 if let Some(error) = &response.error {
159 validate_error_shape(error)?;
160 }
161 Ok(())
162 }
163 ProtocolFrame::Event(event) => validate_non_empty("event name", &event.event),
164 }
165}
166
167fn optional_json_value<'de, D>(deserializer: D) -> Result<Option<Value>, D::Error>
168where
169 D: Deserializer<'de>,
170{
171 Value::deserialize(deserializer).map(Some)
172}
173
174fn optional_non_null<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
175where
176 D: Deserializer<'de>,
177 T: DeserializeOwned,
178{
179 let value = Value::deserialize(deserializer)?;
180 if value.is_null() {
181 return Err(de::Error::custom("null is not allowed for this field"));
182 }
183 T::deserialize(value).map(Some).map_err(de::Error::custom)
184}
185
186fn validate_error_shape(error: &ErrorShape) -> Result<(), String> {
187 validate_non_empty("error code", &error.code)?;
188 validate_non_empty("error message", &error.message)
189}
190
191fn validate_non_empty(name: &str, value: &str) -> Result<(), String> {
192 if value.is_empty() {
193 Err(format!("{name} must not be empty"))
194 } else {
195 Ok(())
196 }
197}
198
199#[cfg(test)]
200mod tests {
201 use serde_json::json;
202
203 use super::*;
204
205 #[test]
206 fn protocol_frames_round_trip() {
207 let frame = ProtocolFrame::Request(RequestFrame {
208 frame_type: RequestFrameType::Req,
209 id: "request-1".to_string(),
210 method: "runtime.invoke".to_string(),
211 params: Some(json!({"endpoint": "guide"})),
212 });
213 let encoded = encode_protocol_frame(&frame).unwrap();
214 assert_eq!(decode_protocol_frame(&encoded).unwrap(), frame);
215 }
216
217 #[test]
218 fn protocol_frames_reject_extra_fields_and_null_typed_options() {
219 assert!(decode_protocol_frame(
220 r#"{"type":"req","id":"request-1","method":"runtime.invoke","extra":true}"#
221 )
222 .is_err());
223 assert!(decode_protocol_frame(r#"{"type":"event","event":"tick","seq":null}"#).is_err());
224 }
225}