use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::RwLock;
use crate::audit::{self, AuditEvent, EventType, EventResult};
use crate::handlers;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RemoteConfig {
pub enabled: bool,
pub bind_addr: String,
pub token: Option<String>,
pub max_connections_per_minute: u32,
pub allowed_methods: Vec<String>,
}
impl Default for RemoteConfig {
fn default() -> Self {
Self {
enabled: false,
bind_addr: "127.0.0.1:9473".to_string(),
token: None,
max_connections_per_minute: 10,
allowed_methods: vec![
"status".to_string(),
"list".to_string(),
"run".to_string(),
],
}
}
}
#[derive(Debug, Deserialize)]
pub struct RemoteRequest {
pub token: String,
pub request: serde_json::Value,
}
#[derive(Debug, Serialize)]
pub struct RemoteResponse {
pub success: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub data: Option<serde_json::Value>,
}
impl RemoteResponse {
pub fn ok(data: serde_json::Value) -> Self {
Self {
success: true,
error: None,
data: Some(data),
}
}
pub fn error(msg: impl Into<String>) -> Self {
Self {
success: false,
error: Some(msg.into()),
data: None,
}
}
}
struct RateLimiter {
connections: RwLock<HashMap<String, Vec<Instant>>>,
max_per_minute: u32,
}
impl RateLimiter {
fn new(max_per_minute: u32) -> Self {
Self {
connections: RwLock::new(HashMap::new()),
max_per_minute,
}
}
async fn check(&self, addr: &str) -> bool {
let now = Instant::now();
let minute_ago = now - Duration::from_secs(60);
let mut connections = self.connections.write().await;
let times = connections.entry(addr.to_string()).or_insert_with(Vec::new);
times.retain(|t| *t > minute_ago);
if times.len() >= self.max_per_minute as usize {
return false;
}
times.push(now);
true
}
}
pub struct RemoteListener {
config: RemoteConfig,
rate_limiter: Arc<RateLimiter>,
}
impl RemoteListener {
pub fn new(config: RemoteConfig) -> Self {
let rate_limiter = Arc::new(RateLimiter::new(config.max_connections_per_minute));
Self { config, rate_limiter }
}
pub async fn start(&self) -> Result<(), Box<dyn std::error::Error>> {
if !self.config.enabled {
tracing::info!("Remote listener disabled");
return Ok(());
}
if self.config.token.is_none() {
return Err("Remote listener enabled but no token configured".into());
}
let addr: SocketAddr = self.config.bind_addr.parse()?;
let listener = TcpListener::bind(addr).await?;
tracing::info!("Remote listener started on {}", addr);
audit::log_event(
AuditEvent::new(EventType::DaemonStart, EventResult::Success)
.with_client(&format!("remote:{}", addr))
);
loop {
match listener.accept().await {
Ok((stream, peer_addr)) => {
let peer_str = peer_addr.to_string();
if !self.rate_limiter.check(&peer_str).await {
tracing::warn!("Rate limit exceeded for {}", peer_str);
audit::log_event(
AuditEvent::new(EventType::AuthFailure, EventResult::Failure)
.with_client(&peer_str)
.with_error("Rate limit exceeded")
);
continue;
}
let config = self.config.clone();
tokio::spawn(async move {
if let Err(e) = handle_remote_connection(stream, &peer_str, &config).await {
tracing::error!("Remote connection error from {}: {}", peer_str, e);
}
});
}
Err(e) => {
tracing::error!("Remote accept error: {}", e);
}
}
}
}
}
async fn handle_remote_connection(
stream: TcpStream,
peer: &str,
config: &RemoteConfig,
) -> Result<(), Box<dyn std::error::Error>> {
let (reader, mut writer) = stream.into_split();
let mut reader = BufReader::new(reader);
let mut line = String::new();
audit::log_event(
AuditEvent::new(EventType::ClientConnect, EventResult::Success)
.with_client(peer)
);
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => break, Ok(_) => {
let response = handle_remote_request(&line, peer, config).await;
let json = serde_json::to_string(&response)?;
writer.write_all(format!("{}\n", json).as_bytes()).await?;
}
Err(e) => {
tracing::error!("Remote read error from {}: {}", peer, e);
break;
}
}
}
audit::log_event(
AuditEvent::new(EventType::ClientDisconnect, EventResult::Success)
.with_client(peer)
);
Ok(())
}
async fn handle_remote_request(
line: &str,
peer: &str,
config: &RemoteConfig,
) -> RemoteResponse {
let remote_req: RemoteRequest = match serde_json::from_str(line) {
Ok(r) => r,
Err(e) => {
audit::log_event(
AuditEvent::new(EventType::InvalidRequest, EventResult::Failure)
.with_client(peer)
.with_error(&e.to_string())
);
return RemoteResponse::error(format!("Invalid request format: {}", e));
}
};
let expected_token = config.token.as_ref().unwrap();
if remote_req.token != *expected_token {
audit::log_event(
AuditEvent::new(EventType::AuthFailure, EventResult::Failure)
.with_client(peer)
.with_error("Invalid token")
);
return RemoteResponse::error("Authentication failed");
}
let method = remote_req.request.get("method")
.and_then(|v| v.as_str())
.unwrap_or("unknown");
if !config.allowed_methods.iter().any(|m| m == method) {
audit::log_event(
AuditEvent::new(EventType::InvalidRequest, EventResult::Failure)
.with_client(peer)
.with_error(&format!("Method '{}' not allowed over remote", method))
);
return RemoteResponse::error(format!(
"Method '{}' not allowed over remote connections. Allowed: {:?}",
method, config.allowed_methods
));
}
audit::log_event(
AuditEvent::new(EventType::AuthSuccess, EventResult::Success)
.with_client(peer)
);
let request_json = serde_json::to_string(&remote_req.request).unwrap_or_default();
let local_response = handlers::handle_request_string(&request_json).await;
match serde_json::from_str::<serde_json::Value>(&local_response) {
Ok(value) => {
let success = value.get("success")
.and_then(|v| v.as_bool())
.unwrap_or(false);
if success {
RemoteResponse::ok(value)
} else {
let error = value.get("error")
.and_then(|v| v.as_str())
.unwrap_or("Unknown error");
RemoteResponse::error(error)
}
}
Err(e) => RemoteResponse::error(format!("Internal error: {}", e)),
}
}
pub fn load_config() -> RemoteConfig {
let config_path = get_config_path();
if config_path.exists() {
match std::fs::read_to_string(&config_path) {
Ok(content) => {
match serde_json::from_str(&content) {
Ok(config) => return config,
Err(e) => {
tracing::warn!("Failed to parse remote config: {}", e);
}
}
}
Err(e) => {
tracing::warn!("Failed to read remote config: {}", e);
}
}
}
RemoteConfig::default()
}
fn get_config_path() -> std::path::PathBuf {
dirs::home_dir().unwrap_or_else(std::env::temp_dir)
.join(".scrt4")
.join("remote-config.json")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = RemoteConfig::default();
assert!(!config.enabled);
assert_eq!(config.bind_addr, "127.0.0.1:9473");
assert!(config.token.is_none());
assert!(config.allowed_methods.contains(&"run".to_string()));
assert!(!config.allowed_methods.contains(&"unlock".to_string()));
assert!(!config.allowed_methods.contains(&"reveal".to_string()));
}
#[test]
fn test_remote_response_ok() {
let response = RemoteResponse::ok(serde_json::json!({"test": "value"}));
assert!(response.success);
assert!(response.error.is_none());
assert!(response.data.is_some());
}
#[test]
fn test_remote_response_error() {
let response = RemoteResponse::error("test error");
assert!(!response.success);
assert_eq!(response.error, Some("test error".to_string()));
assert!(response.data.is_none());
}
}