use std::collections::BTreeMap;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use crate::types::VendorId;
const AKAMAI_STICKY_TTL: Duration = Duration::from_mins(30);
const CLOUDFLARE_STICKY_TTL: Duration = Duration::from_mins(5);
const IMPERVA_STICKY_TTL: Duration = Duration::from_mins(15);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "mode")]
pub enum StickinessPolicy {
StickyForever,
StickyForTtl {
#[serde(with = "serde_duration_secs")]
ttl: Duration,
},
StickyForRequestCount {
max_requests: u32,
},
FreshPerRequest,
FreshPerDomain,
}
impl std::fmt::Display for StickinessPolicy {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::StickyForever => f.write_str("sticky_forever"),
Self::StickyForTtl { ttl } => write!(f, "sticky_for_ttl({}s)", ttl.as_secs()),
Self::StickyForRequestCount { max_requests } => {
write!(f, "sticky_for_request_count({max_requests})")
}
Self::FreshPerDomain => f.write_str("fresh_per_domain"),
Self::FreshPerRequest => f.write_str("fresh_per_request"),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct VendorStickinessMap(BTreeMap<VendorId, StickinessPolicy>);
impl VendorStickinessMap {
#[must_use]
pub const fn new() -> Self {
Self(BTreeMap::new())
}
#[must_use]
pub fn with_builtin_defaults() -> Self {
let mut entries = BTreeMap::new();
entries.insert(
VendorId::Akamai,
StickinessPolicy::StickyForTtl {
ttl: AKAMAI_STICKY_TTL,
},
);
entries.insert(
VendorId::Cloudflare,
StickinessPolicy::StickyForTtl {
ttl: CLOUDFLARE_STICKY_TTL,
},
);
entries.insert(VendorId::DataDome, StickinessPolicy::FreshPerRequest);
entries.insert(
VendorId::Imperva,
StickinessPolicy::StickyForTtl {
ttl: IMPERVA_STICKY_TTL,
},
);
entries.insert(VendorId::Kasada, StickinessPolicy::FreshPerDomain);
entries.insert(VendorId::PerimeterX, StickinessPolicy::FreshPerDomain);
Self(entries)
}
#[must_use]
pub fn for_vendor(&self, vendor: VendorId) -> StickinessPolicy {
self.0
.get(&vendor)
.copied()
.unwrap_or(StickinessPolicy::FreshPerRequest)
}
#[must_use]
pub fn with_override(mut self, vendor: VendorId, policy: StickinessPolicy) -> Self {
self.0.insert(vendor, policy);
self
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
#[must_use]
pub fn len(&self) -> usize {
self.0.len()
}
pub fn iter(&self) -> impl Iterator<Item = (VendorId, StickinessPolicy)> + '_ {
self.0.iter().map(|(v, p)| (*v, *p))
}
}
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,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
mod tests {
use super::*;
#[test]
fn new_is_empty() {
let map = VendorStickinessMap::new();
assert!(map.is_empty());
assert_eq!(map.len(), 0);
}
#[test]
fn default_is_empty() {
let map = VendorStickinessMap::default();
assert!(map.is_empty());
}
#[test]
fn for_vendor_unknown_returns_fresh_per_request() {
let map = VendorStickinessMap::new();
assert_eq!(
map.for_vendor(VendorId::Unknown),
StickinessPolicy::FreshPerRequest
);
assert_eq!(
map.for_vendor(VendorId::Akamai),
StickinessPolicy::FreshPerRequest
);
}
#[test]
fn with_override_inserts_entry() {
let map = VendorStickinessMap::new()
.with_override(VendorId::Akamai, StickinessPolicy::StickyForever);
assert_eq!(map.len(), 1);
assert_eq!(
map.for_vendor(VendorId::Akamai),
StickinessPolicy::StickyForever
);
}
#[test]
fn with_override_replaces_existing_entry() {
let map = VendorStickinessMap::new()
.with_override(VendorId::Akamai, StickinessPolicy::StickyForever)
.with_override(VendorId::Akamai, StickinessPolicy::FreshPerDomain);
assert_eq!(map.len(), 1);
assert_eq!(
map.for_vendor(VendorId::Akamai),
StickinessPolicy::FreshPerDomain
);
}
#[test]
fn built_in_defaults_akamai_is_30min_sticky() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::Akamai),
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(30)
}
);
}
#[test]
fn built_in_defaults_cloudflare_is_5min_sticky() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::Cloudflare),
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(5)
}
);
}
#[test]
fn built_in_defaults_imperva_is_15min_sticky() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::Imperva),
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(15)
}
);
}
#[test]
fn built_in_defaults_perimeter_x_is_fresh_per_domain() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::PerimeterX),
StickinessPolicy::FreshPerDomain
);
}
#[test]
fn built_in_defaults_kasada_is_fresh_per_domain() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::Kasada),
StickinessPolicy::FreshPerDomain
);
}
#[test]
fn built_in_defaults_data_dome_is_fresh_per_request() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::DataDome),
StickinessPolicy::FreshPerRequest
);
}
#[test]
fn built_in_defaults_unknown_vendor_falls_back_to_fresh_per_request() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(
map.for_vendor(VendorId::Unknown),
StickinessPolicy::FreshPerRequest
);
assert_eq!(
map.for_vendor(VendorId::Hcaptcha),
StickinessPolicy::FreshPerRequest
);
assert_eq!(
map.for_vendor(VendorId::ShapeSecurity),
StickinessPolicy::FreshPerRequest
);
}
#[test]
fn built_in_defaults_has_six_entries() {
let map = VendorStickinessMap::with_builtin_defaults();
assert_eq!(map.len(), 6);
}
#[test]
fn built_in_defaults_iterates_in_sorted_order() {
let map = VendorStickinessMap::with_builtin_defaults();
let entries: Vec<_> = map.iter().map(|(v, _)| v).collect();
assert_eq!(
entries,
vec![
VendorId::Akamai,
VendorId::Cloudflare,
VendorId::DataDome,
VendorId::PerimeterX,
VendorId::Kasada,
VendorId::Imperva,
]
);
}
#[test]
fn override_chained_before_builtins_replaces_entry() {
let map = VendorStickinessMap::with_builtin_defaults()
.with_override(VendorId::Akamai, StickinessPolicy::StickyForever);
assert_eq!(
map.for_vendor(VendorId::Akamai),
StickinessPolicy::StickyForever
);
assert_eq!(
map.for_vendor(VendorId::Cloudflare),
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(5)
}
);
}
#[test]
fn stickiness_policy_is_copy() {
let policy = StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(30),
};
let copy = policy;
assert_eq!(policy, copy);
}
#[test]
fn stickiness_policy_is_hash_eq() {
use std::collections::HashSet;
let mut set = HashSet::new();
set.insert(StickinessPolicy::StickyForever);
set.insert(StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(30),
});
set.insert(StickinessPolicy::FreshPerRequest);
assert_eq!(set.len(), 3);
assert!(set.contains(&StickinessPolicy::StickyForever));
}
#[test]
fn stickiness_policy_display_matches_snake_case_label() {
assert_eq!(
format!("{}", StickinessPolicy::StickyForever),
"sticky_forever"
);
assert_eq!(
format!(
"{}",
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(1)
}
),
"sticky_for_ttl(60s)"
);
assert_eq!(
format!(
"{}",
StickinessPolicy::StickyForRequestCount { max_requests: 5 }
),
"sticky_for_request_count(5)"
);
assert_eq!(
format!("{}", StickinessPolicy::FreshPerDomain),
"fresh_per_domain"
);
assert_eq!(
format!("{}", StickinessPolicy::FreshPerRequest),
"fresh_per_request"
);
}
#[test]
fn stickiness_policy_round_trips_through_json() {
let policies = [
StickinessPolicy::StickyForever,
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(30),
},
StickinessPolicy::StickyForRequestCount { max_requests: 7 },
StickinessPolicy::FreshPerDomain,
StickinessPolicy::FreshPerRequest,
];
for policy in policies {
let json = serde_json::to_string(&policy).expect("serialize");
let parsed: StickinessPolicy = serde_json::from_str(&json).expect("deserialize");
assert_eq!(parsed, policy, "round-trip for {policy:?}");
}
}
#[test]
fn stickiness_policy_round_trips_through_toml() {
let policies = [
StickinessPolicy::StickyForever,
StickinessPolicy::StickyForTtl {
ttl: Duration::from_mins(30),
},
StickinessPolicy::StickyForRequestCount { max_requests: 7 },
StickinessPolicy::FreshPerDomain,
StickinessPolicy::FreshPerRequest,
];
for policy in policies {
let toml_str = toml::to_string(&policy).expect("serialize toml");
let parsed: StickinessPolicy = toml::from_str(&toml_str).expect("deserialize toml");
assert_eq!(parsed, policy, "round-trip for {policy:?}");
}
}
#[test]
fn vendor_stickiness_map_round_trips_through_json() {
let map = VendorStickinessMap::with_builtin_defaults()
.with_override(VendorId::Akamai, StickinessPolicy::StickyForever);
let json = serde_json::to_string(&map).expect("serialize");
let parsed: VendorStickinessMap = serde_json::from_str(&json).expect("deserialize");
assert_eq!(parsed, map);
}
#[test]
fn vendor_stickiness_map_round_trips_through_toml() {
let map = VendorStickinessMap::with_builtin_defaults()
.with_override(VendorId::DataDome, StickinessPolicy::StickyForever);
let toml_str = toml::to_string(&map).expect("serialize toml");
let parsed: VendorStickinessMap = toml::from_str(&toml_str).expect("deserialize toml");
assert_eq!(parsed, map);
}
#[test]
fn vendor_stickiness_map_transparent_serde_orders_by_vendor_id() {
let map = VendorStickinessMap::with_builtin_defaults();
let json = serde_json::to_string(&map).expect("serialize");
let akamai_pos = json.find("\"akamai\"").expect("akamai present");
let cloudflare_pos = json.find("\"cloudflare\"").expect("cloudflare present");
assert!(
akamai_pos < cloudflare_pos,
"expected sorted order on wire: {json}"
);
}
#[test]
fn stickiness_policy_for_request_count_variant_exists() {
let policy = StickinessPolicy::StickyForRequestCount { max_requests: 5 };
let json = serde_json::to_string(&policy).expect("serialize");
let parsed: StickinessPolicy = serde_json::from_str(&json).expect("deserialize");
assert_eq!(parsed, policy);
}
}