1use serde::{de, ser::SerializeStruct, Deserialize, Deserializer, Serialize, Serializer};
35use serde_json::Value;
36use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
37
38pub const MAX_FRAME_BYTES: u32 = 16 * 1024 * 1024;
43
44#[derive(Debug, thiserror::Error)]
46pub enum ProtocolError {
47 #[error("io: {0}")]
48 Io(#[from] std::io::Error),
49 #[error("frame exceeds max size: {0} > {MAX_FRAME_BYTES}")]
50 FrameTooLarge(u32),
51 #[error("malformed json frame: {0}")]
52 Json(#[from] serde_json::Error),
53 #[error("connection closed before a full frame was read")]
54 UnexpectedEof,
55}
56
57#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
61pub struct Request {
62 pub id: u64,
63 pub method: String,
64 #[serde(default)]
65 pub params: Value,
66}
67
68impl Request {
69 pub fn new(id: u64, method: impl Into<String>, params: Value) -> Self {
70 Request {
71 id,
72 method: method.into(),
73 params,
74 }
75 }
76}
77
78#[derive(Debug, Clone, PartialEq)]
84pub enum ResponsePayload {
85 Ok(Value),
87 Err(RpcError),
89}
90
91#[derive(Debug, Clone, PartialEq)]
95pub struct Response {
96 pub id: u64,
97 pub payload: ResponsePayload,
98}
99
100impl Response {
101 pub fn ok(id: u64, result: Value) -> Self {
102 Response {
103 id,
104 payload: ResponsePayload::Ok(result),
105 }
106 }
107
108 pub fn err(id: u64, code: ErrorCode, message: impl Into<String>) -> Self {
109 Response {
110 id,
111 payload: ResponsePayload::Err(RpcError {
112 code,
113 message: message.into(),
114 }),
115 }
116 }
117
118 pub fn is_err(&self) -> bool {
120 matches!(self.payload, ResponsePayload::Err(_))
121 }
122
123 pub fn result(&self) -> Option<&Value> {
125 match &self.payload {
126 ResponsePayload::Ok(v) => Some(v),
127 ResponsePayload::Err(_) => None,
128 }
129 }
130
131 pub fn error(&self) -> Option<&RpcError> {
133 match &self.payload {
134 ResponsePayload::Err(e) => Some(e),
135 ResponsePayload::Ok(_) => None,
136 }
137 }
138}
139
140impl Serialize for Response {
141 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
144 where
145 S: Serializer,
146 {
147 let mut st = serializer.serialize_struct("Response", 2)?;
148 st.serialize_field("id", &self.id)?;
149 match &self.payload {
150 ResponsePayload::Ok(v) => st.serialize_field("result", v)?,
151 ResponsePayload::Err(e) => st.serialize_field("error", e)?,
152 }
153 st.end()
154 }
155}
156
157impl<'de> Deserialize<'de> for Response {
158 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
161 where
162 D: Deserializer<'de>,
163 {
164 fn present_value<'de, D>(deserializer: D) -> Result<Option<Value>, D::Error>
172 where
173 D: Deserializer<'de>,
174 {
175 Value::deserialize(deserializer).map(Some)
176 }
177
178 #[derive(Deserialize)]
179 struct Wire {
180 id: u64,
181 #[serde(default, deserialize_with = "present_value")]
182 result: Option<Value>,
183 #[serde(default)]
184 error: Option<RpcError>,
185 }
186 let w = Wire::deserialize(deserializer)?;
187 let payload = match (w.result, w.error) {
188 (Some(_), Some(_)) => {
189 return Err(de::Error::custom(
190 "response carries both `result` and `error`",
191 ))
192 }
193 (Some(r), None) => ResponsePayload::Ok(r),
194 (None, Some(e)) => ResponsePayload::Err(e),
195 (None, None) => {
196 return Err(de::Error::custom(
197 "response carries neither `result` nor `error`",
198 ))
199 }
200 };
201 Ok(Response { id: w.id, payload })
202 }
203}
204
205#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
207pub struct RpcError {
208 pub code: ErrorCode,
209 pub message: String,
210}
211
212#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
216#[serde(rename_all = "snake_case")]
217pub enum ErrorCode {
218 MalformedFrame,
220 UnknownMethod,
222 InvalidParams,
224 AgentNotFound,
226 AgentExists,
228 InvalidStatus,
230 Busy,
232 LockTimeout,
234 SpawnFailed,
236 ChannelUnknown,
238 Internal,
240}
241
242#[derive(Debug, Clone, Copy, PartialEq, Eq)]
244pub enum Namespace {
245 Agent,
247 Channel,
249 Unknown,
251}
252
253impl Namespace {
254 pub fn of(method: &str) -> Namespace {
256 match method.split_once('.') {
257 Some(("agent", _)) => Namespace::Agent,
258 Some(("channel", _)) => Namespace::Channel,
259 _ => Namespace::Unknown,
260 }
261 }
262
263 pub fn verb(method: &str) -> Option<&str> {
265 method.split_once('.').map(|(_, v)| v)
266 }
267}
268
269pub async fn read_frame<R: AsyncRead + Unpin>(reader: &mut R) -> Result<Vec<u8>, ProtocolError> {
278 let mut len_buf = [0u8; 4];
279 match reader.read_exact(&mut len_buf).await {
280 Ok(_) => {}
281 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => {
282 return Err(ProtocolError::UnexpectedEof)
283 }
284 Err(e) => return Err(e.into()),
285 }
286 let len = u32::from_le_bytes(len_buf);
287 if len > MAX_FRAME_BYTES {
288 return Err(ProtocolError::FrameTooLarge(len));
289 }
290 let mut body = vec![0u8; len as usize];
291 reader
292 .read_exact(&mut body)
293 .await
294 .map_err(|e| match e.kind() {
295 std::io::ErrorKind::UnexpectedEof => ProtocolError::UnexpectedEof,
296 _ => ProtocolError::Io(e),
297 })?;
298 Ok(body)
299}
300
301pub async fn write_frame<W: AsyncWrite + Unpin>(
305 writer: &mut W,
306 body: &[u8],
307) -> Result<(), ProtocolError> {
308 let len: u32 = body
309 .len()
310 .try_into()
311 .map_err(|_| ProtocolError::FrameTooLarge(u32::MAX))?;
312 if len > MAX_FRAME_BYTES {
313 return Err(ProtocolError::FrameTooLarge(len));
314 }
315 writer.write_all(&len.to_le_bytes()).await?;
316 writer.write_all(body).await?;
317 writer.flush().await?;
318 Ok(())
319}
320
321pub async fn read_request<R: AsyncRead + Unpin>(reader: &mut R) -> Result<Request, ProtocolError> {
323 let body = read_frame(reader).await?;
324 Ok(serde_json::from_slice(&body)?)
325}
326
327pub async fn write_request<W: AsyncWrite + Unpin>(
329 writer: &mut W,
330 req: &Request,
331) -> Result<(), ProtocolError> {
332 let body = serde_json::to_vec(req)?;
333 write_frame(writer, &body).await
334}
335
336pub async fn read_response<R: AsyncRead + Unpin>(
338 reader: &mut R,
339) -> Result<Response, ProtocolError> {
340 let body = read_frame(reader).await?;
341 Ok(serde_json::from_slice(&body)?)
342}
343
344pub async fn write_response<W: AsyncWrite + Unpin>(
346 writer: &mut W,
347 resp: &Response,
348) -> Result<(), ProtocolError> {
349 let body = serde_json::to_vec(resp)?;
350 write_frame(writer, &body).await
351}
352
353#[cfg(test)]
354mod tests {
355 use super::*;
356 use serde_json::json;
357
358 #[test]
359 fn namespace_classification() {
360 assert_eq!(Namespace::of("agent.spawn"), Namespace::Agent);
361 assert_eq!(
362 Namespace::of("channel.register_channel"),
363 Namespace::Channel
364 );
365 assert_eq!(Namespace::of("bogus.method"), Namespace::Unknown);
366 assert_eq!(Namespace::of("noseparator"), Namespace::Unknown);
367 assert_eq!(Namespace::verb("agent.spawn"), Some("spawn"));
368 assert_eq!(Namespace::verb("nope"), None);
369 }
370
371 #[tokio::test]
372 async fn request_roundtrips_over_duplex() {
373 let (mut a, mut b) = tokio::io::duplex(4096);
374 let req = Request::new(7, "agent.spawn", json!({"name": "worker-A"}));
375 write_request(&mut a, &req).await.unwrap();
376 let got = read_request(&mut b).await.unwrap();
377 assert_eq!(got, req);
378 }
379
380 #[tokio::test]
381 async fn response_ok_and_err_roundtrip() {
382 let (mut a, mut b) = tokio::io::duplex(4096);
383 let ok = Response::ok(7, json!({"status": "live"}));
384 write_response(&mut a, &ok).await.unwrap();
385 let got = read_response(&mut b).await.unwrap();
386 assert_eq!(got, ok);
387 assert!(!got.is_err());
388
389 let err = Response::err(8, ErrorCode::AgentNotFound, "no such agent");
390 write_response(&mut a, &err).await.unwrap();
391 let got = read_response(&mut b).await.unwrap();
392 assert!(got.is_err());
393 assert_eq!(got.error().unwrap().code, ErrorCode::AgentNotFound);
394 }
395
396 #[tokio::test]
397 async fn frame_too_large_is_rejected_not_allocated() {
398 let (mut a, mut b) = tokio::io::duplex(64);
401 let writer = tokio::spawn(async move {
402 let bogus_len = (MAX_FRAME_BYTES + 1).to_le_bytes();
403 a.write_all(&bogus_len).await.unwrap();
404 a.flush().await.unwrap();
405 a
407 });
408 let err = read_frame(&mut b).await.unwrap_err();
409 assert!(matches!(err, ProtocolError::FrameTooLarge(_)));
410 let _a = writer.await.unwrap();
411 }
412
413 #[tokio::test]
414 async fn clean_eof_is_distinguished_from_io_error() {
415 let (a, mut b) = tokio::io::duplex(64);
416 drop(a); let err = read_frame(&mut b).await.unwrap_err();
418 assert!(matches!(err, ProtocolError::UnexpectedEof));
419 }
420
421 #[test]
422 fn response_wire_shape_is_flat() {
423 let ok = Response::ok(7, json!({"status": "live"}));
426 assert_eq!(
427 serde_json::to_value(&ok).unwrap(),
428 json!({"id": 7, "result": {"status": "live"}})
429 );
430 let err = Response::err(8, ErrorCode::AgentNotFound, "no such agent");
431 assert_eq!(
432 serde_json::to_value(&err).unwrap(),
433 json!({"id": 8, "error": {"code": "agent_not_found", "message": "no such agent"}})
434 );
435 }
436
437 #[test]
438 fn response_rejects_both_or_neither_payload() {
439 let both = json!({"id": 1, "result": {}, "error": {"code": "internal", "message": "x"}});
443 assert!(serde_json::from_value::<Response>(both).is_err());
444 let neither = json!({"id": 1});
446 assert!(serde_json::from_value::<Response>(neither).is_err());
447 }
448
449 #[test]
450 fn response_accepts_explicit_null_result() {
451 let parsed: Response =
455 serde_json::from_value(json!({"id": 7, "result": null})).expect("null result parses");
456 assert!(!parsed.is_err());
457 assert_eq!(parsed.result(), Some(&Value::Null));
458 assert_eq!(
459 serde_json::to_value(&parsed).unwrap(),
460 json!({"id": 7, "result": null})
461 );
462 }
463
464 #[tokio::test]
465 async fn malformed_json_body_surfaces_json_error() {
466 let (mut a, mut b) = tokio::io::duplex(4096);
467 write_frame(&mut a, b"{not json").await.unwrap();
468 let err = read_request(&mut b).await.unwrap_err();
469 assert!(matches!(err, ProtocolError::Json(_)));
470 }
471}