feat: caching (#76)

This commit is contained in:
LowderPlay
2026-06-18 03:48:45 +05:00
committed by GitHub
parent be654b3b74
commit fbc0e7a235
7 changed files with 275 additions and 37 deletions
+1
View File
@@ -1,3 +1,4 @@
/target
.idea/
.env
database-cache/
+11 -6
View File
@@ -62,12 +62,17 @@ docker compose up --build
HTTP_PORT=80 docker compose up --build
```
Маршрутизация:
* `/api/v1/*`, `/agency/*`, `/whitelist/*` идут в Rocket
* остальные запросы идут в SvelteKit SSR
---
Маршрутизация:
* `/api/v1/*`, `/agency/*`, `/whitelist/*` идут в Rocket
* остальные запросы идут в SvelteKit SSR
Базы проверок кешируются в `database-cache/`. Путь можно изменить через
`DATABASE_CACHE_DIR`, основной интервал обновления — через
`DATABASE_INTERVAL_SECONDS`, а интервал повтора после ошибки скачивания — через
`DATABASE_RETRY_INTERVAL_SECONDS` (по умолчанию 300 секунд).
---
## Вклад
+4
View File
@@ -29,9 +29,13 @@ services:
ROCKET_ADDRESS: "::"
ROCKET_PORT: 8000
DATABASE_URL: "${DATABASE_URL}"
DATABASE_CACHE_DIR: /var/cache/cheburcheck/databases
DATABASE_RETRY_INTERVAL_SECONDS: "${DATABASE_RETRY_INTERVAL_SECONDS:-300}"
MQTT_ADMIN_TOKEN: "${MQTT_ADMIN_TOKEN}"
MQTT_HOST: rmqtt
MQTT_PORT: 11883
volumes:
- ./database-cache:/var/cache/cheburcheck/databases
expose:
- "8000"
+161
View File
@@ -0,0 +1,161 @@
use crate::Bases;
use crate::geoip::GeoIp;
use crate::lists::{CdnList, RuBlacklist};
use crate::updater::Updatable;
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
use std::fs;
use std::io;
use std::path::{Path, PathBuf};
use std::time::Duration;
const METADATA_FILE: &str = "metadata.json";
const GEO_ASN_FILE: &str = "geo-asn.mmdb";
const GEO_COUNTRY_FILE: &str = "geo-country.mmdb";
const GEO_CITY_FILE: &str = "geo-city.mmdb";
const RKN_NETS_FILE: &str = "rkn-nets.lst";
const RKN_DOMAINS_FILE: &str = "rkn-domains.lst";
const CDN_FILE: &str = "cdn.csv";
#[derive(Debug, Clone)]
pub struct DatabaseCache {
path: PathBuf,
}
#[derive(Deserialize, Serialize)]
struct CacheMetadata {
updated_at: DateTime<Utc>,
sources: Vec<CacheSourceMetadata>,
}
#[derive(Deserialize, Serialize)]
struct CacheSourceMetadata {
name: String,
url: String,
file: String,
bytes: usize,
}
impl DatabaseCache {
pub fn from_env() -> Self {
Self {
path: std::env::var("DATABASE_CACHE_DIR")
.map(PathBuf::from)
.unwrap_or_else(|_| PathBuf::from("database-cache")),
}
}
pub fn path(&self) -> &Path {
&self.path
}
pub fn load(&self) -> io::Result<Bases> {
Ok((
(
fs::read(self.path.join(GEO_ASN_FILE))?,
fs::read(self.path.join(GEO_COUNTRY_FILE))?,
fs::read(self.path.join(GEO_CITY_FILE))?,
),
(
VecDeque::from(fs::read(self.path.join(RKN_NETS_FILE))?),
VecDeque::from(fs::read(self.path.join(RKN_DOMAINS_FILE))?),
VecDeque::from(include_bytes!("../dist-domains.txt").to_vec()),
),
VecDeque::from(fs::read(self.path.join(CDN_FILE))?),
))
}
pub fn updated_at(&self) -> io::Result<DateTime<Utc>> {
let metadata = fs::read(self.path.join(METADATA_FILE))?;
let metadata: CacheMetadata = serde_json::from_slice(&metadata)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
Ok(metadata.updated_at)
}
pub fn refresh_delay(&self, interval: Duration) -> io::Result<Duration> {
let elapsed = Utc::now()
.signed_duration_since(self.updated_at()?)
.to_std()
.unwrap_or(Duration::ZERO);
Ok(interval.saturating_sub(elapsed))
}
pub fn store(&self, bases: &Bases) -> io::Result<()> {
fs::create_dir_all(&self.path)?;
let (geo_ip, ru_blacklist, cdn_list) = bases;
self.write_file(GEO_ASN_FILE, &geo_ip.0)?;
self.write_file(GEO_COUNTRY_FILE, &geo_ip.1)?;
self.write_file(GEO_CITY_FILE, &geo_ip.2)?;
self.write_vec_deque_file(RKN_NETS_FILE, &ru_blacklist.0)?;
self.write_vec_deque_file(RKN_DOMAINS_FILE, &ru_blacklist.1)?;
self.write_vec_deque_file(CDN_FILE, cdn_list)?;
let metadata = CacheMetadata {
updated_at: Utc::now(),
sources: vec![
CacheSourceMetadata {
name: "geo_asn".to_string(),
url: GeoIp::get_url("GEO_ASN", "https://git.io/GeoLite2-ASN.mmdb"),
file: GEO_ASN_FILE.to_string(),
bytes: geo_ip.0.len(),
},
CacheSourceMetadata {
name: "geo_country".to_string(),
url: GeoIp::get_url("GEO_COUNTRY", "https://git.io/GeoLite2-Country.mmdb"),
file: GEO_COUNTRY_FILE.to_string(),
bytes: geo_ip.1.len(),
},
CacheSourceMetadata {
name: "geo_city".to_string(),
url: GeoIp::get_url("GEO_CITY", "https://git.io/GeoLite2-City.mmdb"),
file: GEO_CITY_FILE.to_string(),
bytes: geo_ip.2.len(),
},
CacheSourceMetadata {
name: "rkn_nets".to_string(),
url: RuBlacklist::get_url(
"RKN_NETS",
"https://antifilter.network/download/ipsum.lst",
),
file: RKN_NETS_FILE.to_string(),
bytes: ru_blacklist.0.len(),
},
CacheSourceMetadata {
name: "rkn_domains".to_string(),
url: RuBlacklist::get_url(
"RKN_DOMAINS",
"https://antifilter.download/list/domains.lst",
),
file: RKN_DOMAINS_FILE.to_string(),
bytes: ru_blacklist.1.len(),
},
CacheSourceMetadata {
name: "cdn".to_string(),
url: CdnList::get_url(
"CDN_SOURCE",
"https://raw.githubusercontent.com/123jjck/cdn-ip-ranges/refs/heads/main/all/all.csv",
),
file: CDN_FILE.to_string(),
bytes: cdn_list.len(),
},
],
};
let metadata = serde_json::to_vec_pretty(&metadata)
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
self.write_file(METADATA_FILE, &metadata)
}
fn write_file(&self, file_name: &str, contents: &[u8]) -> io::Result<()> {
let path = self.path.join(file_name);
let tmp_path = self.path.join(format!("{file_name}.tmp"));
fs::write(&tmp_path, contents)?;
fs::rename(tmp_path, path)
}
fn write_vec_deque_file(&self, file_name: &str, contents: &VecDeque<u8>) -> io::Result<()> {
let bytes: Vec<u8> = contents.iter().copied().collect();
self.write_file(file_name, &bytes)
}
}
+1
View File
@@ -16,6 +16,7 @@ use thiserror::Error;
use tokio::sync::watch;
pub mod asn;
pub mod cache;
pub mod geoip;
pub mod lists;
pub mod resolver;
+93
View File
@@ -0,0 +1,93 @@
use log::{info, warn};
use querying::Checker;
use querying::cache::DatabaseCache;
use rocket::tokio;
use rocket::tokio::sync::RwLock;
use rocket::tokio::time;
use std::sync::Arc;
use std::time::Duration;
pub async fn start() -> Arc<RwLock<Checker>> {
let database_interval = Duration::from_secs(
std::env::var("DATABASE_INTERVAL_SECONDS")
.unwrap_or("21600".to_string())
.parse()
.unwrap(),
);
let database_retry_interval = Duration::from_secs(
std::env::var("DATABASE_RETRY_INTERVAL_SECONDS")
.unwrap_or("300".to_string())
.parse()
.unwrap(),
);
let cache = DatabaseCache::from_env();
let checker = Checker::new().await;
let first_refresh_delay = match cache.load() {
Ok(bases) => {
info!("Loading databases from {}", cache.path().display());
checker.update_all(bases).await;
info!("Loaded cached databases");
cache
.refresh_delay(database_interval)
.unwrap_or(Duration::ZERO)
}
Err(e) => {
warn!(
"Failed to load cached databases from {}: {}",
cache.path().display(),
e
);
Duration::ZERO
}
};
let checker = Arc::new(RwLock::new(checker));
spawn_refresh_loop(
checker.clone(),
cache,
database_interval,
database_retry_interval,
first_refresh_delay,
);
checker
}
fn spawn_refresh_loop(
checker: Arc<RwLock<Checker>>,
cache: DatabaseCache,
database_interval: Duration,
database_retry_interval: Duration,
first_refresh_delay: Duration,
) {
tokio::spawn(async move {
info!(
"Refreshing DB every {:?}; retrying failures after {:?}",
database_interval, database_retry_interval
);
if !first_refresh_delay.is_zero() {
info!("Next DB refresh in {:?}", first_refresh_delay);
time::sleep(first_refresh_delay).await;
}
loop {
info!("Updating all DBs");
let next_refresh = match Checker::download_all().await {
Ok(bases) => {
info!("Downloaded, updating...");
if let Err(e) = cache.store(&bases) {
log::error!("Failed to write database cache: {}", e);
}
checker.read().await.update_all(bases).await;
info!("Updated databases");
database_interval
}
Err(e) => {
log::error!("Failed to download all DBs: {}", e);
database_retry_interval
}
};
time::sleep(next_refresh).await;
}
});
}
+4 -31
View File
@@ -2,23 +2,20 @@
extern crate rocket;
mod agency;
mod api;
mod database_refresh;
mod db;
mod mqtt;
mod mqtt_auth;
mod whitelist;
use env_logger::Env;
use log::{LevelFilter, error, info};
use querying::Checker;
use log::{LevelFilter, error};
use rocket::fairing::AdHoc;
use rocket::http::Status;
use rocket::serde::json::Json;
use rocket::tokio::sync::RwLock;
use rocket::tokio::time;
use rocket::{Build, Request, Rocket, fairing, tokio};
use rocket::{Build, Request, Rocket, fairing};
use serde::Serialize;
use sqlx::postgres::PgPool;
use std::sync::Arc;
use std::time::Duration;
#[derive(Debug, Serialize)]
@@ -55,31 +52,7 @@ async fn rocket() -> _ {
.filter_module("querying", LevelFilter::Info)
.init();
let mut interval = time::interval(Duration::from_secs(
std::env::var("DATABASE_INTERVAL_SECONDS")
.unwrap_or("21600".to_string())
.parse()
.unwrap(),
));
let checker = Arc::new(RwLock::new(Checker::new().await));
let checker_clone = checker.clone();
tokio::spawn(async move {
info!("Refreshing DB every {:?}", interval.period());
loop {
interval.tick().await;
info!("Updating all DBs");
match Checker::download_all().await {
Ok(bases) => {
info!("Downloaded, updating...");
checker_clone.read().await.update_all(bases).await;
info!("Updated databases");
}
Err(_) => log::error!("Failed to download all DBs"),
}
}
});
let checker = database_refresh::start().await;
let rate_limit_rpm: u32 = std::env::var("API_RATE_LIMIT_RPM")
.unwrap_or("30".to_string())