1use std::collections::HashMap;
2use std::sync::Arc;
3use std::sync::atomic::{AtomicU64, Ordering};
4use std::time::Instant;
5
6use parking_lot::{Mutex, RwLock};
7use serde::{Deserialize, Serialize};
8use serde_json::Value;
9use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader};
10use tokio::sync::{broadcast, mpsc, oneshot};
11use tokio::task::JoinHandle;
12use tokio_util::sync::CancellationToken;
13use tracing::{Instrument, debug, error, warn};
14
15use crate::{Error, ErrorKind, ProtocolErrorKind};
16
17pub(crate) type InlineResponseCallback =
27 Box<dyn FnOnce(&JsonRpcResponse) -> Result<(), Error> + Send + Sync>;
28
29struct PendingRequest {
32 sender: oneshot::Sender<JsonRpcResponse>,
33 inline_callback: Option<InlineResponseCallback>,
34}
35
36#[derive(Debug, Clone, Serialize, Deserialize)]
38#[serde(rename_all = "camelCase")]
39pub struct JsonRpcRequest {
40 pub jsonrpc: String,
42 pub id: u64,
44 pub method: String,
46 #[serde(skip_serializing_if = "Option::is_none")]
48 pub params: Option<Value>,
49}
50
51#[derive(Debug, Clone, Serialize, Deserialize)]
53#[serde(rename_all = "camelCase")]
54pub struct JsonRpcResponse {
55 pub jsonrpc: String,
57 pub id: u64,
59 #[serde(skip_serializing_if = "Option::is_none")]
61 pub result: Option<Value>,
62 #[serde(skip_serializing_if = "Option::is_none")]
64 pub error: Option<JsonRpcError>,
65}
66
67#[derive(Debug, Clone, Serialize, Deserialize)]
69pub struct JsonRpcError {
70 pub code: i32,
72 pub message: String,
74 #[serde(skip_serializing_if = "Option::is_none")]
76 pub data: Option<Value>,
77}
78
79pub mod error_codes {
81 pub const METHOD_NOT_FOUND: i32 = -32601;
83 pub const INVALID_PARAMS: i32 = -32602;
85 #[allow(dead_code, reason = "standard JSON-RPC code, reserved for future use")]
87 pub const INTERNAL_ERROR: i32 = -32603;
88}
89
90#[derive(Debug, Clone, Serialize, Deserialize)]
92#[serde(rename_all = "camelCase")]
93pub struct JsonRpcNotification {
94 pub jsonrpc: String,
96 pub method: String,
98 #[serde(skip_serializing_if = "Option::is_none")]
100 pub params: Option<Value>,
101}
102
103#[derive(Debug, Clone, Serialize)]
105pub enum JsonRpcMessage {
106 Request(JsonRpcRequest),
108 Response(JsonRpcResponse),
110 Notification(JsonRpcNotification),
112}
113
114impl<'de> Deserialize<'de> for JsonRpcMessage {
123 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
124 where
125 D: serde::Deserializer<'de>,
126 {
127 let mut value = Value::deserialize(deserializer)?;
128 let obj = value
129 .as_object_mut()
130 .ok_or_else(|| serde::de::Error::custom("expected a JSON object"))?;
131
132 let has_id = obj.contains_key("id");
133 let has_method = obj.contains_key("method");
134
135 let payload_key = if has_id && !has_method {
138 "result"
139 } else {
140 "params"
141 };
142 let payload = obj.remove(payload_key).filter(|value| !value.is_null());
143
144 if has_id && has_method {
145 JsonRpcRequest::deserialize(value)
146 .map(|mut request| {
147 request.params = payload;
148 JsonRpcMessage::Request(request)
149 })
150 .map_err(serde::de::Error::custom)
151 } else if has_id {
152 JsonRpcResponse::deserialize(value)
153 .map(|mut response| {
154 response.result = payload;
155 JsonRpcMessage::Response(response)
156 })
157 .map_err(serde::de::Error::custom)
158 } else {
159 JsonRpcNotification::deserialize(value)
160 .map(|mut notification| {
161 notification.params = payload;
162 JsonRpcMessage::Notification(notification)
163 })
164 .map_err(serde::de::Error::custom)
165 }
166 }
167}
168
169impl JsonRpcRequest {
170 pub fn new(id: u64, method: &str, params: Option<Value>) -> Self {
172 Self {
173 jsonrpc: "2.0".to_string(),
174 id,
175 method: method.to_string(),
176 params,
177 }
178 }
179}
180
181impl JsonRpcResponse {
182 #[allow(dead_code)]
184 pub fn is_error(&self) -> bool {
185 self.error.is_some()
186 }
187}
188
189const CONTENT_LENGTH_HEADER: &str = "Content-Length: ";
190
191fn repair_lone_surrogates(body: &[u8]) -> Option<Vec<u8>> {
196 fn hex_escape_at(body: &[u8], index: usize) -> Option<u16> {
197 let digits = body.get(index + 2..index + 6)?;
198 let text = std::str::from_utf8(digits).ok()?;
199 u16::from_str_radix(text, 16).ok()
200 }
201
202 let mut repaired = None;
203 let mut in_string = false;
204 let mut index = 0;
205
206 while index < body.len() {
207 let byte = body[index];
208
209 if !in_string {
210 in_string = byte == b'"';
211 index += 1;
212 continue;
213 }
214
215 match byte {
216 b'"' => {
217 in_string = false;
218 index += 1;
219 }
220 b'\\' if body.get(index + 1) != Some(&b'u') => index += 2,
223 b'\\' => {
224 let Some(unit) = hex_escape_at(body, index) else {
225 index += 2;
226 continue;
227 };
228
229 let is_pair = (0xD800..0xDC00).contains(&unit)
230 && body.get(index + 6) == Some(&b'\\')
231 && body.get(index + 7) == Some(&b'u')
232 && hex_escape_at(body, index + 6)
233 .is_some_and(|low| (0xDC00..0xE000).contains(&low));
234
235 if is_pair {
236 index += 12;
237 continue;
238 }
239
240 if (0xD800..0xE000).contains(&unit) {
241 let output = repaired.get_or_insert_with(|| body.to_vec());
242 output[index..index + 6].copy_from_slice(br"\ufffd");
243 }
244 index += 6;
245 }
246 _ => index += 1,
247 }
248 }
249
250 repaired
251}
252
253struct WriteCommand {
262 frame: Vec<u8>,
263 ack: oneshot::Sender<Result<(), std::io::Error>>,
264}
265
266pub struct JsonRpcClient {
277 request_id: AtomicU64,
278 write_tx: mpsc::UnboundedSender<WriteCommand>,
285 pending_requests: Arc<RwLock<HashMap<u64, PendingRequest>>>,
286 notification_tx: broadcast::Sender<JsonRpcNotification>,
287 request_tx: mpsc::UnboundedSender<JsonRpcRequest>,
288 connection_closed: CancellationToken,
289 pub(crate) confirmation_requests: Arc<crate::installation_confirmation::ConfirmationRequests>,
290 read_task: Mutex<Option<JoinHandle<()>>>,
291 write_task: Mutex<Option<JoinHandle<()>>>,
292}
293
294impl JsonRpcClient {
295 pub fn new(
302 writer: impl AsyncWrite + Unpin + Send + 'static,
303 reader: impl AsyncRead + Unpin + Send + 'static,
304 notification_tx: broadcast::Sender<JsonRpcNotification>,
305 request_tx: mpsc::UnboundedSender<JsonRpcRequest>,
306 ) -> Self {
307 let (write_tx, write_rx) = mpsc::unbounded_channel::<WriteCommand>();
308
309 let writer_span = tracing::error_span!("jsonrpc_write_loop");
310 let write_task = tokio::spawn(Self::write_loop(writer, write_rx).instrument(writer_span));
311
312 let client = Self {
313 request_id: AtomicU64::new(1),
314 write_tx,
315 pending_requests: Arc::new(RwLock::new(HashMap::new())),
316 notification_tx,
317 request_tx,
318 connection_closed: CancellationToken::new(),
319 confirmation_requests: Arc::new(
320 crate::installation_confirmation::ConfirmationRequests::default(),
321 ),
322 read_task: Mutex::new(None),
323 write_task: Mutex::new(Some(write_task)),
324 };
325
326 let pending_requests = client.pending_requests.clone();
327 let notification_tx_clone = client.notification_tx.clone();
328 let request_tx_clone = client.request_tx.clone();
329 let connection_closed = client.connection_closed.clone();
330 let confirmation_requests = client.confirmation_requests.clone();
331 let reader_span = tracing::error_span!("jsonrpc_read_loop");
332
333 let read_task = tokio::spawn(
334 async move {
335 Self::read_loop(
336 reader,
337 pending_requests,
338 notification_tx_clone,
339 request_tx_clone,
340 confirmation_requests.clone(),
341 )
342 .await;
343 connection_closed.cancel();
344 confirmation_requests.clear();
345 }
346 .instrument(reader_span),
347 );
348 *client.read_task.lock() = Some(read_task);
349
350 client
351 }
352
353 pub(crate) fn force_close(&self) {
354 self.connection_closed.cancel();
355 self.confirmation_requests.clear();
356 if let Some(task) = self.read_task.lock().take() {
357 task.abort();
358 }
359 if let Some(task) = self.write_task.lock().take() {
360 task.abort();
361 }
362 self.pending_requests.write().clear();
363 }
364
365 pub(crate) fn connection_closed_token(&self) -> CancellationToken {
366 self.connection_closed.child_token()
367 }
368
369 async fn write_loop(
382 mut writer: impl AsyncWrite + Unpin + Send + 'static,
383 mut rx: mpsc::UnboundedReceiver<WriteCommand>,
384 ) {
385 while let Some(WriteCommand { frame, ack }) = rx.recv().await {
386 let result = async {
387 writer.write_all(&frame).await?;
388 writer.flush().await?;
389 Ok::<_, std::io::Error>(())
390 }
391 .await;
392
393 let _ = ack.send(result);
397 }
398 }
399
400 async fn read_loop(
401 reader: impl AsyncRead + Unpin + Send,
402 pending_requests: Arc<RwLock<HashMap<u64, PendingRequest>>>,
403 notification_tx: broadcast::Sender<JsonRpcNotification>,
404 request_tx: mpsc::UnboundedSender<JsonRpcRequest>,
405 confirmation_requests: Arc<crate::installation_confirmation::ConfirmationRequests>,
406 ) {
407 let mut reader = BufReader::new(reader);
408
409 loop {
410 match Self::read_message(&mut reader).await {
411 Ok(Some(message)) => match message {
412 JsonRpcMessage::Response(mut response) => {
413 let id = response.id;
414 let pending = pending_requests.write().remove(&id);
415 if let Some(PendingRequest {
416 sender,
417 inline_callback,
418 }) = pending
419 {
420 if let Some(cb) = inline_callback
426 && response.error.is_none()
427 {
428 let cb_outcome =
429 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
430 cb(&response)
431 }));
432 match cb_outcome {
433 Ok(Ok(())) => {}
434 Ok(Err(error)) => {
435 response.result = None;
436 response.error = Some(JsonRpcError {
437 code: -32603,
438 message: error.to_string(),
439 data: None,
440 });
441 }
442 Err(panic) => {
443 let message = panic
444 .downcast_ref::<&'static str>()
445 .map(|s| (*s).to_string())
446 .or_else(|| panic.downcast_ref::<String>().cloned())
447 .unwrap_or_else(|| {
448 "inline response callback panicked".to_string()
449 });
450 response.result = None;
451 response.error = Some(JsonRpcError {
452 code: -32603,
453 message,
454 data: None,
455 });
456 }
457 }
458 }
459 if sender.send(response).is_err() {
460 warn!(request_id = %id, "failed to send response for request");
461 }
462 } else {
463 warn!(request_id = %id, "received response for unknown request id");
464 }
465 }
466 JsonRpcMessage::Notification(notification) => {
467 if notification.method == "$/cancelRequest" {
468 if let Some(id) = notification
469 .params
470 .as_ref()
471 .and_then(|params| params.get("id"))
472 .and_then(Value::as_u64)
473 {
474 confirmation_requests.cancel(id);
475 } else {
476 warn!("invalid numeric request cancellation");
477 }
478 }
479 let _ = notification_tx.send(notification);
480 }
481 JsonRpcMessage::Request(request) => {
482 if request.method == crate::installation_confirmation::CONFIRM_METHOD
483 && !confirmation_requests.register(request.id)
484 {
485 warn!("duplicate pending installation confirmation request ID");
486 break;
487 }
488 if request_tx.send(request).is_err() {
489 warn!("failed to forward JSON-RPC request, channel closed");
490 }
491 }
492 },
493 Ok(None) => {
494 break;
495 }
496 Err(e) => {
497 error!(error = %e, "error reading from CLI");
498 break;
499 }
500 }
501 }
502
503 let mut pending = pending_requests.write();
506 if !pending.is_empty() {
507 warn!(
508 count = pending.len(),
509 "draining pending requests after read loop exit"
510 );
511 pending.clear();
512 }
513 }
514
515 async fn read_message(
516 reader: &mut BufReader<impl AsyncRead + Unpin>,
517 ) -> Result<Option<JsonRpcMessage>, Error> {
518 let mut line = String::new();
519 let mut content_length = None;
520
521 loop {
522 line.clear();
523 if reader.read_line(&mut line).await? == 0 {
524 return Ok(None);
525 }
526
527 let trimmed = line.trim();
528 if trimmed.is_empty() {
529 break;
530 }
531
532 if let Some(value) = trimmed.strip_prefix(CONTENT_LENGTH_HEADER) {
533 content_length = Some(value.trim().parse::<usize>().map_err(|_| {
534 Error::from(ErrorKind::Protocol(
535 ProtocolErrorKind::InvalidContentLength(value.trim().to_string()),
536 ))
537 })?);
538 }
539 }
540
541 let Some(length) = content_length else {
542 return Err(ErrorKind::Protocol(ProtocolErrorKind::MissingContentLength).into());
543 };
544
545 let mut body = vec![0u8; length];
546 reader.read_exact(&mut body).await?;
547
548 match serde_json::from_slice::<JsonRpcMessage>(&body) {
549 Ok(message) => Ok(Some(message)),
550 Err(error) => {
551 match repair_lone_surrogates(&body)
554 .and_then(|repaired| serde_json::from_slice::<JsonRpcMessage>(&repaired).ok())
555 {
556 Some(message) => {
557 warn!(
558 error = %error,
559 length,
560 "recovered JSON-RPC frame containing unpaired UTF-16 surrogates"
561 );
562 Ok(Some(message))
563 }
564 None => Err(error.into()),
565 }
566 }
567 }
568 }
569
570 #[allow(dead_code, reason = "public API exported via crate::JsonRpcClient")]
581 pub async fn send_request(
582 &self,
583 method: &str,
584 params: Option<serde_json::Value>,
585 ) -> Result<JsonRpcResponse, Error> {
586 self.send_request_with_inline_callback(method, params, None)
587 .await
588 }
589
590 pub(crate) async fn send_request_with_inline_callback(
607 &self,
608 method: &str,
609 params: Option<serde_json::Value>,
610 inline_callback: Option<InlineResponseCallback>,
611 ) -> Result<JsonRpcResponse, Error> {
612 let request_start = Instant::now();
613 let id = self.request_id.fetch_add(1, Ordering::SeqCst);
614 let request = JsonRpcRequest::new(id, method, params);
615
616 let (tx, rx) = oneshot::channel();
617 self.pending_requests.write().insert(
618 id,
619 PendingRequest {
620 sender: tx,
621 inline_callback,
622 },
623 );
624
625 let mut guard = PendingGuard {
630 map: &self.pending_requests,
631 id,
632 armed: true,
633 };
634
635 if let Err(error) = self.write(&request).await {
639 warn!(
640 elapsed_ms = request_start.elapsed().as_millis(),
641 method = %method,
642 request_id = id,
643 status = "failed",
644 error = %error,
645 "JsonRpcClient::send_request JSON-RPC request finished"
646 );
647 return Err(error);
648 }
649
650 let response = match rx.await {
651 Ok(response) => response,
652 Err(_) => {
653 let error = ErrorKind::Protocol(ProtocolErrorKind::RequestCancelled).into();
654 warn!(
655 elapsed_ms = request_start.elapsed().as_millis(),
656 method = %method,
657 request_id = id,
658 status = "failed",
659 error = %error,
660 "JsonRpcClient::send_request JSON-RPC request finished"
661 );
662 return Err(error);
663 }
664 };
665 guard.disarm();
666 if let Some(error) = &response.error {
667 warn!(
668 elapsed_ms = request_start.elapsed().as_millis(),
669 method = %method,
670 request_id = id,
671 status = "failed",
672 code = error.code,
673 error = %error.message,
674 "JsonRpcClient::send_request JSON-RPC request finished"
675 );
676 } else {
677 debug!(
678 elapsed_ms = request_start.elapsed().as_millis(),
679 method = %method,
680 request_id = id,
681 status = "succeeded",
682 "JsonRpcClient::send_request JSON-RPC request finished"
683 );
684 }
685 Ok(response)
686 }
687
688 pub async fn write<T: serde::Serialize>(&self, message: &T) -> Result<(), Error> {
697 let body = serde_json::to_vec(message)?;
698 let mut frame = Vec::with_capacity(CONTENT_LENGTH_HEADER.len() + 16 + body.len() + 4);
699 frame.extend_from_slice(CONTENT_LENGTH_HEADER.as_bytes());
700 frame.extend_from_slice(body.len().to_string().as_bytes());
701 frame.extend_from_slice(b"\r\n\r\n");
702 frame.extend_from_slice(&body);
703
704 let (ack_tx, ack_rx) = oneshot::channel();
705 self.write_tx
706 .send(WriteCommand { frame, ack: ack_tx })
707 .map_err(|_| {
708 Error::from(std::io::Error::new(
709 std::io::ErrorKind::BrokenPipe,
710 "writer actor has shut down",
711 ))
712 })?;
713
714 match ack_rx.await {
715 Ok(Ok(())) => Ok(()),
716 Ok(Err(e)) => Err(Error::from(e)),
717 Err(_) => Err(Error::from(std::io::Error::new(
718 std::io::ErrorKind::BrokenPipe,
719 "writer actor dropped ack without responding",
720 ))),
721 }
722 }
723}
724
725struct PendingGuard<'a> {
729 map: &'a RwLock<HashMap<u64, PendingRequest>>,
730 id: u64,
731 armed: bool,
732}
733
734impl PendingGuard<'_> {
735 fn disarm(&mut self) {
736 self.armed = false;
737 }
738}
739
740impl Drop for PendingGuard<'_> {
741 fn drop(&mut self) {
742 if self.armed {
743 self.map.write().remove(&self.id);
744 }
745 }
746}
747
748#[cfg(test)]
749mod tests {
750 use super::*;
751
752 #[test]
753 fn deserialize_notification() {
754 let json = r#"{"jsonrpc":"2.0","method":"session.event","params":{"id":"e1"}}"#;
755 let msg: JsonRpcMessage = serde_json::from_str(json).unwrap();
756 assert!(matches!(msg, JsonRpcMessage::Notification(n) if n.method == "session.event"));
757 }
758
759 #[test]
760 fn deserialize_request() {
761 let json =
762 r#"{"jsonrpc":"2.0","id":5,"method":"permission.request","params":{"kind":"shell"}}"#;
763 let msg: JsonRpcMessage = serde_json::from_str(json).unwrap();
764 assert!(
765 matches!(msg, JsonRpcMessage::Request(r) if r.id == 5 && r.method == "permission.request")
766 );
767 }
768
769 #[test]
770 fn deserialize_response_with_result() {
771 let json = r#"{"jsonrpc":"2.0","id":3,"result":{"ok":true}}"#;
772 let msg: JsonRpcMessage = serde_json::from_str(json).unwrap();
773 assert!(matches!(msg, JsonRpcMessage::Response(r) if r.id == 3 && !r.is_error()));
774 }
775
776 #[test]
777 fn deserialize_error_response() {
778 let json = r#"{"jsonrpc":"2.0","id":7,"error":{"code":-32600,"message":"Invalid Request","data":{"nested":[1,{"reason":"invalid"}]}}}"#;
779 let msg: JsonRpcMessage = serde_json::from_str(json).unwrap();
780 match msg {
781 JsonRpcMessage::Response(r) => {
782 assert!(r.is_error());
783 let err = r.error.unwrap();
784 assert_eq!(err.code, -32600);
785 assert_eq!(err.message, "Invalid Request");
786 assert_eq!(
787 err.data,
788 Some(serde_json::json!({"nested": [1, {"reason": "invalid"}]}))
789 );
790 }
791 other => panic!("expected Response, got {other:?}"),
792 }
793 }
794
795 #[test]
796 fn deserialize_rejects_non_object() {
797 let result = serde_json::from_str::<JsonRpcMessage>(r#""not an object""#);
798 assert!(result.is_err());
799 }
800
801 #[test]
802 fn deserialize_preserves_optional_payloads() {
803 for payload in [
804 None,
805 Some(Value::Null),
806 Some(serde_json::json!(false)),
807 Some(serde_json::json!(42)),
808 Some(serde_json::json!("text")),
809 Some(serde_json::json!([{"nested": [1, null, true]}])),
810 Some(serde_json::json!({"rows": [{"content": "result"}]})),
811 ] {
812 for mut envelope in [
813 serde_json::json!({"jsonrpc": "2.0", "method": "notify"}),
814 serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "request"}),
815 serde_json::json!({"jsonrpc": "2.0", "id": 1}),
816 ] {
817 let (payload_key, ignored_key) = if envelope.get("method").is_some() {
818 ("params", "result")
819 } else {
820 ("result", "params")
821 };
822 envelope[ignored_key] = serde_json::json!({"ignored": "opposite payload"});
823 if let Some(payload) = &payload {
824 envelope[payload_key] = payload.clone();
825 }
826 let actual = match serde_json::from_value::<JsonRpcMessage>(envelope).unwrap() {
827 JsonRpcMessage::Request(request) => request.params,
828 JsonRpcMessage::Response(response) => response.result,
829 JsonRpcMessage::Notification(notification) => notification.params,
830 };
831 assert_eq!(actual, payload.clone().filter(|value| !value.is_null()));
832 }
833 }
834 }
835
836 #[test]
837 fn deserialize_rejects_invalid_metadata() {
838 for json in [
839 r#"{"jsonrpc":null,"method":"notify","params":{"nested":[1]}}"#,
840 r#"{"jsonrpc":"2.0","method":42,"params":{"nested":[1]}}"#,
841 r#"{"jsonrpc":"2.0","id":null,"result":{}}"#,
842 r#"{"jsonrpc":"2.0","id":"1","result":{}}"#,
843 r#"{"jsonrpc":"2.0","id":-1,"result":{}}"#,
844 r#"{"jsonrpc":"2.0","id":1,"method":null,"params":{}}"#,
845 r#"{"jsonrpc":"2.0","id":1,"result":{},"error":{"code":"bad","message":"error"}}"#,
846 ] {
847 assert!(
848 serde_json::from_str::<JsonRpcMessage>(json).is_err(),
849 "{json}"
850 );
851 }
852 }
853
854 #[test]
855 fn request_new_sets_version() {
856 let req = JsonRpcRequest::new(42, "test.method", None);
857 assert_eq!(req.jsonrpc, "2.0");
858 assert_eq!(req.id, 42);
859 assert_eq!(req.method, "test.method");
860 assert!(req.params.is_none());
861 }
862
863 #[test]
864 fn request_serializes_camel_case() {
865 let req = JsonRpcRequest::new(1, "ping", Some(serde_json::json!({})));
866 let json = serde_json::to_string(&req).unwrap();
867 assert!(json.contains(r#""jsonrpc":"2.0""#));
868 assert!(json.contains(r#""id":1"#));
869 assert!(json.contains(r#""method":"ping""#));
870 }
871
872 #[test]
873 fn notification_without_params_omits_field() {
874 let n = JsonRpcNotification {
875 jsonrpc: "2.0".into(),
876 method: "ping".into(),
877 params: None,
878 };
879 let json = serde_json::to_string(&n).unwrap();
880 assert!(!json.contains("params"));
881 }
882
883 #[test]
884 fn response_without_error_omits_field() {
885 let r = JsonRpcResponse {
886 jsonrpc: "2.0".into(),
887 id: 1,
888 result: Some(serde_json::json!(true)),
889 error: None,
890 };
891 let json = serde_json::to_string(&r).unwrap();
892 assert!(!json.contains("error"));
893 }
894}