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;