Skip to main content

tau_socket/
lib.rs

1//! Unix socket listener and transport adapters.
2//!
3//! This crate exposes a small transport-agnostic socket peer that reuses the
4//! same self-delimiting CBOR event codec as stdio transports.
5
6use std::io::{self, BufWriter};
7use std::os::fd::OwnedFd;
8use std::os::unix::fs::{FileTypeExt, MetadataExt, PermissionsExt as _};
9use std::os::unix::net::{UnixListener, UnixStream};
10use std::path::{Path, PathBuf};
11use std::sync::mpsc::RecvTimeoutError;
12#[cfg(test)]
13use std::sync::mpsc::SyncSender;
14use std::time::Duration;
15use std::{fmt, fs, net as path_std_net};
16
17use tau_proto::{
18    DecodeError, HarnessInputMessage, HarnessInputReader, HarnessOutputMessage,
19    HarnessOutputWriter, PeerOutputWriter,
20};
21
22use self::reader_worker::ReaderWorker;
23
24mod reader_worker;
25
26/// Errors returned by the Unix socket transport.
27#[derive(Debug)]
28pub enum SocketTransportError {
29    /// Creating the parent directory for a socket path failed.
30    CreateParentDirectory {
31        /// Directory that could not be created.
32        path: PathBuf,
33        /// Underlying filesystem error.
34        source: io::Error,
35    },
36    /// A pre-existing non-socket path blocked binding the listener.
37    RefuseNonSocketPath {
38        /// Path that already existed but was not a Unix socket.
39        path: PathBuf,
40    },
41    /// A pre-existing Unix socket accepted connections and was treated as live.
42    ActiveSocketExists {
43        /// Socket path that appears to already have a listener.
44        path: PathBuf,
45    },
46    /// A pre-existing Unix socket could not be proven inactive.
47    ProbeExistingSocket {
48        /// Socket path whose liveness probe failed.
49        path: PathBuf,
50        /// Underlying probe error.
51        source: io::Error,
52    },
53    /// Removing an inactive stale Unix socket path failed.
54    RemoveStaleSocket {
55        /// Stale socket path that could not be removed.
56        path: PathBuf,
57        /// Underlying filesystem error.
58        source: io::Error,
59    },
60    /// Binding the Unix listener failed.
61    Bind {
62        /// Socket path that could not be bound.
63        path: PathBuf,
64        /// Underlying bind error.
65        source: io::Error,
66    },
67    /// Reading metadata for the bound socket path failed.
68    BoundSocketMetadata {
69        /// Socket path whose metadata could not be read after bind.
70        path: PathBuf,
71        /// Underlying filesystem error.
72        source: io::Error,
73    },
74    /// Accepting an attached Unix socket client failed.
75    Accept {
76        /// Underlying accept error.
77        source: io::Error,
78    },
79    /// Connecting to a listener failed.
80    Connect {
81        /// Socket path that could not be connected.
82        path: PathBuf,
83        /// Underlying connect error.
84        source: io::Error,
85    },
86    /// Cloning a Unix stream failed.
87    Clone {
88        /// Underlying stream clone error.
89        source: io::Error,
90    },
91    /// Spawning the bounded background socket reader failed.
92    SpawnReader {
93        /// Underlying thread admission error.
94        source: io::Error,
95    },
96    /// Encoding a protocol message failed.
97    Encode {
98        /// Underlying protocol encode error.
99        source: tau_proto::EncodeError,
100    },
101    /// Flushing a protocol message failed.
102    Flush {
103        /// Underlying writer flush error.
104        source: io::Error,
105    },
106    /// Decoding a protocol message failed.
107    Decode {
108        /// Underlying protocol decode error.
109        source: DecodeError,
110    },
111}
112
113impl fmt::Display for SocketTransportError {
114    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
115        match self {
116            Self::CreateParentDirectory { path, source } => write!(
117                f,
118                "failed to create socket parent directory {}: {source}",
119                path.display()
120            ),
121            Self::RefuseNonSocketPath { path } => {
122                write!(f, "refusing to replace non-socket path {}", path.display())
123            }
124            Self::ActiveSocketExists { path } => {
125                write!(
126                    f,
127                    "refusing to replace active Unix socket {}",
128                    path.display()
129                )
130            }
131            Self::ProbeExistingSocket { path, source } => write!(
132                f,
133                "refusing to replace Unix socket {} after liveness probe failed: {source}",
134                path.display()
135            ),
136            Self::RemoveStaleSocket { path, source } => write!(
137                f,
138                "failed to remove stale socket {}: {source}",
139                path.display()
140            ),
141            Self::Bind { path, source } => {
142                write!(f, "failed to bind Unix socket {}: {source}", path.display())
143            }
144            Self::BoundSocketMetadata { path, source } => write!(
145                f,
146                "failed to inspect bound Unix socket {}: {source}",
147                path.display()
148            ),
149            Self::Accept { source } => write!(f, "failed to accept Unix socket client: {source}"),
150            Self::Connect { path, source } => {
151                write!(
152                    f,
153                    "failed to connect to Unix socket {}: {source}",
154                    path.display()
155                )
156            }
157            Self::Clone { source } => write!(f, "failed to clone Unix socket stream: {source}"),
158            Self::SpawnReader { source } => {
159                write!(f, "failed to spawn Unix socket reader: {source}")
160            }
161            Self::Encode { source } => write!(f, "failed to encode socket event: {source}"),
162            Self::Flush { source } => write!(f, "failed to flush socket stream: {source}"),
163            Self::Decode { source } => write!(f, "failed to decode socket event: {source}"),
164        }
165    }
166}
167
168impl std::error::Error for SocketTransportError {
169    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
170        match self {
171            Self::CreateParentDirectory { source, .. } => Some(source),
172            Self::RefuseNonSocketPath { .. } | Self::ActiveSocketExists { .. } => None,
173            Self::ProbeExistingSocket { source, .. } => Some(source),
174            Self::RemoveStaleSocket { source, .. } => Some(source),
175            Self::Bind { source, .. } => Some(source),
176            Self::BoundSocketMetadata { source, .. } => Some(source),
177            Self::Accept { source } => Some(source),
178            Self::Connect { source, .. } => Some(source),
179            Self::Clone { source } => Some(source),
180            Self::SpawnReader { source } => Some(source),
181            Self::Encode { source } => Some(source),
182            Self::Flush { source } => Some(source),
183            Self::Decode { source } => Some(source),
184        }
185    }
186}
187
188#[derive(Debug, Clone, Copy, PartialEq, Eq)]
189struct SocketIdentity {
190    // Device number from the bound socket metadata.
191    dev: u64,
192    // Inode number from the bound socket metadata.
193    ino: u64,
194}
195
196impl SocketIdentity {
197    fn from_metadata(metadata: &fs::Metadata) -> Self {
198        Self {
199            dev: metadata.dev(),
200            ino: metadata.ino(),
201        }
202    }
203}
204
205/// Unix socket listener for later-attached protocol clients.
206///
207/// When dropped, a listener removes only the socket path it created and only if
208/// that path still refers to the same device/inode recorded after binding.
209pub struct SocketListener {
210    /// Filesystem path occupied by the listener socket.
211    path: PathBuf,
212    /// Bound Unix listener accepting client streams.
213    listener: UnixListener,
214    /// Device/inode pair of the socket created by this listener.
215    socket_identity: SocketIdentity,
216}
217
218impl SocketListener {
219    /// Binds a listener at an absent path without probing or removing anything.
220    ///
221    /// The caller owns parent-directory preparation and any stale-path
222    /// reclamation under its own stronger coordination boundary.
223    ///
224    /// # Errors
225    ///
226    /// Returns an error when the path already exists, binding fails, or the
227    /// created socket cannot be inspected.
228    pub fn bind_fresh(path: impl Into<PathBuf>) -> Result<Self, SocketTransportError> {
229        let path = path.into();
230        let listener = UnixListener::bind(&path).map_err(|source| SocketTransportError::Bind {
231            path: path.clone(),
232            source,
233        })?;
234        let metadata = fs::symlink_metadata(&path).map_err(|source| {
235            SocketTransportError::BoundSocketMetadata {
236                path: path.clone(),
237                source,
238            }
239        })?;
240        fs::set_permissions(&path, fs::Permissions::from_mode(0o600)).map_err(|source| {
241            SocketTransportError::BoundSocketMetadata {
242                path: path.clone(),
243                source,
244            }
245        })?;
246        Ok(Self {
247            path,
248            listener,
249            socket_identity: SocketIdentity::from_metadata(&metadata),
250        })
251    }
252
253    /// Binds a Unix socket listener at the given path.
254    ///
255    /// Parent directories are created if needed. An inactive stale Unix socket
256    /// may be removed, but non-socket paths and active listeners are refused.
257    /// Active-listener detection opens a short-lived connection that can be
258    /// observed by an already-running daemon. Existing socket paths are treated
259    /// as stale only when that probe fails with
260    /// [`io::ErrorKind::ConnectionRefused`]; other probe failures return
261    /// [`SocketTransportError::ProbeExistingSocket`] without unlinking the
262    /// path.
263    ///
264    /// # Errors
265    ///
266    /// Returns an error when directory creation, stale socket cleanup, binding,
267    /// post-bind socket metadata inspection, non-socket refusal, or
268    /// active/unverified socket refusal fails or applies.
269    pub fn bind(path: impl Into<PathBuf>) -> Result<Self, SocketTransportError> {
270        let path = path.into();
271        create_socket_parent_if_needed(&path)?;
272        remove_inactive_stale_socket(&path)?;
273
274        Self::bind_fresh(path)
275    }
276
277    /// Returns the filesystem path of the listener socket.
278    #[must_use]
279    pub fn path(&self) -> &Path {
280        &self.path
281    }
282
283    /// Clones the raw Unix listener for daemon/internal stream handoff loops.
284    ///
285    /// This is an escape hatch for callers that need to pass accepted raw
286    /// streams to higher-level server code instead of using [`Self::accept`].
287    /// The `SocketListener` still owns identity-checked path cleanup; callers
288    /// must ensure cloned listeners are shut down before dropping it.
289    ///
290    /// # Errors
291    ///
292    /// Returns an error when the underlying listener cannot be cloned.
293    pub fn try_clone_raw_listener(&self) -> Result<UnixListener, SocketTransportError> {
294        self.listener
295            .try_clone()
296            .map_err(|source| SocketTransportError::Clone { source })
297    }
298
299    /// Accepts one attached client for harness/server-side protocol handling.
300    ///
301    /// The returned client reads [`HarnessInputMessage`] values from the peer
302    /// and writes [`HarnessOutputMessage`] values back to it.
303    ///
304    /// # Errors
305    ///
306    /// Returns an error when accepting the Unix stream or cloning it for split
307    /// reader/writer ownership fails.
308    pub fn accept(&self) -> Result<SocketAcceptedClient, SocketTransportError> {
309        let (stream, _) = self
310            .listener
311            .accept()
312            .map_err(|source| SocketTransportError::Accept { source })?;
313        SocketAcceptedClient::new(stream)
314    }
315}
316
317impl Drop for SocketListener {
318    fn drop(&mut self) {
319        let Ok(metadata) = fs::symlink_metadata(&self.path) else {
320            return;
321        };
322        if !metadata.file_type().is_socket() {
323            return;
324        }
325        if SocketIdentity::from_metadata(&metadata) == self.socket_identity {
326            let _ = fs::remove_file(&self.path);
327        }
328    }
329}
330
331/// One server-side accepted Unix socket client speaking the protocol.
332pub struct SocketAcceptedClient {
333    /// Reader for peer/client-to-harness input messages.
334    reader: HarnessInputReader<UnixStream>,
335    /// Writer for harness-to-peer/client output messages.
336    writer: HarnessOutputWriter<BufWriter<UnixStream>>,
337}
338
339impl SocketAcceptedClient {
340    fn new(stream: UnixStream) -> Result<Self, SocketTransportError> {
341        let writer_stream = stream
342            .try_clone()
343            .map_err(|source| SocketTransportError::Clone { source })?;
344        Ok(Self {
345            reader: HarnessInputReader::new(stream),
346            writer: HarnessOutputWriter::new(BufWriter::new(writer_stream)),
347        })
348    }
349
350    /// Reads one peer → harness protocol message from the accepted client.
351    ///
352    /// Returns `Ok(None)` only when the client closes cleanly at a message
353    /// boundary.
354    ///
355    /// # Errors
356    ///
357    /// Returns a decode error for malformed or truncated protocol input.
358    pub fn recv(&mut self) -> Result<Option<HarnessInputMessage>, SocketTransportError> {
359        self.reader
360            .read_message()
361            .map_err(|source| SocketTransportError::Decode { source })
362    }
363
364    /// Sends one harness → peer protocol message to the accepted client.
365    ///
366    /// # Errors
367    ///
368    /// Returns an error when encoding or flushing the message fails.
369    pub fn send(&mut self, message: &HarnessOutputMessage) -> Result<(), SocketTransportError> {
370        self.writer
371            .write_message(message)
372            .map_err(|source| SocketTransportError::Encode { source })?;
373        self.writer
374            .flush()
375            .map_err(|source| SocketTransportError::Flush { source })
376    }
377}
378
379/// Result of attempting to receive a harness → peer message from a socket.
380#[derive(Debug, PartialEq)]
381pub enum SocketReceive {
382    /// A protocol message was received.
383    Message {
384        /// Decoded harness → peer output message.
385        message: HarnessOutputMessage,
386    },
387    /// No message arrived before the requested timeout elapsed.
388    Timeout,
389    /// The socket closed cleanly at a message boundary.
390    Closed,
391}
392
393/// One connected Unix socket peer speaking the protocol.
394///
395/// Each peer owns a bounded background reader thread. Dropping the peer shuts
396/// down the stream, drops the receive queue, and joins that reader thread.
397pub struct SocketPeer {
398    /// Writer for peer/client-to-harness input messages.
399    writer: PeerOutputWriter<BufWriter<UnixStream>>,
400    /// Background reader state, present together until peer shutdown begins.
401    reader_worker: Option<ReaderWorker>,
402    /// Stream clone used to wake the reader thread during peer drop.
403    shutdown_stream: UnixStream,
404}
405
406impl SocketPeer {
407    /// Connects to an existing Unix socket listener.
408    ///
409    /// # Errors
410    ///
411    /// Returns an error when the Unix socket cannot be connected or cloned for
412    /// independent reader/writer ownership.
413    pub fn connect(path: impl Into<PathBuf>) -> Result<Self, SocketTransportError> {
414        let path = path.into();
415        let stream =
416            UnixStream::connect(&path).map_err(|source| SocketTransportError::Connect {
417                path: path.clone(),
418                source,
419            })?;
420        Self::new(stream)
421    }
422
423    /// Connects to an existing Unix socket and bounds subsequent stream I/O.
424    ///
425    /// This is intended for short-lived runtime discovery and control RPCs
426    /// whose caller owns an absolute deadline. The caller must still bound
427    /// how long it waits for a complete protocol response.
428    ///
429    /// # Errors
430    ///
431    /// Returns an error when the socket cannot be connected, its I/O timeouts
432    /// cannot be configured, or it cannot be cloned for split I/O ownership.
433    pub fn connect_with_io_timeout(
434        path: impl Into<PathBuf>,
435        timeout: Duration,
436    ) -> Result<Self, SocketTransportError> {
437        Self::connect_with_timeouts(path, timeout, timeout)
438    }
439
440    /// Connects to an existing Unix socket with separate connect and stream-I/O
441    /// timeouts.
442    ///
443    /// This supports control RPCs that require fast endpoint discovery but
444    /// allow an already-connected peer more time to produce a complete
445    /// response.
446    ///
447    /// # Errors
448    ///
449    /// Returns an error when the socket cannot be connected, its I/O timeouts
450    /// cannot be configured, or it cannot be cloned for split I/O ownership.
451    pub fn connect_with_timeouts(
452        path: impl Into<PathBuf>,
453        connect_timeout: Duration,
454        io_timeout: Duration,
455    ) -> Result<Self, SocketTransportError> {
456        let path = path.into();
457        let stream = connect_unix_with_timeout(&path, connect_timeout).map_err(|source| {
458            SocketTransportError::Connect {
459                path: path.clone(),
460                source,
461            }
462        })?;
463        stream
464            .set_read_timeout(Some(io_timeout))
465            .map_err(|source| SocketTransportError::Connect {
466                path: path.clone(),
467                source,
468            })?;
469        stream
470            .set_write_timeout(Some(io_timeout))
471            .map_err(|source| SocketTransportError::Connect { path, source })?;
472        Self::new(stream)
473    }
474
475    fn new(stream: UnixStream) -> Result<Self, SocketTransportError> {
476        let writer_stream = stream
477            .try_clone()
478            .map_err(|source| SocketTransportError::Clone { source })?;
479        let shutdown_stream = stream
480            .try_clone()
481            .map_err(|source| SocketTransportError::Clone { source })?;
482        let reader_worker = ReaderWorker::spawn(stream)?;
483        Ok(Self {
484            writer: PeerOutputWriter::new(BufWriter::new(writer_stream)),
485            reader_worker: Some(reader_worker),
486            shutdown_stream,
487        })
488    }
489
490    #[cfg(test)]
491    fn new_with_blocked_enqueue_hook(
492        stream: UnixStream,
493        blocked_enqueue: SyncSender<()>,
494    ) -> Result<Self, SocketTransportError> {
495        let writer_stream = stream
496            .try_clone()
497            .map_err(|source| SocketTransportError::Clone { source })?;
498        let shutdown_stream = stream
499            .try_clone()
500            .map_err(|source| SocketTransportError::Clone { source })?;
501        let reader_worker = ReaderWorker::spawn_with_blocked_enqueue_hook(stream, blocked_enqueue)?;
502        Ok(Self {
503            writer: PeerOutputWriter::new(BufWriter::new(writer_stream)),
504            reader_worker: Some(reader_worker),
505            shutdown_stream,
506        })
507    }
508
509    /// Sends one peer → harness protocol message over the Unix socket.
510    ///
511    /// # Errors
512    ///
513    /// Returns an error when encoding or flushing the message fails.
514    pub fn send(&mut self, message: &HarnessInputMessage) -> Result<(), SocketTransportError> {
515        self.writer
516            .write_message(message)
517            .map_err(|source| SocketTransportError::Encode { source })?;
518        self.writer
519            .flush()
520            .map_err(|source| SocketTransportError::Flush { source })
521    }
522
523    /// Updates the write timeout for the next bounded control-plane operation.
524    ///
525    /// # Errors
526    ///
527    /// Returns an error when the connected stream rejects the timeout update.
528    pub fn set_write_timeout(&self, timeout: Duration) -> Result<(), SocketTransportError> {
529        self.writer
530            .get_ref()
531            .get_ref()
532            .set_write_timeout(Some(timeout))
533            .map_err(|source| SocketTransportError::Flush { source })
534    }
535
536    /// Reads one harness → peer protocol message or an explicit timeout/close
537    /// outcome.
538    ///
539    /// # Errors
540    ///
541    /// Returns a decode error for malformed or truncated protocol output.
542    pub fn recv_timeout(
543        &mut self,
544        timeout: Duration,
545    ) -> Result<SocketReceive, SocketTransportError> {
546        let reader_worker = self
547            .reader_worker
548            .as_ref()
549            .expect("socket peer reader missing before drop");
550        match reader_worker.frames.recv_timeout(timeout) {
551            Ok(Ok(frame)) => Ok(SocketReceive::Message { message: frame }),
552            Ok(Err(error)) => Err(SocketTransportError::Decode { source: error }),
553            Err(RecvTimeoutError::Timeout) => Ok(SocketReceive::Timeout),
554            Err(RecvTimeoutError::Disconnected) => Ok(SocketReceive::Closed),
555        }
556    }
557}
558
559fn connect_unix_with_timeout(path: &Path, timeout: Duration) -> io::Result<UnixStream> {
560    let socket = socket2::Socket::new(socket2::Domain::UNIX, socket2::Type::STREAM, None)?;
561    socket.connect_timeout(&socket2::SockAddr::unix(path)?, timeout)?;
562    let fd: OwnedFd = socket.into();
563    Ok(fd.into())
564}
565
566impl Drop for SocketPeer {
567    fn drop(&mut self) {
568        let ReaderWorker { frames, thread } = self
569            .reader_worker
570            .take()
571            .expect("socket peer reader missing during drop");
572        drop(frames);
573        let _ = self.shutdown_stream.shutdown(path_std_net::Shutdown::Both);
574        let _ = thread.join();
575    }
576}
577
578fn create_socket_parent_if_needed(path: &Path) -> Result<(), SocketTransportError> {
579    let Some(parent) = path
580        .parent()
581        .filter(|parent| !parent.as_os_str().is_empty())
582    else {
583        return Ok(());
584    };
585    fs::create_dir_all(parent).map_err(|source| SocketTransportError::CreateParentDirectory {
586        path: parent.to_path_buf(),
587        source,
588    })
589}
590
591fn remove_inactive_stale_socket(path: &Path) -> Result<(), SocketTransportError> {
592    let Ok(metadata) = fs::symlink_metadata(path) else {
593        return Ok(());
594    };
595    if !metadata.file_type().is_socket() {
596        return Err(SocketTransportError::RefuseNonSocketPath {
597            path: path.to_path_buf(),
598        });
599    }
600    match UnixStream::connect(path) {
601        Ok(_) => {
602            return Err(SocketTransportError::ActiveSocketExists {
603                path: path.to_path_buf(),
604            });
605        }
606        Err(error) if error.kind() == io::ErrorKind::ConnectionRefused => {}
607        Err(source) => {
608            return Err(SocketTransportError::ProbeExistingSocket {
609                path: path.to_path_buf(),
610                source,
611            });
612        }
613    }
614    fs::remove_file(path).map_err(|source| SocketTransportError::RemoveStaleSocket {
615        path: path.to_path_buf(),
616        source,
617    })
618}
619
620#[cfg(test)]
621mod tests;