use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::RwLock;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProxyEntry {
pub url: String,
#[serde(default)]
pub region: Option<String>,
#[serde(default)]
pub label: Option<String>,
}
impl ProxyEntry {
pub fn new(url: impl Into<String>) -> Self {
Self {
url: url.into(),
region: None,
label: None,
}
}
pub fn with_region(mut self, region: impl Into<String>) -> Self {
self.region = Some(region.into());
self
}
pub fn with_label(mut self, label: impl Into<String>) -> Self {
self.label = Some(label.into());
self
}
}
pub struct ProxyPool {
entries: Vec<ProxyEntry>,
health: RwLock<HashMap<String, HealthState>>,
sticky: RwLock<HashMap<String, usize>>,
cursor: std::sync::atomic::AtomicUsize,
cooldown: Duration,
}
#[derive(Debug, Clone)]
struct HealthState {
failures: u32,
last_failure: Instant,
}
impl ProxyPool {
pub fn new(entries: Vec<ProxyEntry>) -> Self {
Self {
entries,
health: RwLock::new(HashMap::new()),
sticky: RwLock::new(HashMap::new()),
cursor: std::sync::atomic::AtomicUsize::new(0),
cooldown: Duration::from_secs(60),
}
}
pub fn with_cooldown(mut self, cooldown: Duration) -> Self {
self.cooldown = cooldown;
self
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn for_domain(&self, domain: &str, region: Option<&str>) -> Option<ProxyEntry> {
if self.entries.is_empty() {
return None;
}
if let Some(&idx) = self
.sticky
.read()
.unwrap_or_else(|e| e.into_inner())
.get(domain)
{
if idx < self.entries.len() && !self.is_in_cooldown(&self.entries[idx].url) {
let e = &self.entries[idx];
if region.is_none() || e.region.as_deref() == region {
return Some(e.clone());
}
}
}
let n = self.entries.len();
for offset in 0..n {
let idx = (self.cursor.load(std::sync::atomic::Ordering::SeqCst) + offset) % n;
let e = &self.entries[idx];
if region.is_some() && e.region.as_deref() != region {
continue;
}
if self.is_in_cooldown(&e.url) {
continue;
}
self.cursor.store(idx + 1, std::sync::atomic::Ordering::SeqCst);
self.sticky
.write()
.unwrap_or_else(|e| e.into_inner())
.insert(domain.to_string(), idx);
return Some(e.clone());
}
Some(self.entries[0].clone())
}
fn is_in_cooldown(&self, url: &str) -> bool {
let map = self.health.read().unwrap_or_else(|e| e.into_inner());
match map.get(url) {
Some(s) if s.failures >= 3 => s.last_failure.elapsed() < self.cooldown,
_ => false,
}
}
pub fn record_failure(&self, url: &str) {
let mut map = self.health.write().unwrap_or_else(|e| e.into_inner());
let s = map.entry(url.to_string()).or_insert(HealthState {
failures: 0,
last_failure: Instant::now(),
});
s.failures = s.failures.saturating_add(1);
s.last_failure = Instant::now();
}
pub fn record_success(&self, url: &str) {
let mut map = self.health.write().unwrap_or_else(|e| e.into_inner());
if let Some(s) = map.get_mut(url) {
s.failures = 0;
}
}
pub fn health_snapshot(&self) -> Vec<(String, u32)> {
let map = self.health.read().unwrap_or_else(|e| e.into_inner());
self.entries
.iter()
.map(|e| (e.url.clone(), map.get(&e.url).map(|s| s.failures).unwrap_or(0)))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn pool() -> ProxyPool {
ProxyPool::new(vec![
ProxyEntry::new("http://p1.example:8080").with_region("us"),
ProxyEntry::new("http://p2.example:8080").with_region("us"),
ProxyEntry::new("http://p3.example:8080").with_region("de"),
])
}
#[test]
fn for_domain_returns_sticky_pick() {
let p = pool();
let first = p.for_domain("example.com", None).expect("non-empty");
let second = p.for_domain("example.com", None).expect("sticky");
assert_eq!(first.url, second.url);
}
#[test]
fn region_filter_excludes_other_regions() {
let p = pool();
let de = p.for_domain("example.com", Some("de")).expect("de");
assert_eq!(de.region.as_deref(), Some("de"));
}
#[test]
fn empty_pool_returns_none() {
let p = ProxyPool::new(Vec::new());
assert!(p.for_domain("any", None).is_none());
}
#[test]
fn cooldown_after_three_failures_excludes_entry() {
let p = ProxyPool::new(vec![
ProxyEntry::new("http://bad.example:8080").with_region("x"),
ProxyEntry::new("http://good.example:8080").with_region("x"),
]);
let _ = p.for_domain("d1", Some("x"));
p.record_failure("http://bad.example:8080");
p.record_failure("http://bad.example:8080");
p.record_failure("http://bad.example:8080");
let pick = p.for_domain("d2", Some("x")).unwrap();
assert_ne!(pick.url, "http://bad.example:8080");
}
#[test]
fn record_success_resets_failures() {
let p = ProxyPool::new(vec![ProxyEntry::new("http://x.example:8080")]);
p.record_failure("http://x.example:8080");
p.record_failure("http://x.example:8080");
p.record_success("http://x.example:8080");
let snap = p.health_snapshot();
assert_eq!(snap[0].1, 0);
}
#[test]
fn health_snapshot_contains_every_entry() {
let p = pool();
let snap = p.health_snapshot();
assert_eq!(snap.len(), p.len());
}
}