use std::fmt;
use std::path::PathBuf;
use std::str::FromStr;
use rlmesh_grpc::helpers::{BindTarget, parse_bind_target, parse_env_connect_target};
use crate::{Error, Result};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ConnectAddress {
Tcp(String),
Unix(PathBuf),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BindAddress {
Tcp {
host: String,
port: u16,
},
Unix {
path: PathBuf,
},
}
impl ConnectAddress {
pub fn parse(value: impl AsRef<str>) -> Result<Self> {
let target = parse_env_connect_target(value.as_ref()).map_err(Error::from)?;
Ok(match target.unix_path() {
Some(path) => Self::Unix(path.clone()),
None => Self::Tcp(target.display_address().to_string()),
})
}
pub fn as_str(&self) -> String {
match self {
Self::Tcp(value) => value.clone(),
Self::Unix(path) => format!("unix://{}", path.display()),
}
}
}
impl fmt::Display for ConnectAddress {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.as_str().fmt(f)
}
}
impl FromStr for ConnectAddress {
type Err = Error;
fn from_str(value: &str) -> Result<Self> {
Self::parse(value)
}
}
impl BindAddress {
pub fn parse(value: impl AsRef<str>) -> Result<Self> {
Ok(
match parse_bind_target(value.as_ref()).map_err(Error::from)? {
BindTarget::Tcp { host, port } => Self::Tcp { host, port },
BindTarget::Unix { path } => Self::Unix { path },
},
)
}
pub fn display_address(&self) -> String {
match self {
Self::Tcp { host, port } => format!("tcp://{host}:{port}"),
Self::Unix { path } => format!("unix://{}", path.display()),
}
}
}
impl fmt::Display for BindAddress {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.display_address().fmt(f)
}
}
impl FromStr for BindAddress {
type Err = Error;
fn from_str(value: &str) -> Result<Self> {
Self::parse(value)
}
}
#[cfg(unix)]
pub fn remove_stale_socket(path: &std::path::Path) -> Result<()> {
use std::io::ErrorKind;
use std::os::unix::fs::FileTypeExt;
match std::fs::symlink_metadata(path) {
Ok(metadata) if metadata.file_type().is_socket() => {
match std::os::unix::net::UnixStream::connect(path) {
Ok(_) => Err(Error::Server(format!(
"address already in use: a server is already listening on {}",
path.display()
))),
Err(err) if err.kind() == ErrorKind::ConnectionRefused => {
std::fs::remove_file(path).map_err(|err| {
Error::Server(format!(
"failed to remove stale socket {}: {err}",
path.display()
))
})
}
Err(err) if err.kind() == ErrorKind::NotFound => Ok(()),
Err(err) => Err(Error::Server(format!(
"cannot verify whether socket {} is stale: {err}",
path.display()
))),
}
}
_ => Ok(()),
}
}
#[cfg(all(test, unix))]
mod tests {
use super::*;
#[test]
fn refuses_to_unlink_a_live_socket() {
let dir = std::env::temp_dir().join(format!("rlmesh-live-{}.sock", std::process::id()));
let _ = std::fs::remove_file(&dir);
let _listener = std::os::unix::net::UnixListener::bind(&dir).unwrap();
let result = remove_stale_socket(&dir);
assert!(
matches!(result, Err(Error::Server(_))),
"live socket must not be unlinked, got: {result:?}"
);
assert!(dir.exists(), "live socket file must be left in place");
let _ = std::fs::remove_file(&dir);
}
#[test]
fn removes_a_stale_socket_with_no_listener() {
let path = std::env::temp_dir().join(format!("rlmesh-stale-{}.sock", std::process::id()));
let _ = std::fs::remove_file(&path);
{
let _listener = std::os::unix::net::UnixListener::bind(&path).unwrap();
}
assert!(path.exists(), "precondition: stale socket file present");
remove_stale_socket(&path).expect("stale socket should be removed");
assert!(!path.exists(), "stale socket file must be unlinked");
}
}