Files
martillo-maldito/src/ipset.rs
T
midefos 7342386955 fix: ipset compat with v7.x (plain text, timeout support, arg split)
- 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)
2026-08-31 00:49:29 +02:00

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);
}
}