feat: velocity routing, server registry and RampartVelocity tests

This commit is contained in:
loki5512344 2026-08-03 10:54:28 +02:00
parent fa6de281fb
commit 29ce6fb8e9
Signed by: boba
GPG key ID: 253067914055423B
41 changed files with 1239 additions and 532 deletions

View file

@ -14,4 +14,4 @@ serde_json.workspace = true
tracing.workspace = true
anyhow.workspace = true
clap.workspace = true
reqwest = { version = "0.12", features = ["json"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }

View file

@ -65,17 +65,24 @@ pub struct HmacConfig {
pub secret: String,
#[serde(default = "default_key_rotation")]
pub key_rotation_interval_secs: u64,
#[serde(default = "default_signature_ttl")]
pub signature_ttl_secs: u64,
}
fn default_key_rotation() -> u64 {
3600
}
fn default_signature_ttl() -> u64 {
60
}
impl Default for HmacConfig {
fn default() -> Self {
Self {
secret: String::new(),
key_rotation_interval_secs: 3600,
signature_ttl_secs: 60,
}
}
}
@ -262,7 +269,12 @@ pub struct PowConfig {
}
fn default_pow_enabled() -> bool {
true
// PoW выключен по умолчанию: текущий текстовый challenge отправляется до
// handshake и несовместим с ванильными MC-клиентами — они не умеют его
// решать, поэтому при enabled=true никто не сможет зайти на сервер.
// Включать только после появления клиентского мода или PoW, совместимого
// с протоколом Minecraft.
false
}
fn default_pow_difficulty() -> u8 {
@ -272,7 +284,7 @@ fn default_pow_difficulty() -> u8 {
impl Default for PowConfig {
fn default() -> Self {
Self {
enabled: true,
enabled: false,
difficulty: 4,
}
}

View file

@ -1,73 +1,153 @@
use hmac::{Hmac, Mac};
use sha2::Sha256;
use std::time::{SystemTime, UNIX_EPOCH};
use subtle::ConstantTimeEq;
type HmacSha256 = Hmac<Sha256>;
pub fn sign(hostname: &str, secret: &[u8]) -> String {
let mut mac = HmacSha256::new_from_slice(secret).expect("HMAC accepts any key length");
mac.update(hostname.as_bytes());
fn now_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
fn hmac_hex(key: &[u8], message: &[u8]) -> String {
let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts any key length");
mac.update(message);
hex::encode(mac.finalize().into_bytes())
}
pub fn verify(hostname: &str, provided_sig: &str, secret: &[u8]) -> bool {
let expected = sign(hostname, secret);
expected.as_bytes().ct_eq(provided_sig.as_bytes()).into()
fn derive_key(secret: &[u8], bucket: u64) -> Vec<u8> {
let msg = format!("rampart-key-{bucket}");
let mut mac = HmacSha256::new_from_slice(secret).expect("HMAC accepts any key length");
mac.update(msg.as_bytes());
mac.finalize().into_bytes().to_vec()
}
pub fn sign_hostname(raw: &str, secret: &[u8]) -> String {
/// Подписывает hostname-поле: `domain\0shield\0<ts>\0<sig>`.
///
/// `sig` = HMAC-SHA256(derived_key, "domain|ts"), где
/// `derived_key` = HMAC-SHA256(master_secret, "rampart-key-{bucket}"), bucket = ts / rotation_secs.
pub fn sign_hostname(raw: &str, secret: &[u8], rotation_secs: u64) -> String {
let rotation_secs = rotation_secs.max(1);
let domain = raw.split('\0').next().unwrap_or(raw);
let sig = sign(domain, secret);
format!("{raw}\0shield\0{sig}")
let ts = now_secs();
let bucket = ts / rotation_secs;
let key = derive_key(secret, bucket);
let sig = hmac_hex(&key, format!("{domain}|{ts}").as_bytes());
format!("{domain}\0shield\0{ts}\0{sig}")
}
pub fn parse_hostname(raw: &str) -> (String, Option<String>) {
let parts: Vec<&str> = raw.split('\0').collect();
let domain = parts[0].to_string();
let hmac = parts
.iter()
.position(|&p| p == "shield")
.and_then(|i| parts.get(i + 1))
.map(|s| s.to_string());
(domain, hmac)
/// Проверяет подпись hostname-поля по спецификации.
///
/// Парсит `domain\0shield\0<ts>\0<sig>`, проверяет `0 <= now - ts <= ttl_secs` и
/// сравнивает сигнатуру constant-time для bucket из `{ts_bucket, ts_bucket - 1}`.
pub fn verify_hostname(raw: &str, secret: &[u8], rotation_secs: u64, ttl_secs: u64) -> bool {
let rotation_secs = rotation_secs.max(1);
let mut parts = raw.split('\0');
let (Some(domain), Some(tag), Some(ts_str), Some(sig)) = (parts.next(), parts.next(), parts.next(), parts.next())
else {
return false;
};
if tag != "shield" || parts.next().is_some() {
return false;
}
let ts: u64 = match ts_str.parse() {
Ok(t) => t,
Err(_) => return false,
};
let now = now_secs();
if now < ts || now - ts > ttl_secs {
return false;
}
if sig.len() != 64 {
return false;
}
let bucket = ts / rotation_secs;
for candidate in [bucket, bucket.saturating_sub(1)] {
let key = derive_key(secret, candidate);
let expected = hmac_hex(&key, format!("{domain}|{ts}").as_bytes());
if expected.as_bytes().ct_eq(sig.as_bytes()).into() {
return true;
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
const SECRET: &[u8] = b"test_secret_32_bytes_long_here!!";
fn build_signed(secret: &[u8], domain: &str, ts: u64, rotation_secs: u64) -> String {
let bucket = ts / rotation_secs.max(1);
let key = derive_key(secret, bucket);
let sig = hmac_hex(&key, format!("{domain}|{ts}").as_bytes());
format!("{domain}\0shield\0{ts}\0{sig}")
}
#[test]
fn test_sign_verify() {
let secret = b"test_secret_32_bytes_long_here!!";
let hostname = "play.example.com";
let sig = sign(hostname, secret);
assert!(verify(hostname, &sig, secret));
fn test_sign_verify_roundtrip() {
let signed = sign_hostname("play.example.com", SECRET, 3600);
assert!(verify_hostname(&signed, SECRET, 3600, 60));
}
#[test]
fn test_sign_format() {
let signed = sign_hostname("play.example.com\0ignored", SECRET, 3600);
let parts: Vec<&str> = signed.split('\0').collect();
assert_eq!(parts.len(), 4);
assert_eq!(parts[0], "play.example.com");
assert_eq!(parts[1], "shield");
assert_eq!(parts[3].len(), 64);
assert!(parts[3].chars().all(|c| c.is_ascii_hexdigit()));
}
#[test]
fn test_verify_tampered_domain() {
let signed = sign_hostname("play.example.com", SECRET, 3600);
let tampered = signed.replace("play.example.com", "play.example.co");
assert!(!verify_hostname(&tampered, SECRET, 3600, 60));
}
#[test]
fn test_verify_wrong_secret() {
let secret = b"test_secret_32_bytes_long_here!!";
let signed = sign_hostname("play.example.com", SECRET, 3600);
let wrong = b"wrong_secret_32_bytes_long_here!!!";
let hostname = "play.example.com";
let sig = sign(hostname, wrong);
assert!(!verify(hostname, &sig, secret));
assert!(!verify_hostname(&signed, wrong, 3600, 60));
}
#[test]
fn test_sign_hostname_suffix() {
let secret = b"test_secret";
let result = sign_hostname("play.example.com", secret);
assert!(result.starts_with("play.example.com\0shield\0"));
let sig = result.split("\0shield\0").nth(1).unwrap();
assert_eq!(sig.len(), 64);
fn test_verify_expired_ts() {
let old_ts = now_secs().saturating_sub(120);
let signed = build_signed(SECRET, "play.example.com", old_ts, 3600);
assert!(!verify_hostname(&signed, SECRET, 3600, 60));
}
#[test]
fn test_verify_constant_time() {
let secret = b"test_secret_32_bytes_long_here!!";
let hostname = "play.example.com";
let sig = sign(hostname, secret);
assert!(!verify("play.example.co", &sig, secret));
assert!(verify(hostname, &sig, secret));
fn test_verify_accepts_previous_bucket() {
let rotation = 10u64;
let now = now_secs();
let prev_bucket_ts = (now / rotation).saturating_sub(1) * rotation + 5;
let signed = build_signed(SECRET, "play.example.com", prev_bucket_ts, rotation);
assert!(verify_hostname(&signed, SECRET, rotation, 60));
}
#[test]
fn test_verify_rejects_tampered_ts() {
let ts = now_secs();
let signed = build_signed(SECRET, "play.example.com", ts, 3600);
let parts: Vec<&str> = signed.split('\0').collect();
let tampered = format!("{}\0{}\0{}\0{}", parts[0], parts[1], ts.saturating_sub(1), parts[3]);
assert!(!verify_hostname(&tampered, SECRET, 3600, 60));
}
#[test]
fn test_verify_garbage_input() {
assert!(!verify_hostname("", SECRET, 3600, 60));
assert!(!verify_hostname("no-separators", SECRET, 3600, 60));
assert!(!verify_hostname("a\0shield\0bad\0short", SECRET, 3600, 60));
}
}

View file

@ -1,4 +1,5 @@
use dashmap::DashMap;
use std::net::IpAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
@ -8,7 +9,7 @@ struct BanEntry {
}
pub struct Blacklist {
entries: Arc<DashMap<u32, BanEntry>>,
entries: Arc<DashMap<IpAddr, BanEntry>>,
}
impl Default for Blacklist {
@ -24,7 +25,7 @@ impl Blacklist {
}
}
pub fn is_blocked(&self, ip: u32) -> bool {
pub fn is_blocked(&self, ip: IpAddr) -> bool {
if let Some(entry) = self.entries.get(&ip) {
if entry.expires > Instant::now() {
return true;
@ -35,7 +36,7 @@ impl Blacklist {
false
}
pub fn add(&self, ip: u32, duration: Duration, reason: &str) {
pub fn add(&self, ip: IpAddr, duration: Duration, reason: &str) {
self.entries.insert(
ip,
BanEntry {
@ -45,7 +46,7 @@ impl Blacklist {
);
}
pub fn remove(&self, ip: u32) {
pub fn remove(&self, ip: IpAddr) {
self.entries.remove(&ip);
}
@ -65,34 +66,39 @@ impl Blacklist {
#[cfg(test)]
mod tests {
use super::*;
use std::net::{IpAddr, Ipv4Addr};
fn ip(octets: [u8; 4]) -> IpAddr {
IpAddr::V4(Ipv4Addr::from(octets))
}
#[test]
fn test_blacklist_block() {
let bl = Blacklist::new();
bl.add(0x01020304, Duration::from_secs(60), "test");
assert!(bl.is_blocked(0x01020304));
bl.add(ip([1, 2, 3, 4]), Duration::from_secs(60), "test");
assert!(bl.is_blocked(ip([1, 2, 3, 4])));
}
#[test]
fn test_blacklist_not_blocked() {
let bl = Blacklist::new();
bl.add(0x01020304, Duration::from_secs(60), "test");
assert!(!bl.is_blocked(0x05060708));
bl.add(ip([1, 2, 3, 4]), Duration::from_secs(60), "test");
assert!(!bl.is_blocked(ip([5, 6, 7, 8])));
}
#[test]
fn test_blacklist_expired() {
let bl = Blacklist::new();
bl.add(0x01020304, Duration::from_millis(1), "test");
bl.add(ip([1, 2, 3, 4]), Duration::from_millis(1), "test");
std::thread::sleep(Duration::from_millis(2));
assert!(!bl.is_blocked(0x01020304));
assert!(!bl.is_blocked(ip([1, 2, 3, 4])));
}
#[test]
fn test_blacklist_remove() {
let bl = Blacklist::new();
bl.add(0x01020304, Duration::from_secs(60), "test");
bl.remove(0x01020304);
assert!(!bl.is_blocked(0x01020304));
bl.add(ip([1, 2, 3, 4]), Duration::from_secs(60), "test");
bl.remove(ip([1, 2, 3, 4]));
assert!(!bl.is_blocked(ip([1, 2, 3, 4])));
}
}

View file

@ -1,40 +1,59 @@
use dashmap::DashMap;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
const MAX_BUCKETS: usize = 1_000_000;
const EVICTION_IDLE: Duration = Duration::from_secs(600);
const SWEEP_INTERVAL: Duration = Duration::from_secs(60);
struct Bucket {
tokens: f64,
last_refill: Instant,
last_access: Instant,
}
pub struct RateLimiter {
buckets: Arc<DashMap<u32, Bucket>>,
buckets: Arc<DashMap<IpAddr, Bucket>>,
max_tokens: f64,
refill_rate: f64,
_refill_interval: Duration,
epoch: Instant,
last_sweep_elapsed: AtomicU64,
eviction_idle: Duration,
}
impl RateLimiter {
pub fn new(rate_per_sec: f64, burst: f64) -> Self {
Self::with_eviction(rate_per_sec, burst, EVICTION_IDLE)
}
fn with_eviction(rate_per_sec: f64, burst: f64, eviction_idle: Duration) -> Self {
Self {
buckets: Arc::new(DashMap::new()),
max_tokens: burst,
refill_rate: rate_per_sec,
_refill_interval: Duration::from_secs(1),
epoch: Instant::now(),
last_sweep_elapsed: AtomicU64::new(0),
eviction_idle,
}
}
pub fn check(&self, ip: u32) -> bool {
pub fn check(&self, ip: IpAddr) -> bool {
let now = Instant::now();
let mut entry = self.buckets.entry(ip).or_insert_with(|| Bucket {
tokens: self.max_tokens,
last_refill: Instant::now(),
last_refill: now,
last_access: now,
});
let now = Instant::now();
let elapsed = now.duration_since(entry.last_refill);
let refill = elapsed.as_secs_f64() * self.refill_rate;
entry.tokens = (entry.tokens + refill).min(self.max_tokens);
entry.last_refill = now;
entry.last_access = now;
if entry.tokens >= 1.0 {
entry.tokens -= 1.0;
@ -51,40 +70,83 @@ impl RateLimiter {
pub fn is_empty(&self) -> bool {
self.buckets.is_empty()
}
/// Эвиктит простаивающие бакеты. Запускается, когда бакетов больше
/// MAX_BUCKETS либо по расписанию (раз в SWEEP_INTERVAL).
pub fn sweep(&self) {
let now = Instant::now();
let elapsed_secs = now.duration_since(self.epoch).as_secs();
let last = self.last_sweep_elapsed.load(Ordering::Relaxed);
let due = last == 0 || elapsed_secs.saturating_sub(last) >= SWEEP_INTERVAL.as_secs();
let over_cap = self.buckets.len() > MAX_BUCKETS;
if !over_cap && !due {
return;
}
self.buckets
.retain(|_, b| now.duration_since(b.last_access) < self.eviction_idle);
self.last_sweep_elapsed.store(elapsed_secs, Ordering::Relaxed);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{IpAddr, Ipv4Addr};
fn test_ip(octet: u8) -> IpAddr {
IpAddr::V4(Ipv4Addr::new(10, 0, 0, octet))
}
#[test]
fn test_rate_limit_under() {
let limiter = RateLimiter::new(10.0, 10.0);
assert!(limiter.check(1));
assert!(limiter.check(test_ip(1)));
}
#[test]
fn test_rate_limit_over() {
let limiter = RateLimiter::new(1.0, 1.0);
assert!(limiter.check(1));
assert!(!limiter.check(1));
assert!(limiter.check(test_ip(1)));
assert!(!limiter.check(test_ip(1)));
}
#[test]
fn test_rate_limit_burst() {
let limiter = RateLimiter::new(1.0, 5.0);
for _ in 0..5 {
assert!(limiter.check(2));
assert!(limiter.check(test_ip(2)));
}
assert!(!limiter.check(2));
assert!(!limiter.check(test_ip(2)));
}
#[test]
fn test_rate_limit_refill() {
let limiter = RateLimiter::new(100.0, 1.0);
assert!(limiter.check(3));
assert!(!limiter.check(3));
assert!(limiter.check(test_ip(3)));
assert!(!limiter.check(test_ip(3)));
std::thread::sleep(Duration::from_millis(20));
assert!(limiter.check(3));
assert!(limiter.check(test_ip(3)));
}
#[test]
fn test_sweep_removes_idle_keeps_active() {
let limiter = RateLimiter::with_eviction(1.0, 10.0, Duration::from_millis(20));
limiter.check(test_ip(1));
limiter.check(test_ip(2));
std::thread::sleep(Duration::from_millis(50));
limiter.check(test_ip(2));
limiter.sweep();
assert_eq!(limiter.len(), 1);
assert!(!limiter.buckets.contains_key(&test_ip(1)));
assert!(limiter.buckets.contains_key(&test_ip(2)));
}
#[test]
fn test_sweep_does_not_remove_active() {
let limiter = RateLimiter::with_eviction(1.0, 10.0, Duration::from_millis(50));
limiter.check(test_ip(1));
limiter.sweep();
assert_eq!(limiter.len(), 1);
assert!(limiter.buckets.contains_key(&test_ip(1)));
}
}

View file

@ -6,6 +6,4 @@ pub mod pow;
pub mod proxy;
pub mod store;
pub mod traffic;
#[cfg(feature = "xdp")]
pub mod xdp;

View file

@ -4,11 +4,26 @@ use rampart_core::filter::rate_limit::RateLimiter;
use rampart_core::metrics;
use rampart_core::pow::difficulty::DifficultyAdjuster;
use rampart_core::proxy::listener::ProxyListener;
use rampart_core::store::clickhouse::{ClickHouseEvent, ClickHouseWriter};
use rampart_core::traffic::detector::{AttackDetector, AttackStatus};
use rampart_core::traffic::reputation::IpReputation;
use rampart_core::xdp::XdpFilter;
use std::collections::HashSet;
use std::net::IpAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::watch;
use tracing_subscriber::EnvFilter;
fn attack_status_value(status: AttackStatus) -> i64 {
match status {
AttackStatus::Normal => 0,
AttackStatus::Suspicious => 1,
AttackStatus::UnderAttack => 2,
}
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
tracing_subscriber::fmt()
@ -17,6 +32,7 @@ async fn main() -> anyhow::Result<()> {
let config_path = std::env::var("RAMPART_CONFIG").unwrap_or_else(|_| "/etc/rampart/config.toml".to_string());
let config = Config::from_file(&config_path)?;
let whitelist = build_whitelist(&config)?;
let config = Arc::new(config);
let rate_limiter = Arc::new(RateLimiter::new(
@ -24,6 +40,9 @@ async fn main() -> anyhow::Result<()> {
config.limits.rate_limit_burst,
));
let blacklist = Arc::new(Blacklist::new());
let reputation = Arc::new(IpReputation::new());
let detector = Arc::new(Mutex::new(AttackDetector::new()));
let allowed_1s = Arc::new(AtomicU64::new(0));
let (shutdown_tx, shutdown_rx) = watch::channel(false);
@ -64,35 +83,135 @@ async fn main() -> anyhow::Result<()> {
});
}
#[cfg(feature = "xdp")]
if config.xdp.enabled {
use rampart_core::xdp::{XdpFilter, XdpMetrics};
let clickhouse: Option<Arc<tokio::sync::Mutex<ClickHouseWriter>>> = match &config.store.clickhouse_url {
Some(url) if !url.is_empty() => {
let writer = Arc::new(tokio::sync::Mutex::new(ClickHouseWriter::new(url)));
rampart_core::store::clickhouse::start_flush_task(writer.clone(), shutdown_rx.clone());
Some(writer)
},
_ => None,
};
let mut filter = XdpFilter::new(&config.xdp.interface);
filter.load()?;
#[cfg(feature = "xdp")]
let xdp_filter: Option<Arc<Mutex<XdpFilter>>> = if config.xdp.enabled {
use rampart_core::xdp::XdpMetrics;
let filter = XdpFilter::new(&config.xdp.interface);
let shared = Arc::new(Mutex::new(filter));
shared.lock().expect("xdp lock poisoned").load()?;
let xdp_metrics = XdpMetrics::register()?;
let sd = shutdown_rx.clone();
let shared_thread = shared.clone();
std::thread::spawn(move || {
while !*sd.borrow() {
filter.drain_events();
if let Ok(stats) = filter.get_stats() {
let guard = match shared_thread.lock() {
Ok(g) => g,
Err(_) => break,
};
guard.drain_events();
if let Ok(stats) = guard.get_stats() {
xdp_metrics.update(&stats);
}
drop(guard);
std::thread::sleep(Duration::from_secs(5));
}
filter.unload().ok();
if let Ok(mut guard) = shared_thread.lock() {
guard.unload().ok();
}
});
}
Some(shared)
} else {
None
};
#[cfg(not(feature = "xdp"))]
let xdp_filter: Option<Arc<Mutex<XdpFilter>>> = None;
let rl = rate_limiter.clone();
let bl = blacklist.clone();
let det = detector.clone();
let a1s = allowed_1s.clone();
let ch = clickhouse.clone();
let mut sd = shutdown_rx.clone();
tokio::spawn(async move {
let mut sec_tick = tokio::time::interval(Duration::from_secs(1));
let mut min_tick = tokio::time::interval(Duration::from_secs(60));
sec_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
min_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
let mut was_under_attack = false;
loop {
tokio::select! {
biased;
_ = sd.changed() => {
if *sd.borrow() {
return;
}
}
_ = sec_tick.tick() => {
let pps = a1s.swap(0, Ordering::Relaxed) as f64;
let status = det.lock().expect("detector lock poisoned").analyze(pps);
metrics::ATTACK_STATUS.set(attack_status_value(status));
if status == AttackStatus::UnderAttack {
if !was_under_attack {
was_under_attack = true;
tracing::info!(pps, "attack detected: under attack");
if let Some(writer) = &ch {
let event = ClickHouseEvent {
timestamp: chrono::Utc::now(),
event_type: "attack".to_string(),
ip: String::new(),
data_float: pps,
data_int: 0,
data_string: "under_attack".to_string(),
};
if let Err(e) = writer.lock().await.push(event).await {
tracing::debug!("clickhouse push error: {e}");
}
}
}
} else if was_under_attack {
was_under_attack = false;
}
}
_ = min_tick.tick() => {
rl.sweep();
bl.clear_expired();
}
}
}
});
tracing::info!("Rampart edge starting on {}:{}", config.bind.address, config.bind.port);
tracing::info!("Backend: {}:{}", config.backend.address, config.backend.port);
let adjuster = Arc::new(Mutex::new(DifficultyAdjuster::default()));
let listener = ProxyListener::new(config, rate_limiter, blacklist, adjuster);
let listener = ProxyListener::new(
config,
rate_limiter,
blacklist,
adjuster,
whitelist,
reputation,
xdp_filter,
clickhouse,
allowed_1s,
);
listener.run(shutdown_rx).await
}
fn build_whitelist(config: &Config) -> anyhow::Result<Arc<HashSet<IpAddr>>> {
let mut set = HashSet::with_capacity(config.whitelist.len());
for entry in &config.whitelist {
let ip: IpAddr = match entry.parse() {
Ok(ip) => ip,
Err(_) => anyhow::bail!("invalid whitelist entry: {entry}"),
};
set.insert(ip);
}
Ok(Arc::new(set))
}
async fn wait_for_signal() {
let ctrl_c = tokio::signal::ctrl_c();
let mut term = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())

View file

@ -33,6 +33,14 @@ pub static POW_CURRENT_DIFFICULTY: LazyLock<IntGauge> = LazyLock::new(|| {
register_int_gauge!("rampart_pow_current_difficulty", "Current PoW difficulty").expect("POW_CURRENT_DIFFICULTY")
});
pub static ATTACK_STATUS: LazyLock<IntGauge> = LazyLock::new(|| {
register_int_gauge!(
"rampart_attack_status",
"Attack detector status (0=normal, 1=suspicious, 2=under attack)"
)
.expect("ATTACK_STATUS")
});
pub async fn run_metrics_server(addr: &str) {
let listener = match TcpListener::bind(addr).await {
Ok(l) => l,

View file

@ -124,25 +124,25 @@ mod tests {
#[test]
fn test_varint_zero() {
let buf = vec![0x00];
assert_eq!(read_varint(&buf, 0).unwrap(), (0, 1));
assert_eq!(read_varint(&buf, 0).expect("varint should parse"), (0, 1));
}
#[test]
fn test_varint_single() {
let buf = vec![0x7F];
assert_eq!(read_varint(&buf, 0).unwrap(), (127, 1));
assert_eq!(read_varint(&buf, 0).expect("varint should parse"), (127, 1));
}
#[test]
fn test_varint_multi() {
let buf = vec![0x80, 0x01];
assert_eq!(read_varint(&buf, 0).unwrap(), (128, 2));
assert_eq!(read_varint(&buf, 0).expect("varint should parse"), (128, 2));
}
#[test]
fn test_varint_max() {
let buf = vec![0xFF, 0xFF, 0xFF, 0xFF, 0x07];
assert_eq!(read_varint(&buf, 0).unwrap(), (i32::MAX, 5));
assert_eq!(read_varint(&buf, 0).expect("varint should parse"), (i32::MAX, 5));
}
#[test]
@ -173,7 +173,7 @@ mod tests {
let len = (buf.len() - 1) as u8;
buf[0] = len;
let hs = McHandshake::parse(&buf).unwrap();
let hs = McHandshake::parse(&buf).expect("valid login handshake should parse");
assert_eq!(hs.protocol_version, 765);
assert_eq!(hs.server_address, "play.example.com");
assert_eq!(hs.server_port, 25565);
@ -206,7 +206,7 @@ mod tests {
let len = (buf.len() - 1) as u8;
buf[0] = len;
let hs = McHandshake::parse(&buf).unwrap();
let hs = McHandshake::parse(&buf).expect("valid status handshake should parse");
assert_eq!(hs.protocol_version, 2);
assert_eq!(hs.server_address, "play.example");
assert_eq!(hs.server_port, 25565);

View file

@ -3,9 +3,16 @@ use crate::filter::blacklist::Blacklist;
use crate::filter::rate_limit::RateLimiter;
use crate::pow::difficulty::DifficultyAdjuster;
use crate::proxy::tunnel::ConnectionHandler;
use crate::store::clickhouse::ClickHouseWriter;
use crate::traffic::reputation::IpReputation;
use crate::xdp::XdpFilter;
use socket2::{Domain, Socket, Type};
use std::collections::HashSet;
use std::net::IpAddr;
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, Mutex};
use tokio::net::TcpListener;
use tokio::sync::Mutex as TokioMutex;
use tokio::sync::watch;
pub struct ProxyListener {
@ -13,6 +20,11 @@ pub struct ProxyListener {
rate_limiter: Arc<RateLimiter>,
blacklist: Arc<Blacklist>,
adjuster: Arc<Mutex<DifficultyAdjuster>>,
whitelist: Arc<HashSet<IpAddr>>,
reputation: Arc<IpReputation>,
xdp: Option<Arc<Mutex<XdpFilter>>>,
clickhouse: Option<Arc<TokioMutex<ClickHouseWriter>>>,
allowed_1s: Arc<AtomicU64>,
}
impl ProxyListener {
@ -21,12 +33,22 @@ impl ProxyListener {
rate_limiter: Arc<RateLimiter>,
blacklist: Arc<Blacklist>,
adjuster: Arc<Mutex<DifficultyAdjuster>>,
whitelist: Arc<HashSet<IpAddr>>,
reputation: Arc<IpReputation>,
xdp: Option<Arc<Mutex<XdpFilter>>>,
clickhouse: Option<Arc<TokioMutex<ClickHouseWriter>>>,
allowed_1s: Arc<AtomicU64>,
) -> Self {
Self {
config,
rate_limiter,
blacklist,
adjuster,
whitelist,
reputation,
xdp,
clickhouse,
allowed_1s,
}
}
@ -41,6 +63,11 @@ impl ProxyListener {
let rate_limiter = self.rate_limiter.clone();
let blacklist = self.blacklist.clone();
let adjuster = self.adjuster.clone();
let whitelist = self.whitelist.clone();
let reputation = self.reputation.clone();
let xdp = self.xdp.clone();
let clickhouse = self.clickhouse.clone();
let allowed_1s = self.allowed_1s.clone();
let shutdown = shutdown.clone();
handles.push(tokio::spawn(accept_loop(
listener,
@ -48,6 +75,11 @@ impl ProxyListener {
rate_limiter,
blacklist,
adjuster,
whitelist,
reputation,
xdp,
clickhouse,
allowed_1s,
shutdown,
)));
}
@ -69,12 +101,18 @@ fn build_listener(addr: std::net::SocketAddr) -> anyhow::Result<TcpListener> {
Ok(TcpListener::from_std(socket.into())?)
}
#[allow(clippy::too_many_arguments)]
async fn accept_loop(
listener: TcpListener,
config: Arc<Config>,
rate_limiter: Arc<RateLimiter>,
blacklist: Arc<Blacklist>,
adjuster: Arc<Mutex<DifficultyAdjuster>>,
whitelist: Arc<HashSet<IpAddr>>,
reputation: Arc<IpReputation>,
xdp: Option<Arc<Mutex<XdpFilter>>>,
clickhouse: Option<Arc<TokioMutex<ClickHouseWriter>>>,
allowed_1s: Arc<AtomicU64>,
mut shutdown: watch::Receiver<bool>,
) -> anyhow::Result<()> {
loop {
@ -94,7 +132,17 @@ async fn accept_loop(
continue;
}
};
let handler = ConnectionHandler::new(config.clone(), rate_limiter.clone(), blacklist.clone(), adjuster.clone());
let handler = ConnectionHandler::new(
config.clone(),
rate_limiter.clone(),
blacklist.clone(),
adjuster.clone(),
whitelist.clone(),
reputation.clone(),
xdp.clone(),
clickhouse.clone(),
allowed_1s.clone(),
);
tokio::spawn(async move {
if let Err(e) = handler.handle(stream, peer_addr).await {
tracing::debug!("connection from {peer_addr}: {e}");

View file

@ -1,10 +1,10 @@
use crate::pow::challenge::Challenge;
use std::net::Ipv4Addr;
use std::net::IpAddr;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio::time::{Duration, timeout};
pub async fn handle_pow(stream: &mut TcpStream, peer_ip: Ipv4Addr, difficulty: u8) -> anyhow::Result<bool> {
pub async fn handle_pow(stream: &mut TcpStream, peer_ip: IpAddr, difficulty: u8) -> anyhow::Result<bool> {
if difficulty == 0 {
tracing::debug!("pow: difficulty 0, skipping for {peer_ip}");
return Ok(true);

View file

@ -5,54 +5,71 @@ use crate::filter::death_code;
use crate::filter::rate_limit::RateLimiter;
use crate::metrics;
use crate::pow::difficulty::DifficultyAdjuster;
use crate::proxy::handshake::{McHandshake, read_varint};
use crate::proxy::handshake::{McHandshake, ParseError, read_varint};
use crate::proxy::pow::handle_pow;
use std::net::Ipv4Addr;
use crate::store::clickhouse::{ClickHouseEvent, ClickHouseWriter};
use crate::traffic::reputation::IpReputation;
use crate::xdp::XdpFilter;
use std::collections::HashSet;
use std::net::IpAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio::sync::Mutex as TokioMutex;
const MAX_FRAME_SIZE: usize = 8192;
const READ_CHUNK_SIZE: usize = 512;
pub struct ConnectionHandler {
config: Arc<Config>,
rate_limiter: Arc<RateLimiter>,
blacklist: Arc<Blacklist>,
adjuster: Arc<Mutex<DifficultyAdjuster>>,
whitelist: Arc<HashSet<IpAddr>>,
reputation: Arc<IpReputation>,
xdp: Option<Arc<Mutex<XdpFilter>>>,
clickhouse: Option<Arc<TokioMutex<ClickHouseWriter>>>,
allowed_1s: Arc<AtomicU64>,
}
impl ConnectionHandler {
#[allow(clippy::too_many_arguments)]
pub fn new(
config: Arc<Config>,
rate_limiter: Arc<RateLimiter>,
blacklist: Arc<Blacklist>,
adjuster: Arc<Mutex<DifficultyAdjuster>>,
whitelist: Arc<HashSet<IpAddr>>,
reputation: Arc<IpReputation>,
xdp: Option<Arc<Mutex<XdpFilter>>>,
clickhouse: Option<Arc<TokioMutex<ClickHouseWriter>>>,
allowed_1s: Arc<AtomicU64>,
) -> Self {
Self {
config,
rate_limiter,
blacklist,
adjuster,
}
}
fn ip_to_u32(addr: std::net::SocketAddr) -> u32 {
match addr.ip() {
std::net::IpAddr::V4(ip) => ip.to_bits(),
_ => 0,
whitelist,
reputation,
xdp,
clickhouse,
allowed_1s,
}
}
pub async fn handle(&self, mut client: TcpStream, peer_addr: std::net::SocketAddr) -> anyhow::Result<()> {
let ip_u32 = Self::ip_to_u32(peer_addr);
let peer_ip = peer_addr.ip();
if self.blacklist.is_blocked(ip_u32) {
if self.blacklist.is_blocked(peer_ip) {
metrics::CONNECTIONS_TOTAL.with_label_values(&["blocked"]).inc();
return Ok(());
}
let pow_config = &self.config.pow;
let peer_ip = Ipv4Addr::from_bits(ip_u32);
if pow_config.enabled && pow_config.difficulty > 0 && !self.config.whitelist.contains(&peer_ip.to_string()) {
if pow_config.enabled && pow_config.difficulty > 0 && !self.whitelist.contains(&peer_ip) {
self.adjuster
.lock()
.expect("adjuster lock poisoned")
@ -75,35 +92,45 @@ impl ConnectionHandler {
metrics::POW_CURRENT_DIFFICULTY.set(pow_config.difficulty as i64);
}
if !self.rate_limiter.check(ip_u32) {
metrics::RATE_LIMIT_HITS.with_label_values(&["hit"]).inc();
metrics::CONNECTIONS_TOTAL.with_label_values(&["blocked"]).inc();
if !self.rate_limiter.check(peer_ip) {
self.block_rate_limit(peer_ip).await;
return Ok(());
}
let timeout = Duration::from_secs(self.config.limits.handshake_timeout_secs);
let mut buf = vec![0u8; 4096];
let n = tokio::time::timeout(timeout, client.read(&mut buf)).await??;
if n == 0 {
return Ok(());
let mut buf: Vec<u8> = Vec::new();
match read_full_frame(&mut client, timeout, &mut buf).await {
Ok(false) => return Ok(()),
Ok(true) => {},
Err(e) => {
tracing::debug!("read error from {peer_addr}: {e}");
metrics::CONNECTIONS_TOTAL.with_label_values(&["blocked"]).inc();
self.handle_death_code(peer_addr, &buf).await;
return Ok(());
},
}
let parsed = McHandshake::parse(&buf[..n]);
let parsed = McHandshake::parse(&buf);
match parsed {
Ok(handshake) => {
if !self.rate_limiter.check(ip_u32) {
metrics::RATE_LIMIT_HITS.with_label_values(&["hit"]).inc();
metrics::CONNECTIONS_TOTAL.with_label_values(&["blocked"]).inc();
if !self.rate_limiter.check(peer_ip) {
self.block_rate_limit(peer_ip).await;
return Ok(());
}
metrics::CONNECTIONS_TOTAL.with_label_values(&["allowed"]).inc();
self.allowed_1s.fetch_add(1, Ordering::Relaxed);
self.reputation.record_good(peer_ip);
let backend_addr = format!("{}:{}", self.config.backend.address, self.config.backend.port);
let mut backend = TcpStream::connect(&backend_addr).await?;
let signed = hmac::sign_hostname(&handshake.server_address, self.config.hmac.secret.as_bytes());
let modified = replace_hostname(&buf[..n], &handshake.server_address, &signed)?;
let signed = hmac::sign_hostname(
&handshake.server_address,
self.config.hmac.secret.as_bytes(),
self.config.hmac.key_rotation_interval_secs,
);
let modified = replace_hostname(&buf, &handshake.server_address, &signed)?;
backend.write_all(&modified).await?;
tokio::io::copy_bidirectional(&mut client, &mut backend).await?;
@ -111,19 +138,114 @@ impl ConnectionHandler {
Err(e) => {
tracing::debug!("parse error from {peer_addr}: {e}");
metrics::CONNECTIONS_TOTAL.with_label_values(&["blocked"]).inc();
if self.config.death_code.enabled {
if let Some(code) = death_code::detect(&buf[..n]) {
let duration = Duration::from_secs(self.config.death_code.ban_duration_secs);
self.blacklist.add(ip_u32, duration, code.as_str());
metrics::DEATH_CODE_BANS_TOTAL.with_label_values(&[code.as_str()]).inc();
tracing::info!("death code ban {peer_addr}: {}", code.as_str());
}
}
self.handle_death_code(peer_addr, &buf).await;
},
}
Ok(())
}
async fn block_rate_limit(&self, ip: IpAddr) {
metrics::RATE_LIMIT_HITS.with_label_values(&["hit"]).inc();
metrics::CONNECTIONS_TOTAL.with_label_values(&["blocked"]).inc();
self.reputation.record_bad(ip);
if self.reputation.score(ip) < -40 {
let duration_secs = self.config.death_code.ban_duration_secs;
self.blacklist
.add(ip, Duration::from_secs(duration_secs), "low_reputation");
self.xdp_ban(ip, duration_secs);
self.push_event("block", ip, "low_reputation").await;
tracing::info!("low reputation ban {ip}: rate-limit abuse");
}
}
async fn handle_death_code(&self, peer_addr: std::net::SocketAddr, buf: &[u8]) {
if !self.config.death_code.enabled {
return;
}
if let Some(code) = death_code::detect(buf) {
let duration_secs = self.config.death_code.ban_duration_secs;
let ip = peer_addr.ip();
self.blacklist
.add(ip, Duration::from_secs(duration_secs), code.as_str());
self.reputation.record_bad(ip);
self.xdp_ban(ip, duration_secs);
self.push_event("ban", ip, code.as_str()).await;
metrics::DEATH_CODE_BANS_TOTAL.with_label_values(&[code.as_str()]).inc();
tracing::info!("death code ban {peer_addr}: {}", code.as_str());
}
}
fn xdp_ban(&self, ip: IpAddr, duration_secs: u64) {
let Some(xdp) = &self.xdp else {
return;
};
let IpAddr::V4(ip_v4) = ip else {
return;
};
match xdp.lock().expect("xdp lock poisoned").ban_ip(ip_v4, duration_secs) {
Ok(()) => tracing::debug!("xdp ban {ip_v4} for {duration_secs}s"),
Err(e) => tracing::warn!("xdp ban failed for {ip_v4}: {e}"),
}
}
async fn push_event(&self, event_type: &str, ip: IpAddr, data_string: &str) {
let Some(writer) = &self.clickhouse else {
return;
};
let event = ClickHouseEvent {
timestamp: chrono::Utc::now(),
event_type: event_type.to_string(),
ip: ip.to_string(),
data_float: 0.0,
data_int: 0,
data_string: data_string.to_string(),
};
if let Err(e) = writer.lock().await.push(event).await {
tracing::debug!("clickhouse push error: {e}");
}
}
}
async fn read_full_frame(client: &mut TcpStream, timeout: Duration, buf: &mut Vec<u8>) -> anyhow::Result<bool> {
let mut chunk = [0u8; READ_CHUNK_SIZE];
let first = tokio::time::timeout(timeout, client.read(&mut chunk)).await??;
if first == 0 {
return Ok(false);
}
buf.extend_from_slice(&chunk[..first]);
let total_len = loop {
match read_varint(buf, 0) {
Ok((packet_len, after_len)) => break after_len + packet_len as usize,
Err(ParseError::Incomplete(_)) => {
if buf.len() >= 5 {
anyhow::bail!("length varint incomplete after {} bytes", buf.len());
}
let n = tokio::time::timeout(timeout, client.read(&mut chunk)).await??;
if n == 0 {
anyhow::bail!("connection closed while reading packet length");
}
buf.extend_from_slice(&chunk[..n]);
},
Err(e) => anyhow::bail!("invalid packet length varint: {e}"),
}
};
if total_len > MAX_FRAME_SIZE {
anyhow::bail!("frame too large: {total_len} bytes (max {MAX_FRAME_SIZE})");
}
while buf.len() < total_len {
let n = tokio::time::timeout(timeout, client.read(&mut chunk)).await??;
if n == 0 {
anyhow::bail!("connection closed while reading frame body");
}
buf.extend_from_slice(&chunk[..n]);
}
buf.truncate(total_len);
Ok(true)
}
fn replace_hostname(original: &[u8], _old_hostname: &str, new_hostname: &str) -> anyhow::Result<Vec<u8>> {
@ -197,10 +319,10 @@ mod tests {
fn test_replace_hostname_basic() {
let pkt = build_test_packet("play.example.com");
let new_hostname = "play.example.com\0shield\0abcdef1234567890";
let modified = replace_hostname(&pkt, "play.example.com", new_hostname).unwrap();
let modified = replace_hostname(&pkt, "play.example.com", new_hostname).expect("should replace hostname");
assert!(modified.len() > pkt.len());
let parsed = McHandshake::parse(&modified).unwrap();
let parsed = McHandshake::parse(&modified).expect("signed hostname should parse");
assert_eq!(parsed.server_address, new_hostname);
}
@ -208,10 +330,11 @@ mod tests {
fn test_replace_hostname_shorter() {
let pkt = build_test_packet("very.long.hostname.example.com");
let new_hostname = "short.com";
let modified = replace_hostname(&pkt, "very.long.hostname.example.com", new_hostname).unwrap();
let modified =
replace_hostname(&pkt, "very.long.hostname.example.com", new_hostname).expect("should replace hostname");
assert!(modified.len() < pkt.len());
let parsed = McHandshake::parse(&modified).unwrap();
let parsed = McHandshake::parse(&modified).expect("short hostname should parse");
assert_eq!(parsed.server_address, new_hostname);
}
@ -219,9 +342,9 @@ mod tests {
fn test_replace_hostname_preserves_port_and_protocol() {
let pkt = build_test_packet("mc.example.com");
let new_hostname = "mc.example.com\0shield\x00deadbeef";
let modified = replace_hostname(&pkt, "mc.example.com", new_hostname).unwrap();
let modified = replace_hostname(&pkt, "mc.example.com", new_hostname).expect("should replace hostname");
let parsed = McHandshake::parse(&modified).unwrap();
let parsed = McHandshake::parse(&modified).expect("signed hostname should parse");
assert_eq!(parsed.server_port, 25565);
assert_eq!(parsed.protocol_version, 765);
assert!(parsed.is_login());
@ -232,7 +355,7 @@ mod tests {
let cases = vec![0, 1, 127, 128, 255, 65535, 1000000, i32::MAX];
for val in cases {
let bytes = varint_bytes(val);
let (decoded, _) = read_varint(&bytes, 0).unwrap();
let (decoded, _) = read_varint(&bytes, 0).expect("varint should parse");
assert_eq!(decoded, val, "roundtrip failed for {val}");
}
}

View file

@ -86,23 +86,15 @@ fn handle_event(msg: &Msg, blacklist: &Blacklist) -> anyhow::Result<()> {
let payload: String = msg.get_payload()?;
let event: BlacklistEvent = serde_json::from_str(&payload)?;
let ip_parts: Vec<&str> = event.ip.split('.').collect();
if ip_parts.len() != 4 {
anyhow::bail!("invalid IP: {}", event.ip);
}
let mut ip_u32: u32 = 0;
for part in &ip_parts {
let octet: u32 = part.parse()?;
ip_u32 = (ip_u32 << 8) | octet;
}
let ip: std::net::IpAddr = event.ip.parse()?;
match event.action.as_str() {
"ban" => {
blacklist.add(ip_u32, Duration::from_secs(event.duration_secs), "redis");
blacklist.add(ip, Duration::from_secs(event.duration_secs), "redis");
tracing::info!("blacklist add via Redis: {}", event.ip);
},
"unban" => {
blacklist.remove(ip_u32);
blacklist.remove(ip);
tracing::info!("blacklist remove via Redis: {}", event.ip);
},
a => anyhow::bail!("unknown action: {a}"),

View file

@ -1,9 +1,9 @@
use dashmap::DashMap;
use std::net::Ipv4Addr;
use std::net::IpAddr;
use std::sync::Arc;
pub struct IpReputation {
scores: Arc<DashMap<Ipv4Addr, i32>>,
scores: Arc<DashMap<IpAddr, i32>>,
}
impl Default for IpReputation {
@ -19,21 +19,21 @@ impl IpReputation {
}
}
pub fn record_good(&self, ip: Ipv4Addr) {
pub fn record_good(&self, ip: IpAddr) {
let mut entry = self.scores.entry(ip).or_insert(0);
*entry = (*entry + 1).min(100);
}
pub fn record_bad(&self, ip: Ipv4Addr) {
pub fn record_bad(&self, ip: IpAddr) {
let mut entry = self.scores.entry(ip).or_insert(0);
*entry = (*entry - 10).max(-100);
}
pub fn score(&self, ip: Ipv4Addr) -> i32 {
pub fn score(&self, ip: IpAddr) -> i32 {
self.scores.get(&ip).map(|v| *v).unwrap_or(0)
}
pub fn is_trusted(&self, ip: Ipv4Addr) -> bool {
pub fn is_trusted(&self, ip: IpAddr) -> bool {
self.score(ip) > 50
}
}
@ -41,18 +41,22 @@ impl IpReputation {
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
fn v4(octets: [u8; 4]) -> IpAddr {
IpAddr::V4(Ipv4Addr::from(octets))
}
#[test]
fn test_reputation_initial_score() {
let rep = IpReputation::new();
assert_eq!(rep.score(Ipv4Addr::new(192, 168, 1, 1)), 0);
assert_eq!(rep.score(v4([192, 168, 1, 1])), 0);
}
#[test]
fn test_reputation_good() {
let rep = IpReputation::new();
let ip = Ipv4Addr::new(10, 0, 0, 1);
let ip = v4([10, 0, 0, 1]);
rep.record_good(ip);
assert_eq!(rep.score(ip), 1);
}
@ -60,7 +64,7 @@ mod tests {
#[test]
fn test_reputation_bad() {
let rep = IpReputation::new();
let ip = Ipv4Addr::new(10, 0, 0, 2);
let ip = v4([10, 0, 0, 2]);
rep.record_bad(ip);
assert_eq!(rep.score(ip), -10);
}
@ -68,7 +72,7 @@ mod tests {
#[test]
fn test_reputation_cap_positive() {
let rep = IpReputation::new();
let ip = Ipv4Addr::new(10, 0, 0, 3);
let ip = v4([10, 0, 0, 3]);
for _ in 0..200 {
rep.record_good(ip);
}
@ -78,7 +82,7 @@ mod tests {
#[test]
fn test_reputation_cap_negative() {
let rep = IpReputation::new();
let ip = Ipv4Addr::new(10, 0, 0, 4);
let ip = v4([10, 0, 0, 4]);
for _ in 0..20 {
rep.record_bad(ip);
}
@ -88,11 +92,19 @@ mod tests {
#[test]
fn test_is_trusted() {
let rep = IpReputation::new();
let ip = Ipv4Addr::new(10, 0, 0, 5);
let ip = v4([10, 0, 0, 5]);
assert!(!rep.is_trusted(ip));
for _ in 0..51 {
rep.record_good(ip);
}
assert!(rep.is_trusted(ip));
}
#[test]
fn test_reputation_v6() {
let rep = IpReputation::new();
let ip = IpAddr::V6(Ipv6Addr::LOCALHOST);
rep.record_bad(ip);
assert_eq!(rep.score(ip), -10);
}
}

View file

@ -77,7 +77,7 @@ impl XdpFilter {
.with_context(|| format!("map '{}' not found", name))
}
pub fn ban_ip(&self, ip: Ipv4Addr) -> Result<()> {
pub fn ban_ip(&self, ip: Ipv4Addr, duration_secs: u64) -> Result<()> {
let map = self.find_map("blacklist_map")?;
let mut key = [0u8; 8];
key[0] = 32;
@ -86,7 +86,11 @@ impl XdpFilter {
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
map.update(&key, &(now + 300_000_000_000).to_le_bytes(), MapFlags::ANY)?;
map.update(
&key,
&(now + duration_secs * 1_000_000_000).to_le_bytes(),
MapFlags::ANY,
)?;
Ok(())
}

View file

@ -14,7 +14,7 @@ impl XdpFilter {
Ok(())
}
pub fn drain_events(&self) {}
pub fn ban_ip(&self, _ip: Ipv4Addr) -> Result<()> {
pub fn ban_ip(&self, _ip: Ipv4Addr, _duration_secs: u64) -> Result<()> {
Ok(())
}
pub fn unban_ip(&self, _ip: Ipv4Addr) -> Result<()> {

View file

@ -16,6 +16,7 @@ tracing-subscriber.workspace = true
thiserror.workspace = true
anyhow.workspace = true
dashmap.workspace = true
subtle.workspace = true
prometheus.workspace = true
axum = "0.8"
tower-http = { version = "0.6", features = ["cors"] }

View file

@ -1,21 +1,104 @@
use crate::AppState;
use axum::{Json, extract::State};
use axum::{
Json,
extract::{ConnectInfo, State},
http::StatusCode,
};
use dashmap::DashMap;
use serde::Deserialize;
use std::sync::Arc;
use std::{
net::{IpAddr, SocketAddr},
sync::Arc,
time::{Duration, Instant},
};
use subtle::ConstantTimeEq;
const LOGIN_WINDOW: Duration = Duration::from_secs(60);
const LOGIN_MAX_ATTEMPTS: u32 = 5;
#[derive(Deserialize)]
pub struct LoginRequest {
pub password: String,
}
pub async fn login(State(state): State<Arc<AppState>>, Json(req): Json<LoginRequest>) -> Json<serde_json::Value> {
let api_password = std::env::var("API_PASSWORD").unwrap_or_else(|_| "changeme".to_string());
if req.password != api_password {
return Json(serde_json::json!({"error": "invalid password"}));
pub async fn login(
State(state): State<Arc<AppState>>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
Json(req): Json<LoginRequest>,
) -> Result<Json<serde_json::Value>, (StatusCode, Json<serde_json::Value>)> {
if !allow_login_attempt(&state.login_limiter, addr.ip()) {
return Err((
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({"error": "too many requests"})),
));
}
match crate::auth::create_token(&state.jwt_secret, state.jwt_expiration) {
Ok(token) => Json(serde_json::json!({"token": token})),
Err(_) => Json(serde_json::json!({"error": "token creation failed"})),
if !verify_password(&req.password, &state.api_password) {
return Err((
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": "invalid password"})),
));
}
match crate::auth::create_token(&state.jwt_secret, state.jwt_expiration, &state.jwt_audience) {
Ok(token) => Ok(Json(serde_json::json!({"token": token}))),
Err(_) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": "token creation failed"})),
)),
}
}
fn verify_password(provided: &str, expected: &str) -> bool {
let provided = provided.as_bytes();
let expected = expected.as_bytes();
let len_match = (provided.len() as u64).ct_eq(&(expected.len() as u64));
let min_len = provided.len().min(expected.len());
let bytes_match = provided[..min_len].ct_eq(&expected[..min_len]);
bool::from(len_match & bytes_match)
}
fn allow_login_attempt(limiter: &DashMap<IpAddr, (Instant, u32)>, ip: IpAddr) -> bool {
let now = Instant::now();
let mut slot = limiter.entry(ip).or_insert((now, 0));
let (last_reset, attempts) = &mut *slot;
if now.duration_since(*last_reset) >= LOGIN_WINDOW {
*last_reset = now;
*attempts = 1;
} else if *attempts >= LOGIN_MAX_ATTEMPTS {
return false;
} else {
*attempts += 1;
}
true
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn verify_password_matches() {
assert!(verify_password("s3cret-pass", "s3cret-pass"));
}
#[test]
fn verify_password_rejects_wrong() {
assert!(!verify_password("wrong-pass", "s3cret-pass"));
}
#[test]
fn verify_password_rejects_different_length() {
assert!(!verify_password("short", "a-longer-password"));
}
#[test]
fn allow_login_attempt_respects_limit() {
let limiter = DashMap::new();
let ip = IpAddr::from([127, 0, 0, 1]);
for _ in 0..LOGIN_MAX_ATTEMPTS {
assert!(allow_login_attempt(&limiter, ip));
}
assert!(!allow_login_attempt(&limiter, ip));
}
}

View file

@ -10,26 +10,28 @@ use std::sync::Arc;
#[derive(Debug, Serialize, Deserialize)]
pub struct Claims {
pub sub: String,
pub aud: String,
pub role: String,
pub exp: usize,
pub iat: usize,
}
pub fn create_token(secret: &str, expiration: u64) -> Result<String, jsonwebtoken::errors::Error> {
pub fn create_token(secret: &str, expiration: u64, audience: &str) -> Result<String, jsonwebtoken::errors::Error> {
let now = chrono::Utc::now().timestamp() as usize;
let claims = Claims {
sub: "rampart-admin".to_string(),
aud: audience.to_string(),
role: "admin".to_string(),
exp: now + expiration as usize,
iat: now,
};
encode(&Header::default(), &claims, &EncodingKey::from_secret(secret.as_ref()))
}
pub fn verify_token(token: &str, secret: &str) -> Result<Claims, jsonwebtoken::errors::Error> {
let token_data = decode::<Claims>(
token,
&DecodingKey::from_secret(secret.as_ref()),
&Validation::default(),
)?;
pub fn verify_token(token: &str, secret: &str, audience: &str) -> Result<Claims, jsonwebtoken::errors::Error> {
let mut validation = Validation::default();
validation.set_audience(&[audience]);
let token_data = decode::<Claims>(token, &DecodingKey::from_secret(secret.as_ref()), &validation)?;
Ok(token_data.claims)
}
@ -63,7 +65,7 @@ pub async fn auth_middleware(
},
};
if verify_token(token, &state.jwt_secret).is_err() {
if verify_token(token, &state.jwt_secret, &state.jwt_audience).is_err() {
return Err((
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": "unauthorized"})),
@ -72,3 +74,38 @@ pub async fn auth_middleware(
Ok(next.run(request).await)
}
#[cfg(test)]
mod tests {
use super::*;
fn valid_secret() -> String {
"this-is-a-test-secret-32-bytes-long!".to_string()
}
#[test]
fn create_verify_roundtrip_passes() {
let secret = valid_secret();
let token = create_token(&secret, 3600, "rampart").expect("token creation should succeed");
let claims = verify_token(&token, &secret, "rampart").expect("verification should succeed");
assert_eq!(claims.sub, "rampart-admin");
assert_eq!(claims.aud, "rampart");
assert_eq!(claims.role, "admin");
assert!(claims.exp > claims.iat);
}
#[test]
fn verify_rejects_wrong_audience() {
let secret = valid_secret();
let token = create_token(&secret, 3600, "rampart").expect("token creation should succeed");
assert!(verify_token(&token, &secret, "other").is_err());
}
#[test]
fn verify_rejects_wrong_secret() {
let secret = valid_secret();
let other_secret = "another-test-secret-also-32-bytes-long!".to_string();
let token = create_token(&secret, 3600, "rampart").expect("token creation should succeed");
assert!(verify_token(&token, &other_secret, "rampart").is_err());
}
}

View file

@ -1,8 +1,11 @@
use axum::{
Router, middleware,
Router,
http::HeaderValue,
middleware,
routing::{get, post},
};
use std::sync::Arc;
use dashmap::DashMap;
use std::{net::IpAddr, sync::Arc, time::Instant};
use tower_http::cors::CorsLayer;
use tracing_subscriber::EnvFilter;
@ -13,7 +16,10 @@ mod sync;
pub struct AppState {
pub redis_client: redis::Client,
pub jwt_secret: String,
pub jwt_audience: String,
pub jwt_expiration: u64,
pub api_password: String,
pub login_limiter: DashMap<IpAddr, (Instant, u32)>,
}
#[tokio::main]
@ -26,15 +32,27 @@ async fn main() -> anyhow::Result<()> {
let redis_client = redis::Client::open(redis_url)?;
let jwt_secret = std::env::var("JWT_SECRET").map_err(|_| anyhow::anyhow!("JWT_SECRET must be set"))?;
if jwt_secret.len() < 32 {
return Err(anyhow::anyhow!("JWT_SECRET must be at least 32 bytes"));
}
let jwt_audience = std::env::var("JWT_AUDIENCE").unwrap_or_else(|_| "rampart".to_string());
let jwt_expiration = std::env::var("JWT_EXPIRATION_SECS")
.unwrap_or_else(|_| "86400".to_string())
.parse::<u64>()
.map_err(|_| anyhow::anyhow!("JWT_EXPIRATION_SECS must be a valid u64"))?;
let api_password = std::env::var("API_PASSWORD").map_err(|_| anyhow::anyhow!("API_PASSWORD must be set"))?;
if api_password == "changeme" {
return Err(anyhow::anyhow!("API_PASSWORD must not be the default 'changeme'"));
}
let state = Arc::new(AppState {
redis_client,
jwt_secret,
jwt_audience,
jwt_expiration,
api_password,
login_limiter: DashMap::new(),
});
tokio::spawn(sync::heartbeat::start_heartbeat_check(state.clone()));
@ -52,15 +70,22 @@ async fn main() -> anyhow::Result<()> {
.route("/api/v1/nodes", get(api::nodes::list_nodes))
.route_layer(middleware::from_fn(auth::auth_middleware));
let cors = match std::env::var("CORS_ORIGIN") {
Ok(origin) if origin.is_empty() || origin == "*" => CorsLayer::new().allow_origin(tower_http::cors::Any),
Ok(origin) => CorsLayer::new().allow_origin(HeaderValue::from_str(&origin)?),
Err(_) => CorsLayer::new().allow_origin(HeaderValue::from_static("http://localhost:5173")),
};
let app = Router::new()
.merge(public)
.merge(protected)
.layer(CorsLayer::permissive())
.layer(cors)
.with_state(state);
let addr = "0.0.0.0:8080";
tracing::info!("Manager API listening on {addr}");
let listener = tokio::net::TcpListener::bind(addr).await?;
let app = app.into_make_service_with_connect_info::<std::net::SocketAddr>();
axum::serve(listener, app).await?;
Ok(())
}