diff --git a/.gitignore b/.gitignore index ddbac2f..7cebb46 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,4 @@ /target .idea/ .env +database-cache/ diff --git a/README.md b/README.md index a9651d9..89def46 100644 --- a/README.md +++ b/README.md @@ -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 секунд). + +--- ## Вклад diff --git a/docker-compose.yaml b/docker-compose.yaml index 56b1165..8c53993 100644 --- a/docker-compose.yaml +++ b/docker-compose.yaml @@ -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" diff --git a/querying/src/cache.rs b/querying/src/cache.rs new file mode 100644 index 0000000..e321f29 --- /dev/null +++ b/querying/src/cache.rs @@ -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, + sources: Vec, +} + +#[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 { + 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> { + 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 { + 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) -> io::Result<()> { + let bytes: Vec = contents.iter().copied().collect(); + self.write_file(file_name, &bytes) + } +} diff --git a/querying/src/lib.rs b/querying/src/lib.rs index 4369c73..7c60a99 100644 --- a/querying/src/lib.rs +++ b/querying/src/lib.rs @@ -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; diff --git a/website/src/database_refresh.rs b/website/src/database_refresh.rs new file mode 100644 index 0000000..fee31a1 --- /dev/null +++ b/website/src/database_refresh.rs @@ -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> { + 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>, + 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; + } + }); +} diff --git a/website/src/main.rs b/website/src/main.rs index 75ed31a..d3b6175 100644 --- a/website/src/main.rs +++ b/website/src/main.rs @@ -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())