1use 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#[derive(Debug)]
28pub enum SocketTransportError {
29 CreateParentDirectory {
31 path: PathBuf,
33 source: io::Error,
35 },
36 RefuseNonSocketPath {
38 path: PathBuf,
40 },
41 ActiveSocketExists {
43 path: PathBuf,
45 },
46 ProbeExistingSocket {
48 path: PathBuf,
50 source: io::Error,
52 },
53 RemoveStaleSocket {
55 path: PathBuf,
57 source: io::Error,
59 },
60 Bind {
62 path: PathBuf,
64 source: io::Error,
66 },
67 BoundSocketMetadata {
69 path: PathBuf,
71 source: io::Error,
73 },
74 Accept {
76 source: io::Error,
78 },
79 Connect {
81 path: PathBuf,
83 source: io::Error,
85 },
86 Clone {
88 source: io::Error,
90 },
91 SpawnReader {
93 source: io::Error,
95 },
96 Encode {
98 source: tau_proto::EncodeError,
100 },
101 Flush {
103 source: io::Error,
105 },
106 Decode {
108 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 dev: u64,
192 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
205pub struct SocketListener {
210 path: PathBuf,
212 listener: UnixListener,
214 socket_identity: SocketIdentity,
216}
217
218impl SocketListener {
219 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 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 #[must_use]
279 pub fn path(&self) -> &Path {
280 &self.path
281 }
282
283 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 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
331pub struct SocketAcceptedClient {
333 reader: HarnessInputReader<UnixStream>,
335 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 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 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#[derive(Debug, PartialEq)]
381pub enum SocketReceive {
382 Message {
384 message: HarnessOutputMessage,
386 },
387 Timeout,
389 Closed,
391}
392
393pub struct SocketPeer {
398 writer: PeerOutputWriter<BufWriter<UnixStream>>,
400 reader_worker: Option<ReaderWorker>,
402 shutdown_stream: UnixStream,
404}
405
406impl SocketPeer {
407 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 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 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 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 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 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;