use std::error::Error;
use std::ffi::OsString;
use std::fmt;
use std::path::{Path, PathBuf};
pub const MAX_SOCKET_PATH_BYTES: usize = SUN_PATH_BYTES - 1;
#[cfg(any(
target_vendor = "apple",
target_os = "freebsd",
target_os = "netbsd",
target_os = "openbsd",
target_os = "dragonfly"
))]
const SUN_PATH_BYTES: usize = 104;
#[cfg(not(any(
target_vendor = "apple",
target_os = "freebsd",
target_os = "netbsd",
target_os = "openbsd",
target_os = "dragonfly"
)))]
const SUN_PATH_BYTES: usize = 108;
pub fn resolve_path(
taria_sock: Option<OsString>,
xdg_runtime_dir: Option<OsString>,
temp_dir: &Path,
user: &str,
app_label: &str,
) -> Result<PathBuf, InvalidAppLabel> {
if !is_plain_file_name(app_label) {
return Err(InvalidAppLabel {
label: app_label.to_string(),
});
}
if let Some(path) = taria_sock
&& !path.is_empty()
{
return Ok(PathBuf::from(path));
}
if let Some(dir) = xdg_runtime_dir
&& !dir.is_empty()
{
return Ok(PathBuf::from(dir)
.join("taria")
.join(format!("{app_label}.sock")));
}
Ok(temp_dir
.join(format!("taria-{user}"))
.join(format!("{app_label}.sock")))
}
fn is_plain_file_name(label: &str) -> bool {
!label.is_empty()
&& label != "."
&& label != ".."
&& !label.contains('/')
&& !label.contains('\0')
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InvalidAppLabel {
label: String,
}
impl InvalidAppLabel {
pub fn label(&self) -> &str {
&self.label
}
}
impl fmt::Display for InvalidAppLabel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"app label {:?} is not a file name; it must be a single path component, so no `/`, \
and not `.`, `..` or empty. Pass the socket path itself to use a path of your own.",
self.label,
)
}
}
impl Error for InvalidAppLabel {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SocketPathTooLong {
path: PathBuf,
len: usize,
}
impl SocketPathTooLong {
pub fn path(&self) -> &Path {
&self.path
}
pub fn path_len(&self) -> usize {
self.len
}
}
impl fmt::Display for SocketPathTooLong {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"socket path {} is {} bytes, over the {MAX_SOCKET_PATH_BYTES} byte limit for Unix \
sockets on {}; set $TARIA_SOCK to a shorter path inside a directory only you can \
reach, for example $XDG_RUNTIME_DIR/t.sock, or ~/.taria/t.sock where that variable \
is unset. A shared directory such as /tmp will not work: the adapter binds only \
under a directory you own that grants no group or other access, and creates a \
missing one with mode 0700",
self.path.display(),
self.len,
std::env::consts::OS,
)
}
}
impl Error for SocketPathTooLong {}
pub fn check_path_len(path: &Path) -> Result<(), SocketPathTooLong> {
let len = path.as_os_str().len();
if len > MAX_SOCKET_PATH_BYTES {
return Err(SocketPathTooLong {
path: path.to_path_buf(),
len,
});
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn resolved(
taria_sock: Option<OsString>,
xdg_runtime_dir: Option<OsString>,
label: &str,
) -> PathBuf {
resolve_path(
taria_sock,
xdg_runtime_dir,
Path::new("/tmp"),
"1000",
label,
)
.expect("a plain label resolves")
}
#[test]
fn taria_sock_env_wins() {
let path = resolved(
Some("/custom/app.sock".into()),
Some("/run/user/1000".into()),
"demo",
);
assert_eq!(path, PathBuf::from("/custom/app.sock"));
}
#[test]
fn empty_taria_sock_is_ignored() {
let path = resolved(Some("".into()), Some("/run/user/1000".into()), "demo");
assert_eq!(path, PathBuf::from("/run/user/1000/taria/demo.sock"));
}
#[test]
fn xdg_runtime_dir_is_second_choice() {
let path = resolved(None, Some("/run/user/1000".into()), "demo");
assert_eq!(path, PathBuf::from("/run/user/1000/taria/demo.sock"));
}
#[test]
fn temp_dir_is_the_fallback() {
let cases = [
(None, None),
(None, Some(OsString::from(""))),
(Some(OsString::from("")), Some(OsString::from(""))),
];
for (sock, xdg) in cases {
let path = resolved(sock.clone(), xdg.clone(), "demo");
assert_eq!(
path,
PathBuf::from("/tmp/taria-1000/demo.sock"),
"sock: {sock:?}, xdg: {xdg:?}"
);
}
}
#[test]
fn a_label_that_is_not_a_file_name_is_refused() {
for label in [
"/etc/cron.d/evil",
"../../../tmp/pwn",
"sub/dir",
"app/",
".",
"..",
"",
"nul\0byte",
] {
let err = resolve_path(
None,
Some("/run/user/1000".into()),
Path::new("/tmp"),
"1000",
label,
)
.unwrap_err();
assert_eq!(err.label(), label);
}
}
#[test]
fn a_bad_label_is_refused_even_when_the_override_would_hide_it() {
assert!(
resolve_path(
Some("/custom/app.sock".into()),
None,
Path::new("/tmp"),
"1000",
"../pwn",
)
.is_err()
);
}
#[test]
fn plain_labels_with_awkward_characters_still_resolve() {
for label in ["my app", "app.v2", "..app", "app..", "-", "app:1"] {
let path = resolved(None, Some("/run/user/1000".into()), label);
assert_eq!(
path,
PathBuf::from(format!("/run/user/1000/taria/{label}.sock")),
"label: {label}"
);
}
}
#[test]
fn the_error_names_the_label_and_what_a_label_may_be() {
let err = resolve_path(None, None, Path::new("/tmp"), "1000", "../pwn").unwrap_err();
let message = err.to_string();
assert!(message.contains("../pwn"), "message: {message}");
assert!(
message.contains("single path component"),
"message: {message}"
);
}
#[test]
fn path_length_boundary_is_inclusive() {
let at_limit = PathBuf::from("/".repeat(MAX_SOCKET_PATH_BYTES));
assert_eq!(at_limit.as_os_str().len(), MAX_SOCKET_PATH_BYTES);
assert_eq!(check_path_len(&at_limit), Ok(()));
let over_limit = PathBuf::from("/".repeat(MAX_SOCKET_PATH_BYTES + 1));
let err = check_path_len(&over_limit).unwrap_err();
assert_eq!(err.path_len(), MAX_SOCKET_PATH_BYTES + 1);
assert_eq!(err.path(), over_limit);
}
#[test]
fn the_limit_is_the_platforms_sun_path_minus_the_nul() {
#[cfg(any(
target_vendor = "apple",
target_os = "freebsd",
target_os = "netbsd",
target_os = "openbsd",
target_os = "dragonfly"
))]
assert_eq!(
MAX_SOCKET_PATH_BYTES, 103,
"this platform declares sun_path[104]"
);
#[cfg(not(any(
target_vendor = "apple",
target_os = "freebsd",
target_os = "netbsd",
target_os = "openbsd",
target_os = "dragonfly"
)))]
assert_eq!(
MAX_SOCKET_PATH_BYTES, 107,
"this platform declares sun_path[108]"
);
}
#[test]
fn the_band_between_the_two_sun_path_sizes_follows_the_platform() {
for len in 104..=107 {
let path = PathBuf::from("/".repeat(len));
let checked = check_path_len(&path);
if MAX_SOCKET_PATH_BYTES == 103 {
assert!(checked.is_err(), "{len} bytes does not fit sun_path[104]");
} else {
assert_eq!(checked, Ok(()), "{len} bytes fits sun_path[108]");
}
}
}
#[test]
fn too_long_error_names_path_length_limit_and_platform() {
let path = PathBuf::from(format!("/tmp/{}.sock", "x".repeat(120)));
let err = check_path_len(&path).unwrap_err();
let message = err.to_string();
assert!(message.contains("xxxx"), "message: {message}");
assert!(message.contains("130 bytes"), "message: {message}");
assert!(
message.contains(&format!("{MAX_SOCKET_PATH_BYTES} byte")),
"message: {message}"
);
assert!(
message.contains(std::env::consts::OS),
"the limit is per platform, so the message must name this one: {message}"
);
assert!(message.contains("TARIA_SOCK"), "message: {message}");
}
}