use crate::{Query, Response, Record};
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
struct CacheEntry {
response: Response,
inserted_at: Instant,
expires_at: Instant,
original_ttl: Duration,
}
#[derive(Debug)]
pub struct DnsCache {
cache: Arc<RwLock<HashMap<CacheKey, CacheEntry>>>,
max_ttl: Duration,
stats: Arc<RwLock<CacheStats>>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct CacheKey {
name: String,
qtype: u16,
qclass: u16,
}
#[derive(Debug, Default, Clone)]
pub struct CacheStats {
pub hits: u64,
pub misses: u64,
pub inserts: u64,
pub evictions: u64,
pub current_size: usize,
}
impl CacheKey {
fn from_query(query: &Query) -> Self {
Self {
name: query.name.to_lowercase(),
qtype: query.qtype.into(),
qclass: query.qclass.into(),
}
}
}
impl DnsCache {
pub fn new(max_ttl: Duration) -> Self {
Self {
cache: Arc::new(RwLock::new(HashMap::new())),
max_ttl,
stats: Arc::new(RwLock::new(CacheStats::default())),
}
}
pub fn get(&self, query: &Query) -> Option<Response> {
let key = CacheKey::from_query(query);
let now = Instant::now();
let cache = self.cache.read().ok()?;
if let Some(entry) = cache.get(&key) {
if now < entry.expires_at {
if let Ok(mut stats) = self.stats.write() {
stats.hits += 1;
}
let mut response = entry.response.clone();
let remaining_ttl = entry.expires_at.duration_since(now);
for record in &mut response.answers {
record.ttl = remaining_ttl.as_secs() as u32;
}
for record in &mut response.authorities {
record.ttl = remaining_ttl.as_secs() as u32;
}
for record in &mut response.additionals {
record.ttl = remaining_ttl.as_secs() as u32;
}
return Some(response);
}
}
if let Ok(mut stats) = self.stats.write() {
stats.misses += 1;
}
None
}
pub fn insert(&self, query: Query, response: Response) {
let key = CacheKey::from_query(&query);
let now = Instant::now();
let ttl = self.calculate_ttl(&response);
if ttl.is_zero() {
return; }
let entry = CacheEntry {
response,
inserted_at: now,
expires_at: now + ttl,
original_ttl: ttl,
};
if let Ok(mut cache) = self.cache.write() {
let is_new = !cache.contains_key(&key);
cache.insert(key, entry);
if let Ok(mut stats) = self.stats.write() {
stats.inserts += 1;
if is_new {
stats.current_size = cache.len();
}
}
}
}
fn calculate_ttl(&self, response: &Response) -> Duration {
let mut min_ttl = self.max_ttl;
for record in &response.answers {
let record_ttl = Duration::from_secs(record.ttl as u64);
if record_ttl < min_ttl {
min_ttl = record_ttl;
}
}
for record in &response.authorities {
let record_ttl = Duration::from_secs(record.ttl as u64);
if record_ttl < min_ttl {
min_ttl = record_ttl;
}
}
for record in &response.additionals {
let record_ttl = Duration::from_secs(record.ttl as u64);
if record_ttl < min_ttl {
min_ttl = record_ttl;
}
}
min_ttl.min(self.max_ttl)
}
pub fn cleanup_expired(&self) {
let now = Instant::now();
let mut evicted_count = 0;
if let Ok(mut cache) = self.cache.write() {
let original_size = cache.len();
cache.retain(|_, entry| {
let keep = now < entry.expires_at;
if !keep {
evicted_count += 1;
}
keep
});
if let Ok(mut stats) = self.stats.write() {
stats.evictions += evicted_count;
stats.current_size = cache.len();
}
}
}
pub fn clear(&self) {
if let Ok(mut cache) = self.cache.write() {
cache.clear();
if let Ok(mut stats) = self.stats.write() {
stats.current_size = 0;
}
}
}
pub fn size(&self) -> usize {
self.cache.read().map(|cache| cache.len()).unwrap_or(0)
}
pub fn stats(&self) -> CacheStats {
self.stats.read().map(|stats| stats.clone()).unwrap_or_default()
}
pub fn hit_rate(&self) -> f64 {
let stats = self.stats();
let total = stats.hits + stats.misses;
if total == 0 {
0.0
} else {
stats.hits as f64 / total as f64
}
}
pub fn contains(&self, query: &Query) -> bool {
let key = CacheKey::from_query(query);
let now = Instant::now();
if let Ok(cache) = self.cache.read() {
if let Some(entry) = cache.get(&key) {
return now < entry.expires_at;
}
}
false
}
pub fn remove(&self, query: &Query) -> bool {
let key = CacheKey::from_query(query);
if let Ok(mut cache) = self.cache.write() {
let removed = cache.remove(&key).is_some();
if removed {
if let Ok(mut stats) = self.stats.write() {
stats.current_size = cache.len();
}
}
return removed;
}
false
}
pub fn get_cached_queries(&self) -> Vec<Query> {
if let Ok(cache) = self.cache.read() {
cache.keys().map(|key| Query {
name: key.name.clone(),
qtype: key.qtype.into(),
qclass: key.qclass.into(),
}).collect()
} else {
Vec::new()
}
}
pub fn set_max_ttl(&mut self, max_ttl: Duration) {
self.max_ttl = max_ttl;
}
pub fn max_ttl(&self) -> Duration {
self.max_ttl
}
}
pub struct CacheCleanupTask {
cache: Arc<DnsCache>,
interval: Duration,
}
impl CacheCleanupTask {
pub fn new(cache: Arc<DnsCache>, interval: Duration) -> Self {
Self { cache, interval }
}
pub async fn start(self) {
let mut interval_timer = tokio::time::interval(self.interval);
loop {
interval_timer.tick().await;
self.cache.cleanup_expired();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::*;
use std::net::Ipv4Addr;
fn create_test_query() -> Query {
Query {
name: "example.com".to_string(),
qtype: RecordType::A,
qclass: QClass::IN,
}
}
fn create_test_response() -> Response {
Response {
id: 12345,
flags: Flags::default(),
queries: vec![create_test_query()],
answers: vec![Record {
name: "example.com".to_string(),
rtype: RecordType::A,
class: QClass::IN,
ttl: 300,
data: RecordData::A(Ipv4Addr::new(93, 184, 216, 34)),
}],
authorities: vec![],
additionals: vec![],
}
}
#[test]
fn test_cache_insert_and_get() {
let cache = DnsCache::new(Duration::from_secs(3600));
let query = create_test_query();
let response = create_test_response();
cache.insert(query.clone(), response.clone());
let cached_response = cache.get(&query);
assert!(cached_response.is_some());
let cached = cached_response.unwrap();
assert_eq!(cached.answers.len(), 1);
assert_eq!(cached.answers[0].name, "example.com");
}
#[test]
fn test_cache_expiration() {
let cache = DnsCache::new(Duration::from_millis(100));
let query = create_test_query();
let mut response = create_test_response();
response.answers[0].ttl = 0;
cache.insert(query.clone(), response);
let cached_response = cache.get(&query);
assert!(cached_response.is_none());
}
}