use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use serde::{Deserialize, Serialize};
use tokio::sync::RwLock;
use uuid::Uuid;
use crate::stickiness::{StickinessPolicy, VendorStickinessMap};
use crate::types::VendorId;
const DEFAULT_TTL_SECS: u64 = 300;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "mode")]
#[non_exhaustive]
pub enum StickyPolicy {
#[default]
Disabled,
Domain {
#[serde(with = "serde_duration_secs")]
ttl: Duration,
},
}
impl StickyPolicy {
#[must_use]
pub const fn domain(ttl: Duration) -> Self {
Self::Domain { ttl }
}
#[must_use]
pub const fn domain_default() -> Self {
Self::Domain {
ttl: Duration::from_secs(DEFAULT_TTL_SECS),
}
}
#[must_use]
pub const fn is_disabled(&self) -> bool {
matches!(self, Self::Disabled)
}
}
#[derive(Debug, Clone)]
struct ProxySession {
proxy_id: Uuid,
bound_at: Instant,
ttl: Duration,
}
impl ProxySession {
fn is_expired(&self) -> bool {
self.bound_at.elapsed() >= self.ttl
}
}
#[derive(Debug, Clone)]
pub struct SessionMap {
inner: Arc<RwLock<HashMap<String, ProxySession>>>,
}
impl Default for SessionMap {
fn default() -> Self {
Self::new()
}
}
impl SessionMap {
#[must_use]
pub fn new() -> Self {
Self {
inner: Arc::new(RwLock::new(HashMap::new())),
}
}
#[must_use]
pub fn lookup(&self, key: &str) -> Option<Uuid> {
let guard = self.inner.try_read().ok()?;
guard
.get(key)
.filter(|s| !s.is_expired())
.map(|s| s.proxy_id)
}
pub fn bind(&self, key: &str, proxy_id: Uuid, ttl: Duration) {
let session = ProxySession {
proxy_id,
bound_at: Instant::now(),
ttl,
};
if let Ok(mut guard) = self.inner.try_write() {
guard.insert(key.to_string(), session);
}
}
#[must_use]
pub fn purge_expired(&self) -> usize {
let Ok(mut guard) = self.inner.try_write() else {
return 0;
};
let before = guard.len();
guard.retain(|_, s| !s.is_expired());
before - guard.len()
}
pub fn unbind(&self, key: &str) {
if let Ok(mut guard) = self.inner.try_write() {
guard.remove(key);
}
}
#[must_use]
pub fn active_count(&self) -> usize {
let Ok(guard) = self.inner.try_read() else {
return 0;
};
guard.values().filter(|s| !s.is_expired()).count()
}
#[must_use]
pub fn acquire_session(
&self,
domain: &str,
vendor: VendorId,
policy_map: &VendorStickinessMap,
) -> SessionDecision {
let policy = policy_map.for_vendor(vendor);
let ttl = match policy {
StickinessPolicy::StickyForever => Some(Duration::MAX),
StickinessPolicy::StickyForTtl { ttl } => Some(ttl),
StickinessPolicy::StickyForRequestCount { .. }
| StickinessPolicy::FreshPerRequest
| StickinessPolicy::FreshPerDomain => None,
};
let evict_for_fresh_domain = matches!(policy, StickinessPolicy::FreshPerDomain);
ttl.map_or_else(
|| {
if evict_for_fresh_domain {
self.unbind(domain);
}
SessionDecision::AcquireFresh
},
|ttl| {
self.lookup(domain).map_or(
SessionDecision::AcquireAndBind(ttl),
SessionDecision::UseSticky,
)
},
)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SessionDecision {
UseSticky(Uuid),
AcquireFresh,
AcquireAndBind(Duration),
}
mod serde_duration_secs {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::time::Duration;
pub fn serialize<S: Serializer>(d: &Duration, s: S) -> Result<S::Ok, S::Error> {
d.as_secs().serialize(s)
}
pub fn deserialize<'de, D: Deserializer<'de>>(d: D) -> Result<Duration, D::Error> {
Ok(Duration::from_secs(u64::deserialize(d)?))
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn same_domain_returns_same_proxy() {
let map = SessionMap::new();
let id = Uuid::new_v4();
map.bind("example.com", id, Duration::from_mins(1));
assert_eq!(map.lookup("example.com"), Some(id));
assert_eq!(map.lookup("example.com"), Some(id));
}
#[test]
fn different_domains_independent() {
let map = SessionMap::new();
let id_a = Uuid::new_v4();
let id_b = Uuid::new_v4();
map.bind("a.com", id_a, Duration::from_mins(1));
map.bind("b.com", id_b, Duration::from_mins(1));
assert_eq!(map.lookup("a.com"), Some(id_a));
assert_eq!(map.lookup("b.com"), Some(id_b));
}
#[test]
fn expired_session_returns_none() {
let map = SessionMap::new();
let id = Uuid::new_v4();
map.bind("example.com", id, Duration::ZERO);
std::thread::sleep(Duration::from_millis(1));
assert_eq!(map.lookup("example.com"), None);
}
#[test]
fn purge_removes_expired() {
let map = SessionMap::new();
map.bind("expired.com", Uuid::new_v4(), Duration::ZERO);
map.bind("active.com", Uuid::new_v4(), Duration::from_mins(5));
std::thread::sleep(Duration::from_millis(1));
let removed = map.purge_expired();
assert_eq!(removed, 1);
assert_eq!(map.active_count(), 1);
}
#[test]
fn unbind_removes_session() {
let map = SessionMap::new();
map.bind("example.com", Uuid::new_v4(), Duration::from_mins(1));
map.unbind("example.com");
assert_eq!(map.lookup("example.com"), None);
}
#[test]
fn rebind_overwrites_previous() {
let map = SessionMap::new();
let old_id = Uuid::new_v4();
let new_id = Uuid::new_v4();
map.bind("example.com", old_id, Duration::from_mins(1));
map.bind("example.com", new_id, Duration::from_mins(1));
assert_eq!(map.lookup("example.com"), Some(new_id));
}
#[test]
fn policy_domain_default_ttl() {
let policy = StickyPolicy::domain_default();
assert!(matches!(policy, StickyPolicy::Domain { ttl } if ttl == Duration::from_mins(5)));
}
#[test]
fn policy_disabled_by_default() {
let policy = StickyPolicy::default();
assert!(policy.is_disabled());
}
#[test]
fn policy_serde_roundtrip() -> std::result::Result<(), Box<dyn std::error::Error>> {
let policy = StickyPolicy::domain(Duration::from_mins(2));
let json = serde_json::to_string(&policy)?;
let back: StickyPolicy = serde_json::from_str(&json)?;
assert!(matches!(back, StickyPolicy::Domain { ttl } if ttl == Duration::from_mins(2)));
Ok(())
}
use crate::stickiness::{StickinessPolicy, VendorStickinessMap};
fn vendor_policy_map() -> VendorStickinessMap {
VendorStickinessMap::with_builtin_defaults()
}
#[test]
fn acquire_session_akamai_no_binding_returns_acquire_and_bind_30min() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let decision = map.acquire_session("example.com", VendorId::Akamai, &policy);
assert_eq!(
decision,
SessionDecision::AcquireAndBind(Duration::from_mins(30))
);
}
#[test]
fn acquire_session_akamai_with_existing_binding_returns_sticky() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let proxy_id = Uuid::new_v4();
map.bind("example.com", proxy_id, Duration::from_mins(30));
let decision = map.acquire_session("example.com", VendorId::Akamai, &policy);
assert_eq!(decision, SessionDecision::UseSticky(proxy_id));
}
#[test]
fn acquire_session_akamai_100_calls_within_ttl_return_same_proxy() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let proxy_id = Uuid::new_v4();
map.bind("example.com", proxy_id, Duration::from_mins(30));
for _ in 0..100 {
assert_eq!(
map.acquire_session("example.com", VendorId::Akamai, &policy),
SessionDecision::UseSticky(proxy_id)
);
}
}
#[test]
fn acquire_session_akamai_expired_binding_returns_acquire_and_bind() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let stale_id = Uuid::new_v4();
map.bind("example.com", stale_id, Duration::ZERO);
std::thread::sleep(Duration::from_millis(1));
let decision = map.acquire_session("example.com", VendorId::Akamai, &policy);
assert_eq!(
decision,
SessionDecision::AcquireAndBind(Duration::from_mins(30))
);
}
#[test]
fn acquire_session_cloudflare_no_binding_returns_acquire_and_bind_5min() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let decision = map.acquire_session("example.com", VendorId::Cloudflare, &policy);
assert_eq!(
decision,
SessionDecision::AcquireAndBind(Duration::from_mins(5))
);
}
#[test]
fn acquire_session_imperva_no_binding_returns_acquire_and_bind_15min() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let decision = map.acquire_session("example.com", VendorId::Imperva, &policy);
assert_eq!(
decision,
SessionDecision::AcquireAndBind(Duration::from_mins(15))
);
}
#[test]
fn acquire_session_data_dome_always_returns_acquire_fresh() {
let map = SessionMap::new();
let policy = vendor_policy_map();
map.bind("example.com", Uuid::new_v4(), Duration::from_hours(1));
let decision = map.acquire_session("example.com", VendorId::DataDome, &policy);
assert_eq!(decision, SessionDecision::AcquireFresh);
}
#[test]
fn acquire_session_perimeter_x_evicts_existing_binding() {
let map = SessionMap::new();
let policy = vendor_policy_map();
map.bind("example.com", Uuid::new_v4(), Duration::from_hours(1));
let decision = map.acquire_session("example.com", VendorId::PerimeterX, &policy);
assert_eq!(decision, SessionDecision::AcquireFresh);
assert_eq!(map.lookup("example.com"), None);
}
#[test]
fn acquire_session_perimeter_x_no_existing_binding_returns_fresh() {
let map = SessionMap::new();
let policy = vendor_policy_map();
let decision = map.acquire_session("example.com", VendorId::PerimeterX, &policy);
assert_eq!(decision, SessionDecision::AcquireFresh);
}
#[test]
fn acquire_session_kasada_evicts_existing_binding() {
let map = SessionMap::new();
let policy = vendor_policy_map();
map.bind("example.com", Uuid::new_v4(), Duration::from_hours(1));
let decision = map.acquire_session("example.com", VendorId::Kasada, &policy);
assert_eq!(decision, SessionDecision::AcquireFresh);
assert_eq!(map.lookup("example.com"), None);
}
#[test]
fn acquire_session_unknown_vendor_defaults_to_fresh() {
let map = SessionMap::new();
let policy = vendor_policy_map();
map.bind("example.com", Uuid::new_v4(), Duration::from_hours(1));
let decision = map.acquire_session("example.com", VendorId::Unknown, &policy);
assert_eq!(decision, SessionDecision::AcquireFresh);
assert!(map.lookup("example.com").is_some());
}
#[test]
fn acquire_session_sticky_forever_uses_max_duration() {
let map = SessionMap::new();
let custom = VendorStickinessMap::new()
.with_override(VendorId::Akamai, StickinessPolicy::StickyForever);
let decision = map.acquire_session("example.com", VendorId::Akamai, &custom);
assert_eq!(decision, SessionDecision::AcquireAndBind(Duration::MAX));
}
#[test]
fn acquire_session_sticky_for_request_count_treated_as_fresh() {
let map = SessionMap::new();
let custom = VendorStickinessMap::new().with_override(
VendorId::Akamai,
StickinessPolicy::StickyForRequestCount { max_requests: 5 },
);
let decision = map.acquire_session("example.com", VendorId::Akamai, &custom);
assert_eq!(decision, SessionDecision::AcquireFresh);
}
#[test]
fn acquire_session_sticky_forever_uses_existing_binding() {
let map = SessionMap::new();
let custom = VendorStickinessMap::new()
.with_override(VendorId::Akamai, StickinessPolicy::StickyForever);
let proxy_id = Uuid::new_v4();
map.bind("example.com", proxy_id, Duration::from_hours(1));
let decision = map.acquire_session("example.com", VendorId::Akamai, &custom);
assert_eq!(decision, SessionDecision::UseSticky(proxy_id));
}
#[test]
fn acquire_session_override_replaces_builtin_akamai_policy() {
let map = SessionMap::new();
let policy = vendor_policy_map().with_override(
VendorId::Akamai,
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(2),
},
);
let decision = map.acquire_session("example.com", VendorId::Akamai, &policy);
assert_eq!(
decision,
SessionDecision::AcquireAndBind(Duration::from_mins(2))
);
}
#[test]
fn acquire_session_empty_map_defaults_all_to_fresh() {
let map = SessionMap::new();
let empty = VendorStickinessMap::new();
map.bind("example.com", Uuid::new_v4(), Duration::from_hours(1));
for vendor in [
VendorId::Akamai,
VendorId::Cloudflare,
VendorId::DataDome,
VendorId::PerimeterX,
VendorId::Kasada,
VendorId::Imperva,
VendorId::Unknown,
VendorId::Hcaptcha,
] {
assert_eq!(
map.acquire_session("example.com", vendor, &empty),
SessionDecision::AcquireFresh,
"{vendor:?} should default to fresh when no entry exists"
);
}
}
}