forked from midefos/martillo-maldito
- list() now parses plain text (works on ipset 6/7/8, not just 8) - Strip " timeout N" suffix from member lines (kernel adds it when set has timeout support) - ensure_exists() creates set with 'timeout 0' so add() can use --timeout later - Fix ensure_ipset_match_rule: split rule string into separate argv tokens (Command::args with whitespace string was treated as one arg, breaking nf_tables) - Add #[serial] to tests sharing TEST_SET (race condition fix) - Add 3 new tests: list_parses_plain_format_correctly, list_empty_set_returns_empty_vec, list_filters_non_ip_lines - Simplify integration_netns.sh to use iptables-only assertions (ipset is host-global)
378 lines
11 KiB
Rust
378 lines
11 KiB
Rust
use std::path::Path;
|
|
use std::process::Command;
|
|
|
|
pub const DEFAULT_SET_NAME: &str = "banned";
|
|
pub const DEFAULT_TIMEOUT_SECS: u32 = 3600;
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct Ipset {
|
|
name: String,
|
|
}
|
|
|
|
#[derive(Debug, thiserror::Error)]
|
|
pub enum IpsetError {
|
|
#[error("ipset binary not found in PATH")]
|
|
NotInstalled,
|
|
#[error("ipset command failed: {0}")]
|
|
CommandFailed(String),
|
|
#[error("ipset output parse error: {0}")]
|
|
ParseError(String),
|
|
#[error("invalid IP address: {0}")]
|
|
InvalidIp(String),
|
|
#[error("io error: {0}")]
|
|
Io(#[from] std::io::Error),
|
|
}
|
|
|
|
pub type Result<T> = std::result::Result<T, IpsetError>;
|
|
|
|
impl Ipset {
|
|
pub fn new(name: impl Into<String>) -> Self {
|
|
Self { name: name.into() }
|
|
}
|
|
|
|
pub fn name(&self) -> &str {
|
|
&self.name
|
|
}
|
|
|
|
pub fn is_available() -> bool {
|
|
Command::new("ipset")
|
|
.arg("--version")
|
|
.output()
|
|
.map(|o| o.status.success())
|
|
.unwrap_or(false)
|
|
}
|
|
|
|
pub fn ensure_exists(&self) -> Result<()> {
|
|
let output = Command::new("ipset")
|
|
.args(["create", &self.name, "hash:ip", "timeout", "0", "-exist"])
|
|
.output()?;
|
|
if !output.status.success() {
|
|
return Err(IpsetError::CommandFailed(
|
|
String::from_utf8_lossy(&output.stderr).into_owned(),
|
|
));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
pub fn add(&self, ip: &str, timeout_secs: Option<u32>) -> Result<()> {
|
|
validate_ip(ip)?;
|
|
let mut cmd = Command::new("ipset");
|
|
cmd.args(["add", &self.name, ip, "-exist"]);
|
|
if let Some(secs) = timeout_secs {
|
|
cmd.args(["timeout", &secs.to_string()]);
|
|
}
|
|
let output = cmd.output()?;
|
|
if !output.status.success() {
|
|
return Err(IpsetError::CommandFailed(
|
|
String::from_utf8_lossy(&output.stderr).into_owned(),
|
|
));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
pub fn del(&self, ip: &str) -> Result<()> {
|
|
validate_ip(ip)?;
|
|
let output = Command::new("ipset")
|
|
.args(["del", &self.name, ip, "-exist"])
|
|
.output()?;
|
|
if !output.status.success() {
|
|
return Err(IpsetError::CommandFailed(
|
|
String::from_utf8_lossy(&output.stderr).into_owned(),
|
|
));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
pub fn contains(&self, ip: &str) -> Result<bool> {
|
|
validate_ip(ip)?;
|
|
let output = Command::new("ipset")
|
|
.args(["test", &self.name, ip])
|
|
.output()?;
|
|
Ok(output.status.success())
|
|
}
|
|
|
|
pub fn list(&self) -> Result<Vec<String>> {
|
|
let output = Command::new("ipset")
|
|
.args(["list", &self.name])
|
|
.output()?;
|
|
if !output.status.success() {
|
|
let stderr = String::from_utf8_lossy(&output.stderr).into_owned();
|
|
if stderr.contains("does not exist") {
|
|
return Ok(Vec::new());
|
|
}
|
|
return Err(IpsetError::CommandFailed(stderr));
|
|
}
|
|
let stdout = String::from_utf8_lossy(&output.stdout);
|
|
let mut in_members = false;
|
|
let mut ips = Vec::new();
|
|
for line in stdout.lines() {
|
|
if line.starts_with("Members:") {
|
|
in_members = true;
|
|
continue;
|
|
}
|
|
if in_members && !line.is_empty() {
|
|
let entry = line.split_whitespace().next().unwrap_or("");
|
|
if entry.is_empty() {
|
|
continue;
|
|
}
|
|
if entry.parse::<std::net::Ipv4Addr>().is_ok()
|
|
|| entry.parse::<std::net::Ipv6Addr>().is_ok()
|
|
|| entry.contains('/')
|
|
{
|
|
ips.push(entry.to_string());
|
|
}
|
|
}
|
|
}
|
|
Ok(ips)
|
|
}
|
|
|
|
pub fn save(&self, path: &Path) -> Result<()> {
|
|
let output = Command::new("ipset")
|
|
.args(["save", &self.name])
|
|
.output()?;
|
|
if !output.status.success() {
|
|
return Err(IpsetError::CommandFailed(
|
|
String::from_utf8_lossy(&output.stderr).into_owned(),
|
|
));
|
|
}
|
|
std::fs::write(path, &output.stdout)?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn flush(&self) -> Result<()> {
|
|
let output = Command::new("ipset")
|
|
.args(["flush", &self.name])
|
|
.output()?;
|
|
if !output.status.success() {
|
|
return Err(IpsetError::CommandFailed(
|
|
String::from_utf8_lossy(&output.stderr).into_owned(),
|
|
));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
pub fn destroy(&self) -> Result<()> {
|
|
let output = Command::new("ipset")
|
|
.args(["destroy", &self.name])
|
|
.output()?;
|
|
if !output.status.success() {
|
|
let stderr = String::from_utf8_lossy(&output.stderr).into_owned();
|
|
if stderr.contains("does not exist") {
|
|
return Ok(());
|
|
}
|
|
return Err(IpsetError::CommandFailed(stderr));
|
|
}
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
fn validate_ip(ip: &str) -> Result<()> {
|
|
if ip.parse::<std::net::Ipv4Addr>().is_err() && ip.parse::<std::net::Ipv6Addr>().is_err() {
|
|
return Err(IpsetError::InvalidIp(ip.to_string()));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use serial_test::serial;
|
|
|
|
const TEST_SET: &str = "banned_test_rs";
|
|
|
|
fn setup() -> Option<Ipset> {
|
|
if !Ipset::is_available() {
|
|
eprintln!("ipset binary not available, skipping tests");
|
|
return None;
|
|
}
|
|
let set = Ipset::new(TEST_SET);
|
|
let _ = set.destroy();
|
|
match set.ensure_exists() {
|
|
Ok(_) => Some(set),
|
|
Err(e) => {
|
|
eprintln!("Cannot create ipset (need root?): {}", e);
|
|
None
|
|
}
|
|
}
|
|
}
|
|
|
|
fn teardown(set: &Ipset) {
|
|
let _ = set.destroy();
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn ensure_exists_idempotent() {
|
|
let Some(set) = setup() else { return };
|
|
for _ in 0..5 {
|
|
set.ensure_exists().expect("ensure_exists failed");
|
|
}
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn add_then_list_contains_ip() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("192.0.2.1", None).expect("add failed");
|
|
let ips = set.list().expect("list failed");
|
|
assert!(ips.contains(&"192.0.2.1".to_string()));
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn add_twice_no_error() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("192.0.2.2", None).expect("first add");
|
|
set.add("192.0.2.2", None).expect("second add should be idempotent");
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn del_then_list_empty() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("192.0.2.3", None).expect("add");
|
|
set.del("192.0.2.3").expect("del");
|
|
let ips = set.list().expect("list");
|
|
assert!(!ips.contains(&"192.0.2.3".to_string()));
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn del_nonexistent_no_error() {
|
|
let Some(set) = setup() else { return };
|
|
set.del("192.0.2.4").expect("del nonexistent should not fail");
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn contains_true_after_add_false_after_del() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("192.0.2.5", None).expect("add");
|
|
assert!(set.contains("192.0.2.5").expect("contains"));
|
|
set.del("192.0.2.5").expect("del");
|
|
assert!(!set.contains("192.0.2.5").expect("contains after del"));
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn save_and_restore_roundtrip() {
|
|
use std::env;
|
|
let Some(set) = setup() else { return };
|
|
set.add("192.0.2.10", None).unwrap();
|
|
set.add("192.0.2.11", None).unwrap();
|
|
|
|
let tmp = env::temp_dir().join("martillo_ipset_test.restore");
|
|
set.save(&tmp).expect("save");
|
|
|
|
let content = std::fs::read_to_string(&tmp).expect("read");
|
|
assert!(content.contains("192.0.2.10"));
|
|
assert!(content.contains("192.0.2.11"));
|
|
|
|
let _ = std::fs::remove_file(&tmp);
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn list_parses_plain_format_correctly() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("198.51.100.1", None).unwrap();
|
|
set.add("198.51.100.2", None).unwrap();
|
|
set.add("198.51.100.3", None).unwrap();
|
|
let ips = set.list().unwrap();
|
|
assert_eq!(ips.len(), 3);
|
|
assert!(ips.contains(&"198.51.100.1".to_string()));
|
|
assert!(ips.contains(&"198.51.100.2".to_string()));
|
|
assert!(ips.contains(&"198.51.100.3".to_string()));
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn list_empty_set_returns_empty_vec() {
|
|
let Some(set) = setup() else { return };
|
|
let ips = set.list().unwrap();
|
|
assert_eq!(ips.len(), 0);
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn list_filters_non_ip_lines() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("203.0.113.1", None).unwrap();
|
|
set.add("203.0.113.2", None).unwrap();
|
|
let ips = set.list().unwrap();
|
|
for ip in &ips {
|
|
assert!(
|
|
ip.parse::<std::net::Ipv4Addr>().is_ok()
|
|
|| ip.parse::<std::net::Ipv6Addr>().is_ok()
|
|
|| ip.contains('/'),
|
|
"got non-IP entry: {ip}"
|
|
);
|
|
}
|
|
assert_eq!(ips.len(), 2);
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
fn invalid_ip_rejected() {
|
|
let result = validate_ip("not-an-ip");
|
|
assert!(matches!(result, Err(IpsetError::InvalidIp(_))));
|
|
assert!(validate_ip("1.2.3.4").is_ok());
|
|
assert!(validate_ip("::1").is_ok());
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn flush_clears_all() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("203.0.113.1", None).unwrap();
|
|
set.add("203.0.113.2", None).unwrap();
|
|
assert_eq!(set.list().unwrap().len(), 2);
|
|
set.flush().unwrap();
|
|
assert_eq!(set.list().unwrap().len(), 0);
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn add_with_timeout() {
|
|
let Some(set) = setup() else { return };
|
|
set.add("203.0.113.10", Some(2)).expect("add with timeout");
|
|
assert!(set.contains("203.0.113.10").unwrap());
|
|
std::thread::sleep(std::time::Duration::from_secs(3));
|
|
assert!(!set.contains("203.0.113.10").unwrap(), "timeout should expire");
|
|
teardown(&set);
|
|
}
|
|
|
|
#[test]
|
|
fn list_on_nonexistent_set_returns_empty() {
|
|
let bogus = Ipset::new("nonexistent_set_zzz_999");
|
|
let result = bogus.list();
|
|
assert!(result.is_ok(), "list on nonexistent should not error");
|
|
assert_eq!(result.unwrap().len(), 0);
|
|
}
|
|
|
|
#[test]
|
|
#[serial]
|
|
fn bulk_1000_ips_performance() {
|
|
let Some(set) = setup() else { return };
|
|
let start = std::time::Instant::now();
|
|
for i in 0..1000u32 {
|
|
let ip = format!("10.99.{}.{}", i / 256, i % 256);
|
|
set.add(&ip, None).expect("add");
|
|
}
|
|
let elapsed = start.elapsed();
|
|
assert!(elapsed.as_secs() < 5, "1000 adds took {:?}", elapsed);
|
|
assert_eq!(set.list().unwrap().len(), 1000);
|
|
teardown(&set);
|
|
}
|
|
}
|