use nix::sys::socket::{getsockopt, sockopt::PeerCredentials as NixPeerCred};
use std::os::fd::AsFd;
use super::error::{ProtocolError, ProtocolResult};
#[cfg(test)]
use nix::unistd::{getuid, getgid};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PeerCredentials {
pub pid: i32,
pub uid: u32,
pub gid: u32,
}
impl PeerCredentials {
pub fn from_socket<F: AsFd>(socket: &F) -> ProtocolResult<Self> {
let creds = getsockopt(socket, NixPeerCred)
.map_err(|e| ProtocolError::Internal(format!("Failed to get peer credentials: {}", e)))?;
Ok(Self {
pid: creds.pid(),
uid: creds.uid(),
gid: creds.gid(),
})
}
pub fn is_root(&self) -> bool {
self.uid == 0
}
pub fn has_uid(&self, uid: u32) -> bool {
self.uid == uid
}
pub fn has_gid(&self, gid: u32) -> bool {
self.gid == gid
}
#[cfg(target_os = "linux")]
pub fn process_info(&self) -> Option<ProcessInfo> {
ProcessInfo::from_pid(self.pid)
}
}
#[cfg(target_os = "linux")]
#[derive(Debug, Clone)]
pub struct ProcessInfo {
pub pid: i32,
pub name: String,
pub exe_path: Option<String>,
}
#[cfg(target_os = "linux")]
impl ProcessInfo {
pub fn from_pid(pid: i32) -> Option<Self> {
use std::fs;
let comm_path = format!("/proc/{}/comm", pid);
let name = fs::read_to_string(comm_path).ok()?.trim().to_string();
let exe_path = fs::read_link(format!("/proc/{}/exe", pid))
.ok()
.and_then(|p| p.to_str().map(String::from));
Some(Self {
pid,
name,
exe_path,
})
}
}
pub fn verify_peer_credentials<F: AsFd>(
socket: &F,
expected_uid: Option<u32>,
expected_gid: Option<u32>,
) -> ProtocolResult<PeerCredentials> {
let creds = PeerCredentials::from_socket(socket)?;
if let Some(uid) = expected_uid {
if creds.uid != uid {
return Err(ProtocolError::PermissionDenied(format!(
"Expected UID {}, got {}",
uid, creds.uid
)));
}
}
if let Some(gid) = expected_gid {
if creds.gid != gid {
return Err(ProtocolError::PermissionDenied(format!(
"Expected GID {}, got {}",
gid, creds.gid
)));
}
}
Ok(creds)
}
pub fn require_root<F: AsFd>(socket: &F) -> ProtocolResult<PeerCredentials> {
let creds = PeerCredentials::from_socket(socket)?;
if !creds.is_root() {
return Err(ProtocolError::PermissionDenied(format!(
"Root access required (peer UID: {})",
creds.uid
)));
}
Ok(creds)
}
pub fn require_user<F: AsFd>(socket: &F, uid: u32) -> ProtocolResult<PeerCredentials> {
verify_peer_credentials(socket, Some(uid), None)
}
pub fn allow_uids<F: AsFd>(socket: &F, allowed_uids: &[u32]) -> ProtocolResult<PeerCredentials> {
let creds = PeerCredentials::from_socket(socket)?;
if !allowed_uids.contains(&creds.uid) {
return Err(ProtocolError::PermissionDenied(format!(
"UID {} not in allowed list: {:?}",
creds.uid, allowed_uids
)));
}
Ok(creds)
}
#[cfg(test)]
mod tests {
use super::*;
use std::os::unix::net::UnixStream;
#[test]
fn test_peer_credentials() {
let (sock1, sock2) = UnixStream::pair().unwrap();
let creds = PeerCredentials::from_socket(&sock1).unwrap();
assert_eq!(creds.pid, std::process::id() as i32);
assert_eq!(creds.uid, getuid().as_raw());
assert_eq!(creds.gid, getgid().as_raw());
let creds2 = PeerCredentials::from_socket(&sock2).unwrap();
assert_eq!(creds, creds2);
}
#[test]
fn test_is_root() {
let (sock, _) = UnixStream::pair().unwrap();
let creds = PeerCredentials::from_socket(&sock).unwrap();
assert!(!creds.is_root());
}
#[test]
fn test_verify_credentials() {
let (sock, _) = UnixStream::pair().unwrap();
let my_uid = getuid().as_raw();
let result = verify_peer_credentials(&sock, Some(my_uid), None);
assert!(result.is_ok());
let result = verify_peer_credentials(&sock, Some(0), None);
assert!(result.is_err());
}
#[test]
fn test_allow_uids() {
let (sock, _) = UnixStream::pair().unwrap();
let my_uid = getuid().as_raw();
let result = allow_uids(&sock, &[my_uid, 1000, 1001]);
assert!(result.is_ok());
let other_uids: Vec<u32> = (0..3).map(|i| my_uid + 1000 + i).collect();
let result = allow_uids(&sock, &other_uids);
assert!(result.is_err());
}
#[cfg(target_os = "linux")]
#[test]
fn test_process_info() {
let (sock, _) = UnixStream::pair().unwrap();
let creds = PeerCredentials::from_socket(&sock).unwrap();
if let Some(info) = creds.process_info() {
assert_eq!(info.pid, creds.pid);
assert!(!info.name.is_empty());
println!("Process: {} (PID: {})", info.name, info.pid);
}
}
}