use super::{
credentials::PeerCredentials,
error::{ProtocolError, ProtocolResult},
framing::MessageFraming,
jsonrpc::{
JsonRpcRequest, JsonRpcResponse, JsonRpcErrorResponse, JsonRpcMessage,
JsonRpcNotification,
},
};
use async_trait::async_trait;
use serde_json::Value;
use std::path::Path;
use std::sync::Arc;
use tokio::{
net::{UnixListener, UnixStream},
sync::Semaphore,
};
use tracing::{debug, error, info, warn};
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub socket_path: String,
pub socket_mode: u32,
pub max_connections: usize,
pub allowed_uids: Option<Vec<u32>>,
pub require_root: bool,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
socket_path: "/var/run/crrouterd.sock".to_string(),
socket_mode: 0o666, max_connections: 100,
allowed_uids: None,
require_root: false,
}
}
}
impl ServerConfig {
pub fn from_env() -> Self {
Self {
socket_path: std::env::var("DAEMON_SOCKET_PATH")
.unwrap_or_else(|_| "/var/run/crrouterd.sock".to_string()),
socket_mode: std::env::var("DAEMON_SOCKET_MODE")
.ok()
.and_then(|v| u32::from_str_radix(&v, 8).ok())
.unwrap_or(0o666),
max_connections: std::env::var("DAEMON_MAX_CONNECTIONS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(100),
allowed_uids: None,
require_root: false,
}
}
}
#[async_trait]
pub trait ConnectionHandler: Send + Sync {
async fn handle_request(
&self,
request: JsonRpcRequest,
credentials: PeerCredentials,
) -> ProtocolResult<Value>;
async fn handle_notification(
&self,
notification: JsonRpcNotification,
credentials: PeerCredentials,
) -> ProtocolResult<()> {
debug!(
"Received notification: {} from PID {}",
notification.method, credentials.pid
);
Ok(())
}
async fn on_connect(&self, credentials: PeerCredentials) -> ProtocolResult<()> {
info!(
"Client connected: PID {} (UID {}, GID {})",
credentials.pid, credentials.uid, credentials.gid
);
Ok(())
}
async fn on_disconnect(&self, credentials: PeerCredentials) {
info!(
"Client disconnected: PID {} (UID {}, GID {})",
credentials.pid, credentials.uid, credentials.gid
);
}
}
pub struct UnixServer<H: ConnectionHandler + 'static> {
config: ServerConfig,
handler: Arc<H>,
connection_semaphore: Arc<Semaphore>,
}
impl<H: ConnectionHandler + 'static> UnixServer<H> {
pub fn new(config: ServerConfig, handler: H) -> Self {
let semaphore = Arc::new(Semaphore::new(config.max_connections));
Self {
config,
handler: Arc::new(handler),
connection_semaphore: semaphore,
}
}
pub async fn run(self: Arc<Self>) -> ProtocolResult<()> {
let socket_path = Path::new(&self.config.socket_path);
if socket_path.exists() {
std::fs::remove_file(socket_path).map_err(|e| {
ProtocolError::Internal(format!("Failed to remove existing socket: {}", e))
})?;
}
let listener = UnixListener::bind(socket_path).map_err(|e| {
ProtocolError::Internal(format!("Failed to bind Unix socket: {}", e))
})?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let permissions = std::fs::Permissions::from_mode(self.config.socket_mode);
std::fs::set_permissions(socket_path, permissions).map_err(|e| {
ProtocolError::Internal(format!("Failed to set socket permissions: {}", e))
})?;
}
info!(
"Server listening on {} (mode: {:o})",
self.config.socket_path, self.config.socket_mode
);
loop {
match listener.accept().await {
Ok((stream, _addr)) => {
let server = Arc::clone(&self);
tokio::spawn(async move {
if let Err(e) = server.handle_connection(stream).await {
error!("Connection error: {}", e);
}
});
}
Err(e) => {
error!("Failed to accept connection: {}", e);
}
}
}
}
async fn handle_connection(&self, stream: UnixStream) -> ProtocolResult<()> {
let _permit = self
.connection_semaphore
.acquire()
.await
.map_err(|e| ProtocolError::Internal(format!("Semaphore error: {}", e)))?;
let credentials = PeerCredentials::from_socket(&stream)?;
self.verify_credentials(&stream, &credentials)?;
if let Err(e) = self.handler.on_connect(credentials).await {
warn!("Connection rejected by handler: {}", e);
return Err(e);
}
let result = self.handle_messages(stream, credentials).await;
self.handler.on_disconnect(credentials).await;
result
}
fn verify_credentials<F: std::os::fd::AsFd>(
&self,
_socket: &F,
credentials: &PeerCredentials,
) -> ProtocolResult<()> {
if self.config.require_root && !credentials.is_root() {
return Err(ProtocolError::PermissionDenied(format!(
"Root required (peer UID: {})",
credentials.uid
)));
}
if let Some(ref allowed) = self.config.allowed_uids {
if !allowed.contains(&credentials.uid) {
return Err(ProtocolError::PermissionDenied(format!(
"UID {} not in allowed list",
credentials.uid
)));
}
}
Ok(())
}
async fn handle_messages(
&self,
mut stream: UnixStream,
credentials: PeerCredentials,
) -> ProtocolResult<()> {
loop {
let message: JsonRpcMessage = match MessageFraming::recv_json(&mut stream).await {
Ok(msg) => msg,
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => {
debug!("Connection closed by peer");
break;
}
Err(e) => {
return Err(ProtocolError::Io(e));
}
};
match message {
JsonRpcMessage::Request(request) => {
let response = self.process_request(request, credentials).await;
MessageFraming::send_json(&mut stream, &response)
.await
.map_err(ProtocolError::Io)?;
}
JsonRpcMessage::Notification(notification) => {
if let Err(e) = self.handler.handle_notification(notification, credentials).await {
warn!("Notification handler error: {}", e);
}
}
_ => {
warn!("Unexpected message type from client");
}
}
}
Ok(())
}
async fn process_request(
&self,
request: JsonRpcRequest,
credentials: PeerCredentials,
) -> JsonRpcMessage {
debug!("Processing request: {}", request.method);
match self.handler.handle_request(request.clone(), credentials).await {
Ok(result) => {
JsonRpcMessage::Response(JsonRpcResponse::new(result, request.id))
}
Err(e) => {
error!("Request handler error: {}", e);
JsonRpcMessage::ErrorResponse(JsonRpcErrorResponse::new(
e.to_jsonrpc_error(),
request.id,
))
}
}
}
}
#[cfg(test)]
pub struct EchoHandler;
#[cfg(test)]
#[async_trait]
impl ConnectionHandler for EchoHandler {
async fn handle_request(
&self,
request: JsonRpcRequest,
_credentials: PeerCredentials,
) -> ProtocolResult<Value> {
Ok(request.params.unwrap_or(Value::Null))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_server_config_default() {
let config = ServerConfig::default();
assert_eq!(config.socket_path, "/var/run/crrouterd.sock");
assert_eq!(config.socket_mode, 0o666);
assert_eq!(config.max_connections, 100);
}
#[tokio::test]
async fn test_server_creation() {
let config = ServerConfig {
socket_path: "/tmp/test-server.sock".to_string(),
..Default::default()
};
let handler = EchoHandler;
let server = UnixServer::new(config, handler);
assert_eq!(server.config.socket_path, "/tmp/test-server.sock");
}
}