Skip to main content

cgraph/ipc/
mod.rs

1#![doc = include_str!("README.md")]
2
3use std::{
4    collections::HashMap,
5    fs, io,
6    os::unix::fs::PermissionsExt,
7    path::{Path, PathBuf},
8    sync::{
9        Arc,
10        atomic::{AtomicU64, AtomicUsize, Ordering},
11    },
12};
13
14use anyhow::{Context, Result, bail};
15use serde::Serialize;
16use serde_json::Value;
17use tokio::{
18    io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader},
19    net::{UnixListener, UnixStream, unix::OwnedWriteHalf},
20    runtime::Handle,
21    sync::{mpsc, oneshot},
22    task::JoinHandle,
23};
24
25use crate::{
26    ipc::protocol::{Envelope, IpcEvent, IpcRequest, IpcResponse, PROTOCOL_VERSION},
27    state::SourceLocation,
28};
29
30pub mod protocol;
31mod socket;
32
33use socket::{
34    SocketGuard, prepare_socket_path, remove_socket_if_identity_matches, socket_identity,
35    validate_socket_parent,
36};
37
38const CLIENT_QUEUE_CAPACITY: usize = 16;
39const COMMAND_QUEUE_CAPACITY: usize = 64;
40const MAX_FRAME_BYTES: usize = 1024 * 1024;
41
42#[derive(Clone, Debug)]
43pub struct IpcEventSender {
44    sender: mpsc::UnboundedSender<IpcEvent>,
45    client_count: Arc<AtomicUsize>,
46}
47
48impl IpcEventSender {
49    pub fn send_open_location(&self, location: &SourceLocation) -> Result<usize> {
50        let Some(line) = location.line else {
51            bail!("selected node has no source line");
52        };
53        let Some(character) = location.character else {
54            bail!("selected node has no source column");
55        };
56        if location.uri.is_empty() {
57            bail!("selected node has an empty source URI");
58        }
59        let client_count = self.client_count.load(Ordering::Acquire);
60        if client_count == 0 {
61            bail!("no IPC editor client is connected");
62        }
63        self.sender
64            .send(IpcEvent::OpenLocation {
65                uri: location.uri.clone(),
66                line,
67                character,
68            })
69            .map_err(|_| anyhow::anyhow!("IPC server is no longer running"))?;
70        Ok(client_count)
71    }
72
73    pub fn connected_clients(&self) -> usize {
74        self.client_count.load(Ordering::Acquire)
75    }
76}
77
78#[derive(Debug)]
79pub struct IpcCommand {
80    request_id: u64,
81    request: IpcRequest,
82    responder: IpcResponder,
83}
84
85impl IpcCommand {
86    pub(crate) fn new(request_id: u64, request: IpcRequest, responder: IpcResponder) -> Self {
87        Self {
88            request_id,
89            request,
90            responder,
91        }
92    }
93
94    pub fn request_id(&self) -> u64 {
95        self.request_id
96    }
97
98    pub fn request(&self) -> &IpcRequest {
99        &self.request
100    }
101
102    pub fn into_parts(self) -> (IpcRequest, IpcResponder) {
103        (self.request, self.responder)
104    }
105
106    #[cfg(test)]
107    pub(crate) fn test_command(
108        request_id: u64,
109        request: IpcRequest,
110    ) -> (Self, mpsc::Receiver<Arc<[u8]>>) {
111        let (sender, receiver) = mpsc::channel(1);
112        let responder = IpcResponder { request_id, sender };
113        (Self::new(request_id, request, responder), receiver)
114    }
115}
116
117#[derive(Clone, Debug)]
118pub struct IpcResponder {
119    request_id: u64,
120    sender: mpsc::Sender<Arc<[u8]>>,
121}
122
123impl IpcResponder {
124    pub fn respond(self, response: IpcResponse) -> Result<()> {
125        let frame = encode_frame(Some(self.request_id), response)?;
126        self.sender.try_send(frame).map_err(|error| match error {
127            mpsc::error::TrySendError::Full(_) => {
128                anyhow::anyhow!("IPC client response queue is full")
129            }
130            mpsc::error::TrySendError::Closed(_) => {
131                anyhow::anyhow!("IPC client disconnected before its response")
132            }
133        })
134    }
135}
136
137#[derive(Debug)]
138pub struct IpcServer {
139    socket_path: PathBuf,
140    event_sender: IpcEventSender,
141    command_receiver: Option<mpsc::Receiver<IpcCommand>>,
142    shutdown_sender: Option<oneshot::Sender<()>>,
143    task: Option<JoinHandle<io::Result<()>>>,
144    socket_guard: Option<SocketGuard>,
145}
146
147impl IpcServer {
148    pub fn start(socket_path: impl Into<PathBuf>) -> Result<Self> {
149        let socket_path = socket_path.into();
150        let runtime = Handle::try_current().context("IPC server requires a Tokio runtime")?;
151        validate_socket_parent(&socket_path)?;
152        prepare_socket_path(&socket_path)?;
153        let listener = UnixListener::bind(&socket_path)
154            .with_context(|| format!("failed to bind IPC socket {}", socket_path.display()))?;
155        let bound_identity = socket_identity(&socket_path)
156            .context("bound IPC socket disappeared before permission setup")?;
157        if let Err(error) = fs::set_permissions(&socket_path, fs::Permissions::from_mode(0o600)) {
158            remove_socket_if_identity_matches(&socket_path, bound_identity);
159            return Err(error).with_context(|| {
160                format!(
161                    "failed to restrict IPC socket permissions for {}",
162                    socket_path.display()
163                )
164            });
165        }
166        let socket_guard = SocketGuard::create(socket_path.clone())?;
167        let (event_sender, event_receiver) = mpsc::unbounded_channel();
168        let (command_sender, command_receiver) = mpsc::channel(COMMAND_QUEUE_CAPACITY);
169        let (shutdown_sender, shutdown_receiver) = oneshot::channel();
170        let client_count = Arc::new(AtomicUsize::new(0));
171        let task_client_count = Arc::clone(&client_count);
172        let task = runtime.spawn(async move {
173            let result = run_server(
174                listener,
175                event_receiver,
176                command_sender,
177                shutdown_receiver,
178                Arc::clone(&task_client_count),
179            )
180            .await;
181            task_client_count.store(0, Ordering::Release);
182            result
183        });
184
185        Ok(Self {
186            socket_path,
187            event_sender: IpcEventSender {
188                sender: event_sender,
189                client_count,
190            },
191            command_receiver: Some(command_receiver),
192            shutdown_sender: Some(shutdown_sender),
193            task: Some(task),
194            socket_guard: Some(socket_guard),
195        })
196    }
197
198    pub fn socket_path(&self) -> &Path {
199        &self.socket_path
200    }
201
202    pub fn event_sender(&self) -> IpcEventSender {
203        self.event_sender.clone()
204    }
205
206    pub fn take_command_receiver(&mut self) -> Option<mpsc::Receiver<IpcCommand>> {
207        self.command_receiver.take()
208    }
209
210    pub async fn shutdown(mut self) -> Result<()> {
211        if let Some(sender) = self.shutdown_sender.take() {
212            let _ = sender.send(());
213        }
214        if let Some(task) = self.task.take() {
215            task.await
216                .context("IPC server task failed")?
217                .context("IPC server stopped with an I/O error")?;
218        }
219        self.socket_guard.take();
220        Ok(())
221    }
222}
223
224impl Drop for IpcServer {
225    fn drop(&mut self) {
226        if let Some(sender) = self.shutdown_sender.take() {
227            let _ = sender.send(());
228        }
229        if let Some(task) = self.task.take() {
230            task.abort();
231        }
232    }
233}
234
235async fn run_server(
236    listener: UnixListener,
237    mut event_receiver: mpsc::UnboundedReceiver<IpcEvent>,
238    command_sender: mpsc::Sender<IpcCommand>,
239    mut shutdown_receiver: oneshot::Receiver<()>,
240    client_count: Arc<AtomicUsize>,
241) -> io::Result<()> {
242    let (disconnected_sender, mut disconnected_receiver) = mpsc::unbounded_channel();
243    let mut clients = HashMap::<u64, mpsc::Sender<Arc<[u8]>>>::new();
244    let next_client_id = AtomicU64::new(1);
245
246    loop {
247        tokio::select! {
248            accepted = listener.accept() => {
249                let (stream, _) = accepted?;
250                let client_id = next_client_id.fetch_add(1, Ordering::Relaxed);
251                let (sender, receiver) = mpsc::channel(CLIENT_QUEUE_CAPACITY);
252                clients.insert(client_id, sender.clone());
253                client_count.store(clients.len(), Ordering::Release);
254                spawn_client_connection(
255                    client_id,
256                    stream,
257                    sender,
258                    receiver,
259                    command_sender.clone(),
260                    disconnected_sender.clone(),
261                );
262            }
263            Some(event) = event_receiver.recv() => {
264                let frame = encode_frame(None, event)?;
265                clients.retain(|_, sender| sender.try_send(Arc::clone(&frame)).is_ok());
266                client_count.store(clients.len(), Ordering::Release);
267            }
268            Some(client_id) = disconnected_receiver.recv() => {
269                clients.remove(&client_id);
270                client_count.store(clients.len(), Ordering::Release);
271            }
272            _ = &mut shutdown_receiver => break,
273        }
274    }
275
276    clients.clear();
277    client_count.store(0, Ordering::Release);
278    Ok(())
279}
280
281fn spawn_client_connection(
282    client_id: u64,
283    stream: UnixStream,
284    response_sender: mpsc::Sender<Arc<[u8]>>,
285    receiver: mpsc::Receiver<Arc<[u8]>>,
286    command_sender: mpsc::Sender<IpcCommand>,
287    disconnected_sender: mpsc::UnboundedSender<u64>,
288) {
289    let _connection_task = tokio::spawn(async move {
290        let (reader, writer) = stream.into_split();
291        let read = read_client_requests(reader, response_sender, command_sender);
292        let mut writer_task = tokio::spawn(write_client_frames(writer, receiver));
293        tokio::pin!(read);
294        tokio::select! {
295            _ = &mut read => {
296                let _ = disconnected_sender.send(client_id);
297                let _ = writer_task.await;
298            }
299            _ = &mut writer_task => {
300                let _ = disconnected_sender.send(client_id);
301            }
302        }
303    });
304}
305
306async fn read_client_requests(
307    reader: tokio::net::unix::OwnedReadHalf,
308    response_sender: mpsc::Sender<Arc<[u8]>>,
309    command_sender: mpsc::Sender<IpcCommand>,
310) -> io::Result<()> {
311    let mut reader = BufReader::new(reader);
312    loop {
313        let mut frame = Vec::new();
314        let bytes_read = (&mut reader)
315            .take((MAX_FRAME_BYTES + 1) as u64)
316            .read_until(b'\n', &mut frame)
317            .await?;
318        if bytes_read == 0 {
319            return Ok(());
320        }
321        if frame.len() > MAX_FRAME_BYTES {
322            send_protocol_error(&response_sender, None, "IPC frame exceeds 1 MiB").await?;
323            return Ok(());
324        }
325        if frame.last() != Some(&b'\n') {
326            send_protocol_error(&response_sender, None, "IPC frame must end with a newline")
327                .await?;
328            return Ok(());
329        }
330        frame.pop();
331        if frame.last() == Some(&b'\r') {
332            frame.pop();
333        }
334
335        let envelope: Envelope<Value> = match serde_json::from_slice(&frame) {
336            Ok(envelope) => envelope,
337            Err(error) => {
338                send_protocol_error(
339                    &response_sender,
340                    None,
341                    &format!("invalid IPC JSON: {error}"),
342                )
343                .await?;
344                continue;
345            }
346        };
347        if envelope.version != PROTOCOL_VERSION {
348            send_protocol_error(
349                &response_sender,
350                envelope.request_id,
351                &format!(
352                    "unsupported IPC protocol version {}; expected {}",
353                    envelope.version, PROTOCOL_VERSION
354                ),
355            )
356            .await?;
357            continue;
358        }
359        let Some(request_id) = envelope.request_id else {
360            send_protocol_error(&response_sender, None, "IPC requests require a request_id")
361                .await?;
362            continue;
363        };
364        let request = match serde_json::from_value(envelope.payload) {
365            Ok(request) => request,
366            Err(error) => {
367                send_protocol_error(
368                    &response_sender,
369                    Some(request_id),
370                    &format!("invalid IPC request: {error}"),
371                )
372                .await?;
373                continue;
374            }
375        };
376        let responder = IpcResponder {
377            request_id,
378            sender: response_sender.clone(),
379        };
380        if command_sender
381            .send(IpcCommand::new(request_id, request, responder))
382            .await
383            .is_err()
384        {
385            send_protocol_error(
386                &response_sender,
387                Some(request_id),
388                "cgraph command loop is no longer available",
389            )
390            .await?;
391            return Ok(());
392        }
393    }
394}
395
396async fn write_client_frames(
397    mut writer: OwnedWriteHalf,
398    mut receiver: mpsc::Receiver<Arc<[u8]>>,
399) -> io::Result<()> {
400    while let Some(frame) = receiver.recv().await {
401        writer.write_all(&frame).await?;
402    }
403    Ok(())
404}
405
406async fn send_protocol_error(
407    sender: &mpsc::Sender<Arc<[u8]>>,
408    request_id: Option<u64>,
409    message: &str,
410) -> io::Result<()> {
411    sender
412        .send(encode_frame(
413            request_id,
414            IpcResponse::Error {
415                message: message.to_owned(),
416            },
417        )?)
418        .await
419        .map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "IPC client disconnected"))
420}
421
422fn encode_frame<T: Serialize>(request_id: Option<u64>, payload: T) -> io::Result<Arc<[u8]>> {
423    let mut frame =
424        serde_json::to_vec(&Envelope::new(request_id, payload)).map_err(io::Error::other)?;
425    frame.push(b'\n');
426    Ok(Arc::from(frame))
427}
428
429#[cfg(test)]
430mod tests;