use std::io::{Cursor, Read as _, Write as _};
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use ed25519_dalek::{Signature, SigningKey, VerifyingKey, PUBLIC_KEY_LENGTH, SIGNATURE_LENGTH};
use sha2::{Digest, Sha256};
use tirith_core::policy;
use tirith_core::threatdb::{ThreatDb, ThreatDbWriter, ThreatSource};
use tirith_core::threatdb_feeds::{
parse_domain_blocklist, parse_phishtank_csv, parse_threatfox_zip, parse_tor_exit_list,
parse_urlhaus_csv,
};
static VERIFY_KEY_BYTES: &[u8; PUBLIC_KEY_LENGTH] =
include_bytes!("../../assets/keys/threatdb-verify.pub");
const MANIFEST_URL_PRIMARY: &str =
"https://raw.githubusercontent.com/sheeki03/tirith/main/threatdb-manifest.json";
const MANIFEST_URL_FALLBACK: &str =
"https://github.com/sheeki03/tirith/releases/latest/download/threatdb-manifest.json";
const MAX_MANIFEST_SIZE: u64 = 64 * 1024;
const MAX_DB_SIZE: u64 = 256 * 1024 * 1024;
const MANIFEST_TIMEOUT_SECS: u64 = 15;
const DB_DOWNLOAD_TIMEOUT_SECS: u64 = 120;
const SUPPLEMENTAL_DOWNLOAD_TIMEOUT_SECS: u64 = 120;
const MAX_SUPPLEMENTAL_FEED_SIZE: u64 = 256 * 1024 * 1024;
const LOCKFILE_NAME: &str = "threatdb-update.lock";
const NEXT_CHECK_FILE: &str = "threatdb-next-check-at";
const SPAWNED_AT_FILE: &str = "threatdb-spawned-at";
const SPAWNED_AT_DEDUP_SECS: u64 = 30;
const BACKOFF_SECS: u64 = 3600;
const URLHAUS_EXPORT_TEMPLATE: &str =
"https://urlhaus-api.abuse.ch/files/exports/full.csv?auth-key={auth_key}";
const THREATFOX_EXPORT_TEMPLATE: &str =
"https://threatfox-api.abuse.ch/files/exports/full.csv.zip?auth-key={auth_key}";
const PHISHING_ARMY_URL: &str =
"https://phishing.army/download/phishing_army_blocklist_extended.txt";
const PHISHTANK_URL: &str = "https://data.phishtank.com/data/online-valid.csv";
const TOR_EXIT_URL: &str = "https://check.torproject.org/torbulkexitlist";
#[derive(Debug, serde::Deserialize)]
struct Manifest {
sha256: String,
size: u64,
url: String,
version: u64,
signature: String,
}
impl Manifest {
fn canonical_payload(&self) -> String {
let mut map = std::collections::BTreeMap::new();
map.insert("sha256", serde_json::Value::String(self.sha256.clone()));
map.insert("size", serde_json::json!(self.size));
map.insert("url", serde_json::Value::String(self.url.clone()));
map.insert("version", serde_json::json!(self.version));
serde_json::to_string(&map).expect("canonical payload serialization")
}
fn verify_signature(&self) -> Result<(), String> {
let sig_bytes =
base64::Engine::decode(&base64::engine::general_purpose::STANDARD, &self.signature)
.map_err(|e| format!("invalid manifest signature encoding: {e}"))?;
if sig_bytes.len() != SIGNATURE_LENGTH {
return Err(format!(
"manifest signature wrong length: {} (expected {})",
sig_bytes.len(),
SIGNATURE_LENGTH
));
}
let signature = Signature::from_slice(&sig_bytes)
.map_err(|e| format!("invalid manifest signature: {e}"))?;
let verify_key = VerifyingKey::from_bytes(VERIFY_KEY_BYTES)
.map_err(|e| format!("invalid embedded public key: {e}"))?;
let payload = self.canonical_payload();
use ed25519_dalek::Verifier;
verify_key
.verify(payload.as_bytes(), &signature)
.map_err(|_| "manifest signature verification failed".to_string())
}
}
pub fn update(force: bool, background: bool) -> i32 {
if background {
return run_background_update();
}
match do_update(force) {
Ok(()) => 0,
Err(e) => {
eprintln!("tirith: threat-db update failed: {e}");
1
}
}
}
fn do_update(force: bool) -> Result<(), String> {
let manifest = fetch_manifest()?;
manifest.verify_signature()?;
if !force {
if let Some(db) = ThreatDb::cached() {
let current_seq = db.build_sequence();
if manifest.version < current_seq {
return Err(format!(
"rollback protection: manifest version {} < current {}",
manifest.version, current_seq
));
}
if manifest.version == current_seq {
eprintln!(
"tirith: threat DB is already up to date (version {})",
manifest.version
);
return Ok(());
}
}
}
eprintln!(
"tirith: downloading threat DB v{} ({} bytes)...",
manifest.version, manifest.size
);
let data = download_db(&manifest)?;
let computed_hash = hex::encode(Sha256::digest(&data));
if computed_hash != manifest.sha256 {
return Err(format!(
"SHA-256 mismatch: expected {}, got {}",
manifest.sha256, computed_hash
));
}
let min_seq = if force { 0 } else { current_sequence() };
let db =
ThreatDb::from_bytes(data.clone(), min_seq).map_err(|e| format!("invalid DB file: {e}"))?;
db.verify_signature()
.map_err(|e| format!("DB file internal signature verification failed: {e}"))?;
let dest =
ThreatDb::default_path().ok_or_else(|| "cannot determine data directory".to_string())?;
atomic_write(&dest, &data)?;
ThreatDb::refresh_cache();
let stats = db.stats();
let total_entries = stats.package_count
+ stats.hostname_count
+ stats.ip_count
+ stats.typosquat_count
+ stats.popular_count;
eprintln!(
"tirith: threat DB updated to v{} ({} entries)",
manifest.version, total_entries
);
if let Err(e) = update_supplemental_db(&policy::Policy::discover(None)) {
eprintln!("tirith: warning: supplemental threat DB update failed: {e}");
}
Ok(())
}
#[derive(Default)]
struct SupplementalEntries {
hostnames: Vec<(String, ThreatSource)>,
ips: Vec<(std::net::Ipv4Addr, ThreatSource)>,
}
impl SupplementalEntries {
fn is_empty(&self) -> bool {
self.hostnames.is_empty() && self.ips.is_empty()
}
fn ingest(
&mut self,
entries: tirith_core::threatdb_feeds::FeedEntries,
source: ThreatSource,
) -> usize {
let count = entries.hostnames.len() + entries.ips.len();
self.hostnames
.extend(entries.hostnames.into_iter().map(|h| (h, source)));
self.ips
.extend(entries.ips.into_iter().map(|ip| (ip, source)));
count
}
}
fn update_supplemental_db(policy: &policy::Policy) -> Result<(), String> {
let supplemental_path = match ThreatDb::supplemental_path() {
Some(path) => path,
None => return Ok(()),
};
let abusech_enabled = policy
.threat_intel
.abusech_auth_key
.as_deref()
.is_some_and(|key| !key.trim().is_empty());
let phishing_enabled = policy.threat_intel.phishing_army_enabled;
if !abusech_enabled && !phishing_enabled {
let _ = std::fs::remove_file(&supplemental_path);
ThreatDb::refresh_cache();
return Ok(());
}
let client = reqwest::blocking::Client::builder()
.timeout(std::time::Duration::from_secs(
SUPPLEMENTAL_DOWNLOAD_TIMEOUT_SECS,
))
.build()
.map_err(|e| format!("supplemental feed HTTP client error: {e}"))?;
let mut supplemental = SupplementalEntries::default();
let mut attempted_feeds = 0usize;
if let Some(auth_key) = policy.threat_intel.abusech_auth_key.as_deref() {
if !auth_key.trim().is_empty() {
attempted_feeds += 1;
log_feed_result(
"URLhaus",
fetch_urlhaus_feed(&client, auth_key.trim(), &mut supplemental),
);
attempted_feeds += 1;
log_feed_result(
"ThreatFox",
fetch_threatfox_feed(&client, auth_key.trim(), &mut supplemental),
);
}
}
if policy.threat_intel.phishing_army_enabled {
attempted_feeds += 1;
log_feed_result(
"Phishing Army",
fetch_phishing_army_feed(&client, &mut supplemental),
);
attempted_feeds += 1;
log_feed_result(
"PhishTank",
fetch_phishtank_feed(&client, &mut supplemental),
);
}
attempted_feeds += 1;
log_feed_result("Tor exit", fetch_tor_exit_feed(&client, &mut supplemental));
if supplemental.is_empty() {
eprintln!(
"tirith: warning: supplemental feeds produced no IOC data across {attempted_feeds} attempted feed(s); leaving existing supplemental threat DB unchanged"
);
return Ok(());
}
let mut writer = ThreatDbWriter::new(unix_now(), 0);
for (host, source) in &supplemental.hostnames {
writer.add_hostname(host, *source);
}
for (ip, source) in &supplemental.ips {
writer.add_ip(*ip, *source);
}
if let Some(parent) = supplemental_path.parent() {
std::fs::create_dir_all(parent)
.map_err(|e| format!("failed to create supplemental DB directory: {e}"))?;
}
let data = writer
.build(&local_overlay_signing_key())
.map_err(|e| format!("failed to build supplemental threat DB: {e}"))?;
atomic_write(&supplemental_path, &data)?;
ThreatDb::refresh_cache();
eprintln!(
"tirith: supplemental threat DB updated ({} hostnames, {} IPs)",
supplemental.hostnames.len(),
supplemental.ips.len()
);
Ok(())
}
fn log_feed_result(feed_name: &str, result: Result<usize, String>) {
match result {
Ok(0) => eprintln!("tirith: warning: {feed_name} feed returned no entries"),
Ok(_) => {}
Err(e) => eprintln!("tirith: warning: {feed_name} feed failed: {e}"),
}
}
fn fetch_urlhaus_feed(
client: &reqwest::blocking::Client,
auth_key: &str,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let url = URLHAUS_EXPORT_TEMPLATE.replace("{auth_key}", auth_key);
let body = fetch_text(client, &url)?;
let entries = parse_urlhaus_csv(Cursor::new(body.into_bytes()))
.map_err(|e| format!("URLhaus parse failed: {e}"))?;
Ok(supplemental.ingest(entries, ThreatSource::Urlhaus))
}
fn fetch_threatfox_feed(
client: &reqwest::blocking::Client,
auth_key: &str,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let url = THREATFOX_EXPORT_TEMPLATE.replace("{auth_key}", auth_key);
let zip_bytes = fetch_bytes(client, &url)?;
let entries = parse_threatfox_zip(Cursor::new(zip_bytes))?;
Ok(supplemental.ingest(entries, ThreatSource::ThreatFoxIoc))
}
fn fetch_phishing_army_feed(
client: &reqwest::blocking::Client,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let body = fetch_text(client, PHISHING_ARMY_URL)?;
let entries = parse_domain_blocklist(&body);
Ok(supplemental.ingest(entries, ThreatSource::PhishingArmy))
}
fn fetch_phishtank_feed(
client: &reqwest::blocking::Client,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let body = fetch_text(client, PHISHTANK_URL)?;
let entries = parse_phishtank_csv(Cursor::new(body.into_bytes()))
.map_err(|e| format!("PhishTank parse failed: {e}"))?;
Ok(supplemental.ingest(entries, ThreatSource::PhishTank))
}
fn fetch_tor_exit_feed(
client: &reqwest::blocking::Client,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let body = fetch_text(client, TOR_EXIT_URL)?;
let entries = parse_tor_exit_list(&body);
Ok(supplemental.ingest(entries, ThreatSource::TorExit))
}
fn redact_url(url: &str) -> String {
if let Some(q) = url.find('?') {
format!("{}?<redacted>", &url[..q])
} else {
url.to_string()
}
}
fn fetch_text(client: &reqwest::blocking::Client, url: &str) -> Result<String, String> {
let bytes = fetch_bytes(client, url)?;
let safe = redact_url(url);
String::from_utf8(bytes)
.map_err(|e| format!("failed to decode UTF-8 response body for {safe}: {e}"))
}
fn fetch_bytes(client: &reqwest::blocking::Client, url: &str) -> Result<Vec<u8>, String> {
let safe = redact_url(url);
let response = client
.get(url)
.header(
"User-Agent",
format!("tirith/{}", env!("CARGO_PKG_VERSION")),
)
.send()
.and_then(|resp| resp.error_for_status())
.map_err(|e| format!("fetch failed for {safe}: {e}"))?;
let content_length = response.content_length();
read_bounded_bytes(response, &safe, content_length, MAX_SUPPLEMENTAL_FEED_SIZE)
}
fn read_bounded_bytes<R: std::io::Read>(
reader: R,
url: &str,
content_length: Option<u64>,
max_size: u64,
) -> Result<Vec<u8>, String> {
if content_length.is_some_and(|len| len > max_size) {
return Err(format!(
"response body for {url} is too large: {content_length:?} bytes exceeds {max_size}"
));
}
let mut limited = reader.take(max_size + 1);
let mut bytes = Vec::new();
limited
.read_to_end(&mut bytes)
.map_err(|e| format!("failed to read response body for {url}: {e}"))?;
if bytes.len() as u64 > max_size {
return Err(format!(
"response body for {url} exceeded max size of {max_size} bytes"
));
}
Ok(bytes)
}
fn local_overlay_signing_key() -> SigningKey {
let digest = Sha256::digest(b"tirith-local-supplemental-threatdb-v1");
let mut key_bytes = [0u8; 32];
key_bytes.copy_from_slice(&digest[..32]);
SigningKey::from_bytes(&key_bytes)
}
fn run_background_update() -> i32 {
let state = match policy::state_dir() {
Some(d) => d,
None => return 1,
};
if let Err(e) = std::fs::create_dir_all(&state) {
eprintln!(
"tirith: warning: failed to create state directory {}: {e}",
state.display()
);
return 1;
}
let lock_path = state.join(LOCKFILE_NAME);
let lock_file = match std::fs::OpenOptions::new()
.create(true)
.truncate(false)
.write(true)
.open(&lock_path)
{
Ok(f) => f,
Err(e) => {
eprintln!(
"tirith: warning: failed to open lock file {}: {e}",
lock_path.display()
);
return 1;
}
};
use fs2::FileExt;
if lock_file.try_lock_exclusive().is_err() {
return 0;
}
let policy = policy::Policy::discover(None);
let auto_hours = policy.threat_intel.auto_update_hours;
if auto_hours == 0 {
let _ = fs2::FileExt::unlock(&lock_file);
return 0;
}
let result = do_update(false);
let next_check_path = state.join(NEXT_CHECK_FILE);
let now = unix_now();
let success = result.is_ok();
if success {
let next = now + auto_hours * 3600;
if let Err(e) = std::fs::write(&next_check_path, next.to_string()) {
eprintln!("tirith: warning: failed to write next-check-at: {e}");
}
} else {
if let Err(ref e) = result {
eprintln!("tirith: background update failed: {e}");
}
let next = now + BACKOFF_SECS;
if let Err(e) = std::fs::write(&next_check_path, next.to_string()) {
eprintln!("tirith: warning: failed to write next-check-at: {e}");
}
}
let _ = fs2::FileExt::unlock(&lock_file);
if success {
0
} else {
1
}
}
pub fn status(json: bool) -> i32 {
let info = gather_status();
if json {
match serde_json::to_string_pretty(&info) {
Ok(s) => println!("{s}"),
Err(e) => {
eprintln!("tirith: JSON serialization failed: {e}");
return 1;
}
}
} else {
print_status_human(&info);
}
0
}
#[derive(Debug, serde::Serialize)]
struct ThreatDbStatus {
installed: bool,
path: Option<String>,
age_hours: Option<f64>,
build_timestamp: Option<u64>,
build_sequence: Option<u64>,
package_count: Option<u32>,
hostname_count: Option<u32>,
ip_count: Option<u32>,
typosquat_count: Option<u32>,
popular_count: Option<u32>,
total_entries: Option<u32>,
skipped_range_only: Option<u32>,
signature_valid: Option<bool>,
stale: bool,
error: Option<String>,
}
fn gather_status() -> ThreatDbStatus {
let db_path = ThreatDb::default_path();
let path_str = db_path.as_ref().map(|p| p.display().to_string());
let db_path_ref = match db_path {
Some(ref p) if p.exists() => p,
_ => {
return ThreatDbStatus {
installed: false,
path: path_str,
age_hours: None,
build_timestamp: None,
build_sequence: None,
package_count: None,
hostname_count: None,
ip_count: None,
typosquat_count: None,
popular_count: None,
total_entries: None,
skipped_range_only: None,
signature_valid: None,
stale: true,
error: None,
};
}
};
match ThreatDb::load_from_path(db_path_ref, 0) {
Ok(db) => {
let sig_valid = db.verify_signature().is_ok();
let stats = db.stats();
let now = unix_now();
let age_secs = now.saturating_sub(stats.build_timestamp);
let age_hours = age_secs as f64 / 3600.0;
let total = stats.package_count
+ stats.hostname_count
+ stats.ip_count
+ stats.typosquat_count
+ stats.popular_count;
let policy = policy::Policy::discover(None);
let stale_hours = policy.threat_intel.auto_update_hours;
let is_stale = if stale_hours == 0 {
false
} else {
age_hours > (stale_hours as f64 * 2.0)
};
ThreatDbStatus {
installed: true,
path: path_str,
age_hours: Some(age_hours),
build_timestamp: Some(stats.build_timestamp),
build_sequence: Some(stats.build_sequence),
package_count: Some(stats.package_count),
hostname_count: Some(stats.hostname_count),
ip_count: Some(stats.ip_count),
typosquat_count: Some(stats.typosquat_count),
popular_count: Some(stats.popular_count),
total_entries: Some(total),
skipped_range_only: None,
signature_valid: Some(sig_valid),
stale: is_stale,
error: None,
}
}
Err(e) => ThreatDbStatus {
installed: true,
path: path_str,
age_hours: None,
build_timestamp: None,
build_sequence: None,
package_count: None,
hostname_count: None,
ip_count: None,
typosquat_count: None,
popular_count: None,
total_entries: None,
skipped_range_only: None,
signature_valid: None,
stale: true,
error: Some(format!("{e}")),
},
}
}
fn print_status_human(info: &ThreatDbStatus) {
if !info.installed {
println!("threat DB: not installed — run 'tirith threat-db update'");
if let Some(ref path) = info.path {
println!(" expected at: {path}");
}
return;
}
if let Some(ref err) = info.error {
println!("threat DB: ERROR: {err}");
if let Some(ref path) = info.path {
println!(" path: {path}");
}
println!(" Hint: re-download with 'tirith threat-db update --force'");
return;
}
if info.signature_valid == Some(false) {
println!(
"threat DB: INVALID SIGNATURE — re-download with 'tirith threat-db update --force'"
);
if let Some(ref path) = info.path {
println!(" path: {path}");
}
return;
}
let path = info.path.as_deref().unwrap_or("unknown");
let age_str = match info.age_hours {
Some(h) if h < 1.0 => format!("{:.0}m old", h * 60.0),
Some(h) if h < 48.0 => format!("{:.0}h old", h),
Some(h) => format!("{:.0}d old", h / 24.0),
None => "unknown age".to_string(),
};
let total = info.total_entries.unwrap_or(0);
if info.stale {
println!("threat DB: STALE ({age_str}) — run 'tirith threat-db update'");
} else {
let sig_label = if info.signature_valid == Some(true) {
"signature ok"
} else {
"signature unknown"
};
println!("threat DB: {path} ({age_str}, {total} entries, {sig_label})");
}
if let Some(seq) = info.build_sequence {
println!(" version: {seq}");
}
if let (Some(pkg), Some(host), Some(ip), Some(typo), Some(pop)) = (
info.package_count,
info.hostname_count,
info.ip_count,
info.typosquat_count,
info.popular_count,
) {
println!(
" entries: {pkg} packages, {host} hostnames, {ip} IPs, {typo} typosquats, {pop} popular"
);
}
println!(
" update: auto-update checks main manifest, falls back to release asset if stale"
);
println!(" (fallback may hit GitHub API rate limits for unauthenticated users)");
}
static UPDATE_ATTEMPTED: AtomicBool = AtomicBool::new(false);
pub fn maybe_background_update() {
if UPDATE_ATTEMPTED.swap(true, Ordering::Relaxed) {
return;
}
let policy = policy::Policy::discover(None);
if policy.threat_intel.auto_update_hours == 0 {
return;
}
let state = match policy::state_dir() {
Some(d) => d,
None => return,
};
let next_check_path = state.join(NEXT_CHECK_FILE);
let now = unix_now();
if let Ok(content) = std::fs::read_to_string(&next_check_path) {
if let Ok(next_ts) = content.trim().parse::<u64>() {
if now < next_ts {
return;
}
}
}
let spawned_at_path = state.join(SPAWNED_AT_FILE);
if let Ok(content) = std::fs::read_to_string(&spawned_at_path) {
if let Ok(spawned_ts) = content.trim().parse::<u64>() {
if now.saturating_sub(spawned_ts) < SPAWNED_AT_DEDUP_SECS {
return;
}
}
}
if let Err(e) = std::fs::create_dir_all(&state) {
eprintln!("tirith: warning: failed to create state directory: {e}");
return;
}
let _ = std::fs::write(&spawned_at_path, now.to_string());
let exe = match std::env::current_exe() {
Ok(e) => e,
Err(_) => return,
};
match std::process::Command::new(&exe)
.args(["threat-db", "update", "--background"])
.stdin(std::process::Stdio::null())
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.spawn()
{
Ok(_) => {}
Err(e) => {
eprintln!("tirith: warning: failed to spawn background update: {e}");
let _ = std::fs::remove_file(&spawned_at_path);
}
}
}
fn fetch_manifest() -> Result<Manifest, String> {
match fetch_manifest_from(MANIFEST_URL_PRIMARY) {
Ok(m) => {
if let Some(db) = ThreatDb::cached() {
if m.version <= db.build_sequence() {
eprintln!("tirith: primary manifest is stale (v{} <= current v{}), trying fallback...",
m.version, db.build_sequence());
match fetch_manifest_from(MANIFEST_URL_FALLBACK) {
Ok(fallback) if fallback.version > db.build_sequence() => {
return Ok(fallback)
}
_ => {}
}
}
}
Ok(m)
}
Err(primary_err) => {
eprintln!("tirith: primary manifest unavailable ({primary_err}), trying fallback...");
fetch_manifest_from(MANIFEST_URL_FALLBACK).map_err(|fallback_err| {
format!("manifest fetch failed: primary: {primary_err}; fallback: {fallback_err}")
})
}
}
}
#[derive(Debug, PartialEq)]
enum CacheResolution {
Fresh(String),
Cached(String),
RetryNeeded,
}
fn resolve_cache(
http_status: u16,
response_body: Option<&str>,
cached_body: Option<&str>,
) -> Result<CacheResolution, String> {
if http_status == 304 {
if let Some(body) = cached_body {
if serde_json::from_str::<Manifest>(body).is_ok() {
return Ok(CacheResolution::Cached(body.to_string()));
}
}
return Ok(CacheResolution::RetryNeeded);
}
if !(200..300).contains(&http_status) {
return Err(format!("HTTP {http_status}"));
}
match response_body {
Some(body) => Ok(CacheResolution::Fresh(body.to_string())),
None => Err("empty response body".to_string()),
}
}
fn manifest_cache_key(url: &str) -> String {
use sha2::{Digest, Sha256};
let hash = Sha256::digest(url.as_bytes());
let hex: String = hash.iter().take(8).map(|b| format!("{b:02x}")).collect();
format!("threatdb-manifest-{hex}")
}
fn fetch_manifest_from(url: &str) -> Result<Manifest, String> {
fetch_manifest_from_with_state(url, tirith_core::policy::state_dir())
}
fn fetch_manifest_from_with_state(
url: &str,
state: Option<std::path::PathBuf>,
) -> Result<Manifest, String> {
let client = reqwest::blocking::Client::builder()
.timeout(std::time::Duration::from_secs(MANIFEST_TIMEOUT_SECS))
.build()
.map_err(|e| format!("HTTP client error: {e}"))?;
let cache_key = manifest_cache_key(url);
let etag_path = state.as_ref().map(|d| d.join(format!("{cache_key}-etag")));
let body_path = state.as_ref().map(|d| d.join(format!("{cache_key}-body")));
let mut req = client.get(url).header(
"User-Agent",
format!("tirith/{}", env!("CARGO_PKG_VERSION")),
);
if let Some(ref ep) = etag_path {
if let Ok(etag) = std::fs::read_to_string(ep) {
let etag = etag.trim();
if !etag.is_empty() {
req = req.header("If-None-Match", etag);
}
}
}
let resp = req
.send()
.map_err(|e| format!("manifest fetch failed: {e}"))?;
let status = resp.status().as_u16();
let resp_etag = resp
.headers()
.get("etag")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let resp_body = if status != 304 {
let content_len = resp.content_length().unwrap_or(0);
if content_len > MAX_MANIFEST_SIZE {
return Err(format!(
"manifest too large: {} bytes (max {})",
content_len, MAX_MANIFEST_SIZE
));
}
let body = resp
.text()
.map_err(|e| format!("failed to read manifest body: {e}"))?;
if body.len() as u64 > MAX_MANIFEST_SIZE {
return Err(format!("manifest body too large: {} bytes", body.len()));
}
Some(body)
} else {
None
};
let cached_body = if status == 304 {
body_path.as_ref().and_then(|bp| {
if let Ok(meta) = std::fs::metadata(bp) {
if meta.len() > MAX_MANIFEST_SIZE {
eprintln!(
"tirith: warning: cached manifest too large ({} bytes), ignoring",
meta.len()
);
return None;
}
}
let content = std::fs::read_to_string(bp).ok()?;
Some(content)
})
} else {
None
};
match resolve_cache(status, resp_body.as_deref(), cached_body.as_deref()) {
Ok(CacheResolution::Fresh(body)) => {
let manifest = serde_json::from_str::<Manifest>(&body)
.map_err(|e| format!("invalid manifest JSON: {e}"))?;
persist_cache_files(&etag_path, resp_etag.as_deref(), &body_path, &body);
Ok(manifest)
}
Ok(CacheResolution::Cached(body)) => serde_json::from_str::<Manifest>(&body)
.map_err(|e| format!("cached manifest parse error: {e}")),
Ok(CacheResolution::RetryNeeded) => {
if let Some(ref ep) = etag_path {
let _ = std::fs::remove_file(ep);
}
if let Some(ref bp) = body_path {
let _ = std::fs::remove_file(bp);
}
let retry_resp = client
.get(url)
.header(
"User-Agent",
format!("tirith/{}", env!("CARGO_PKG_VERSION")),
)
.send()
.map_err(|e| format!("manifest retry fetch failed: {e}"))?;
if !retry_resp.status().is_success() {
return Err(format!("manifest retry HTTP {}", retry_resp.status()));
}
let retry_content_len = retry_resp.content_length().unwrap_or(0);
if retry_content_len > MAX_MANIFEST_SIZE {
return Err(format!(
"manifest too large on retry: {} bytes (max {})",
retry_content_len, MAX_MANIFEST_SIZE
));
}
let retry_etag = retry_resp
.headers()
.get("etag")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let retry_body = retry_resp
.text()
.map_err(|e| format!("failed to read retry body: {e}"))?;
if retry_body.len() as u64 > MAX_MANIFEST_SIZE {
return Err(format!(
"manifest body too large on retry: {} bytes",
retry_body.len()
));
}
let manifest = serde_json::from_str::<Manifest>(&retry_body)
.map_err(|e| format!("invalid manifest JSON on retry: {e}"))?;
persist_cache_files(&etag_path, retry_etag.as_deref(), &body_path, &retry_body);
Ok(manifest)
}
Err(e) => Err(e),
}
}
fn persist_cache_files(
etag_path: &Option<std::path::PathBuf>,
etag_val: Option<&str>,
body_path: &Option<std::path::PathBuf>,
body: &str,
) {
if let (Some(ep), Some(val)) = (etag_path, etag_val) {
if let Some(parent) = ep.parent() {
let _ = std::fs::create_dir_all(parent);
}
let _ = std::fs::write(ep, val);
}
if let Some(bp) = body_path {
if let Some(parent) = bp.parent() {
let _ = std::fs::create_dir_all(parent);
}
let _ = std::fs::write(bp, body);
}
}
fn download_db(manifest: &Manifest) -> Result<Vec<u8>, String> {
if manifest.size > MAX_DB_SIZE {
return Err(format!(
"DB file too large: {} bytes (max {})",
manifest.size, MAX_DB_SIZE
));
}
let client = reqwest::blocking::Client::builder()
.timeout(std::time::Duration::from_secs(DB_DOWNLOAD_TIMEOUT_SECS))
.build()
.map_err(|e| format!("HTTP client error: {e}"))?;
let resp = client
.get(&manifest.url)
.header(
"User-Agent",
format!("tirith/{}", env!("CARGO_PKG_VERSION")),
)
.send()
.map_err(|e| format!("DB download failed: {e}"))?;
if !resp.status().is_success() {
return Err(format!("DB download HTTP {}", resp.status()));
}
let bytes = resp
.bytes()
.map_err(|e| format!("failed to read DB body: {e}"))?;
if bytes.len() as u64 > MAX_DB_SIZE {
return Err(format!("DB body too large: {} bytes", bytes.len()));
}
Ok(bytes.to_vec())
}
fn atomic_write(dest: &PathBuf, data: &[u8]) -> Result<(), String> {
let parent = dest
.parent()
.ok_or_else(|| "cannot determine parent directory".to_string())?;
std::fs::create_dir_all(parent).map_err(|e| format!("failed to create directory: {e}"))?;
let mut tmp = tempfile::NamedTempFile::new_in(parent)
.map_err(|e| format!("failed to create temp file: {e}"))?;
tmp.write_all(data)
.map_err(|e| format!("failed to write temp file: {e}"))?;
tmp.flush()
.map_err(|e| format!("failed to flush temp file: {e}"))?;
tmp.persist(dest)
.map_err(|e| format!("failed to rename temp file: {e}"))?;
Ok(())
}
fn unix_now() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
fn current_sequence() -> u64 {
ThreatDb::cached()
.map(|db| db.build_sequence())
.unwrap_or(0)
}
mod hex {
pub fn encode(data: impl AsRef<[u8]>) -> String {
data.as_ref().iter().map(|b| format!("{b:02x}")).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::Path;
use std::sync::atomic::Ordering;
static TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn is_next_check_in_future(state_dir: &Path, now: u64) -> bool {
let next_check_path = state_dir.join(NEXT_CHECK_FILE);
if let Ok(content) = std::fs::read_to_string(&next_check_path) {
if let Ok(next_ts) = content.trim().parse::<u64>() {
return now < next_ts;
}
}
false
}
fn is_spawned_at_recent(state_dir: &Path, now: u64) -> bool {
let spawned_at_path = state_dir.join(SPAWNED_AT_FILE);
if let Ok(content) = std::fs::read_to_string(&spawned_at_path) {
if let Ok(spawned_ts) = content.trim().parse::<u64>() {
return now.saturating_sub(spawned_ts) < SPAWNED_AT_DEDUP_SECS;
}
}
false
}
fn try_acquire_update_lock(state_dir: &Path) -> Option<std::fs::File> {
let lock_path = state_dir.join(LOCKFILE_NAME);
let lock_file = std::fs::OpenOptions::new()
.create(true)
.truncate(false)
.write(true)
.open(&lock_path)
.ok()?;
use fs2::FileExt;
if lock_file.try_lock_exclusive().is_err() {
return None;
}
Some(lock_file)
}
#[test]
fn auto_update_hours_zero_disables_background_child() {
let _guard = TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let policy_dir = tmp.path().join(".tirith");
std::fs::create_dir_all(&policy_dir).unwrap();
std::fs::write(
policy_dir.join("policy.yaml"),
"threat_intel:\n auto_update_hours: 0\n",
)
.unwrap();
unsafe { std::env::set_var("TIRITH_POLICY_ROOT", tmp.path()) };
let policy = policy::Policy::discover(Some(tmp.path().to_str().unwrap()));
assert_eq!(
policy.threat_intel.auto_update_hours, 0,
"policy should reflect auto_update_hours=0"
);
unsafe { std::env::remove_var("TIRITH_POLICY_ROOT") };
}
#[test]
fn next_check_at_future_skips_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let future_ts = unix_now() + 3600;
std::fs::write(state.join(NEXT_CHECK_FILE), future_ts.to_string()).unwrap();
let now = unix_now();
assert!(
is_next_check_in_future(state, now),
"should skip when next-check-at is in the future"
);
}
#[test]
fn next_check_at_past_allows_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let past_ts = unix_now().saturating_sub(3600);
std::fs::write(state.join(NEXT_CHECK_FILE), past_ts.to_string()).unwrap();
let now = unix_now();
assert!(
!is_next_check_in_future(state, now),
"should proceed when next-check-at is in the past"
);
}
#[test]
fn next_check_at_missing_allows_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let now = unix_now();
assert!(
!is_next_check_in_future(state, now),
"should proceed when next-check-at file does not exist"
);
}
#[test]
fn next_check_at_corrupt_allows_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
std::fs::write(state.join(NEXT_CHECK_FILE), "not-a-number").unwrap();
let now = unix_now();
assert!(
!is_next_check_in_future(state, now),
"should proceed when next-check-at is unparseable"
);
}
#[test]
fn spawned_at_recent_skips_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let recent_ts = unix_now().saturating_sub(5);
std::fs::write(state.join(SPAWNED_AT_FILE), recent_ts.to_string()).unwrap();
let now = unix_now();
assert!(
is_spawned_at_recent(state, now),
"should skip when spawned-at is recent (within 30s window)"
);
}
#[test]
fn spawned_at_old_allows_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let old_ts = unix_now().saturating_sub(60);
std::fs::write(state.join(SPAWNED_AT_FILE), old_ts.to_string()).unwrap();
let now = unix_now();
assert!(
!is_spawned_at_recent(state, now),
"should proceed when spawned-at is older than 30s"
);
}
#[test]
fn spawned_at_missing_allows_update() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let now = unix_now();
assert!(
!is_spawned_at_recent(state, now),
"should proceed when spawned-at file does not exist"
);
}
#[test]
fn update_attempted_guard_fires_once() {
let guard = AtomicBool::new(false);
let first = guard.swap(true, Ordering::Relaxed);
assert!(
!first,
"first swap should return false, allowing the update"
);
let second = guard.swap(true, Ordering::Relaxed);
assert!(second, "second swap should return true, blocking re-entry");
let third = guard.swap(true, Ordering::Relaxed);
assert!(third, "third swap should also return true");
}
#[test]
fn lock_dedup_second_acquire_fails() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let lock1 = try_acquire_update_lock(state);
assert!(lock1.is_some(), "first lock acquisition should succeed");
let lock2 = try_acquire_update_lock(state);
assert!(
lock2.is_none(),
"second lock acquisition should fail while first is held"
);
let l1 = lock1.unwrap();
fs2::FileExt::unlock(&l1).expect("unlock lock1");
drop(l1);
let lock3 = try_acquire_update_lock(state);
assert!(
lock3.is_some(),
"lock acquisition should succeed after previous lock is released"
);
}
#[test]
fn lock_file_is_created_in_state_dir() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let lock = try_acquire_update_lock(state);
assert!(lock.is_some());
assert!(
state.join(LOCKFILE_NAME).exists(),
"lock file should be created at the expected path"
);
}
#[test]
fn failure_backoff_sets_one_hour() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let next_check_path = state.join(NEXT_CHECK_FILE);
let now = unix_now();
let backoff_ts = now + BACKOFF_SECS;
std::fs::write(&next_check_path, backoff_ts.to_string()).unwrap();
let content = std::fs::read_to_string(&next_check_path).unwrap();
let written_ts: u64 = content.trim().parse().unwrap();
let diff = written_ts.saturating_sub(now);
assert_eq!(
diff, BACKOFF_SECS,
"backoff should set next-check-at to now + {} seconds, got diff={}",
BACKOFF_SECS, diff
);
assert_eq!(
BACKOFF_SECS, 3600,
"BACKOFF_SECS constant should be 3600 (1 hour)"
);
}
#[test]
fn success_sets_next_check_at_auto_update_hours() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let next_check_path = state.join(NEXT_CHECK_FILE);
let auto_hours: u64 = 24;
let now = unix_now();
let next = now + auto_hours * 3600;
std::fs::write(&next_check_path, next.to_string()).unwrap();
let content = std::fs::read_to_string(&next_check_path).unwrap();
let written_ts: u64 = content.trim().parse().unwrap();
let diff = written_ts.saturating_sub(now);
assert_eq!(
diff,
auto_hours * 3600,
"success should set next-check-at to now + auto_update_hours*3600"
);
}
#[test]
fn backoff_differs_from_normal_interval() {
let default_config = policy::ThreatIntelConfig::default();
let normal_interval_secs = default_config.auto_update_hours * 3600;
assert_ne!(
BACKOFF_SECS, normal_interval_secs,
"backoff interval ({BACKOFF_SECS}s) must differ from normal interval ({normal_interval_secs}s)"
);
assert!(
BACKOFF_SECS < normal_interval_secs,
"backoff ({BACKOFF_SECS}s) should be shorter than normal interval ({normal_interval_secs}s) for faster retry"
);
}
#[test]
fn canonical_payload_format_sorted_keys_no_whitespace() {
let manifest = Manifest {
sha256: "abcdef1234567890abcdef1234567890abcdef1234567890abcdef1234567890".to_string(),
size: 12345,
url: "https://example.com/tirith-threatdb.dat".to_string(),
version: 42,
signature: String::new(),
};
let payload = manifest.canonical_payload();
assert_eq!(
payload,
r#"{"sha256":"abcdef1234567890abcdef1234567890abcdef1234567890abcdef1234567890","size":12345,"url":"https://example.com/tirith-threatdb.dat","version":42}"#,
"canonical payload should have alphabetically sorted keys with no whitespace"
);
}
#[test]
fn canonical_payload_no_whitespace() {
let manifest = Manifest {
sha256: "deadbeef".to_string(),
size: 1,
url: "https://x.com/db.dat".to_string(),
version: 1,
signature: String::new(),
};
let payload = manifest.canonical_payload();
assert!(
!payload.contains(' '),
"canonical payload must not contain spaces"
);
assert!(
!payload.contains('\n'),
"canonical payload must not contain newlines"
);
assert!(
!payload.contains('\t'),
"canonical payload must not contain tabs"
);
assert!(
!payload.ends_with('\n'),
"canonical payload must not have trailing newline"
);
}
#[test]
fn canonical_payload_is_valid_utf8_json() {
let manifest = Manifest {
sha256: "0123456789abcdef".to_string(),
size: 999,
url: "https://example.com/db.dat".to_string(),
version: 7,
signature: String::new(),
};
let payload = manifest.canonical_payload();
assert!(
std::str::from_utf8(payload.as_bytes()).is_ok(),
"canonical payload must be valid UTF-8"
);
let parsed: serde_json::Value =
serde_json::from_str(&payload).expect("canonical payload must be valid JSON");
let obj = parsed.as_object().expect("payload should be a JSON object");
let keys: Vec<&String> = obj.keys().collect();
assert_eq!(
keys,
&["sha256", "size", "url", "version"],
"keys must be in alphabetical order"
);
}
#[test]
fn canonical_payload_excludes_signature_field() {
let manifest = Manifest {
sha256: "abc".to_string(),
size: 1,
url: "https://x.com/db.dat".to_string(),
version: 1,
signature: "should-not-appear-in-payload".to_string(),
};
let payload = manifest.canonical_payload();
assert!(
!payload.contains("signature"),
"canonical payload must not include the 'signature' field"
);
assert!(
!payload.contains("should-not-appear-in-payload"),
"canonical payload must not include the signature value"
);
}
#[test]
fn canonical_payload_round_trips_through_json_parse() {
let manifest = Manifest {
sha256: "abc123".to_string(),
size: 42,
url: "https://example.com/db.dat".to_string(),
version: 99,
signature: "ignored".to_string(),
};
let payload = manifest.canonical_payload();
let parsed: serde_json::Value = serde_json::from_str(&payload).unwrap();
assert_eq!(parsed["sha256"], "abc123");
assert_eq!(parsed["size"], 42);
assert_eq!(parsed["url"], "https://example.com/db.dat");
assert_eq!(parsed["version"], 99);
}
#[test]
fn spawned_at_exactly_at_boundary_skips() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let now = 1000000u64;
let ts = now - (SPAWNED_AT_DEDUP_SECS - 1);
std::fs::write(state.join(SPAWNED_AT_FILE), ts.to_string()).unwrap();
assert!(
is_spawned_at_recent(state, now),
"29 seconds ago should still be within the dedup window"
);
}
#[test]
fn spawned_at_exactly_at_boundary_allows() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let now = 1000000u64;
let ts = now - SPAWNED_AT_DEDUP_SECS;
std::fs::write(state.join(SPAWNED_AT_FILE), ts.to_string()).unwrap();
assert!(
!is_spawned_at_recent(state, now),
"exactly 30 seconds ago should be outside the dedup window"
);
}
#[test]
fn next_check_at_exactly_now_allows() {
let tmp = tempfile::tempdir().unwrap();
let state = tmp.path();
let now = 1000000u64;
std::fs::write(state.join(NEXT_CHECK_FILE), now.to_string()).unwrap();
assert!(
!is_next_check_in_future(state, now),
"next-check-at == now should allow the update (not strictly in the future)"
);
}
#[test]
fn manifest_cache_key_is_url_specific() {
let k1 = super::manifest_cache_key("https://example.com/manifest.json");
let k2 = super::manifest_cache_key("https://other.com/manifest.json");
assert_ne!(k1, k2, "different URLs must produce different cache keys");
assert!(
k1.starts_with("threatdb-manifest-"),
"cache key should have expected prefix"
);
}
#[test]
fn manifest_cache_key_is_deterministic() {
let url = "https://example.com/manifest.json";
assert_eq!(
super::manifest_cache_key(url),
super::manifest_cache_key(url),
"same URL must produce same cache key"
);
}
#[test]
fn cached_body_round_trips_through_json() {
let json = r#"{"sha256":"abc123","size":42,"url":"https://example.com/db.dat","version":99,"signature":"sig"}"#;
let parsed: Manifest = serde_json::from_str(json).unwrap();
assert_eq!(parsed.sha256, "abc123");
assert_eq!(parsed.version, 99);
assert_eq!(parsed.size, 42);
let tmp = tempfile::tempdir().unwrap();
let body_file = tmp.path().join("cached-body");
std::fs::write(&body_file, json).unwrap();
let reloaded = std::fs::read_to_string(&body_file).unwrap();
let reparsed: Manifest = serde_json::from_str(&reloaded).unwrap();
assert_eq!(reparsed.sha256, "abc123");
assert_eq!(reparsed.version, 99);
}
#[test]
fn etag_and_body_files_are_per_url() {
let url1 = "https://primary.example.com/m.json";
let url2 = "https://fallback.example.com/m.json";
let k1 = super::manifest_cache_key(url1);
let k2 = super::manifest_cache_key(url2);
let etag1 = format!("{k1}-etag");
let etag2 = format!("{k2}-etag");
assert_ne!(etag1, etag2, "etag files must be per-URL");
let body1 = format!("{k1}-body");
let body2 = format!("{k2}-body");
assert_ne!(body1, body2, "body cache files must be per-URL");
}
#[test]
fn cache_200_then_304_round_trip() {
let tmp = tempfile::tempdir().unwrap();
let url = "https://example.com/manifest.json";
let key = super::manifest_cache_key(url);
let etag_path = tmp.path().join(format!("{key}-etag"));
let body_path = tmp.path().join(format!("{key}-body"));
let manifest_json = r#"{"sha256":"dead","size":100,"url":"https://x.com/db.dat","version":5,"signature":"sig"}"#;
std::fs::write(&etag_path, "\"etag-value-abc\"").unwrap();
std::fs::write(&body_path, manifest_json).unwrap();
let cached = std::fs::read_to_string(&body_path).unwrap();
let m: Manifest = serde_json::from_str(&cached).unwrap();
assert_eq!(m.sha256, "dead");
assert_eq!(m.version, 5);
let etag = std::fs::read_to_string(&etag_path).unwrap();
assert_eq!(etag.trim(), "\"etag-value-abc\"");
}
#[test]
fn cache_304_with_missing_body_cleans_etag() {
let tmp = tempfile::tempdir().unwrap();
let url = "https://example.com/manifest.json";
let key = super::manifest_cache_key(url);
let etag_path = tmp.path().join(format!("{key}-etag"));
let body_path = tmp.path().join(format!("{key}-body"));
std::fs::write(&etag_path, "\"stale-etag\"").unwrap();
assert!(!body_path.exists(), "body should not exist for this test");
let body_ok = body_path
.exists()
.then(|| std::fs::read_to_string(&body_path).ok())
.flatten()
.and_then(|s| serde_json::from_str::<Manifest>(&s).ok());
if body_ok.is_none() {
let _ = std::fs::remove_file(&etag_path);
let _ = std::fs::remove_file(&body_path);
}
assert!(
!etag_path.exists(),
"ETag should be deleted after 304 with missing body"
);
}
#[test]
fn cache_304_with_corrupt_body_cleans_etag() {
let tmp = tempfile::tempdir().unwrap();
let url = "https://example.com/manifest.json";
let key = super::manifest_cache_key(url);
let etag_path = tmp.path().join(format!("{key}-etag"));
let body_path = tmp.path().join(format!("{key}-body"));
std::fs::write(&etag_path, "\"some-etag\"").unwrap();
std::fs::write(&body_path, "this is not json").unwrap();
let body_ok = std::fs::read_to_string(&body_path)
.ok()
.and_then(|s| serde_json::from_str::<Manifest>(&s).ok());
if body_ok.is_none() {
let _ = std::fs::remove_file(&etag_path);
let _ = std::fs::remove_file(&body_path);
}
assert!(
!etag_path.exists(),
"ETag should be deleted after 304 with corrupt body"
);
assert!(
!body_path.exists(),
"Corrupt body should be deleted after recovery"
);
}
#[test]
fn primary_and_fallback_independent_cache_state() {
let tmp = tempfile::tempdir().unwrap();
let primary =
"https://raw.githubusercontent.com/sheeki03/tirith/main/threatdb-manifest.json";
let fallback =
"https://github.com/sheeki03/tirith/releases/latest/download/threatdb-manifest.json";
let pk = super::manifest_cache_key(primary);
let fk = super::manifest_cache_key(fallback);
let p_etag = tmp.path().join(format!("{pk}-etag"));
let f_etag = tmp.path().join(format!("{fk}-etag"));
let p_body = tmp.path().join(format!("{pk}-body"));
let f_body = tmp.path().join(format!("{fk}-body"));
std::fs::write(&p_etag, "\"primary-etag\"").unwrap();
std::fs::write(
&p_body,
r#"{"sha256":"p","size":1,"url":"p","version":10,"signature":"s"}"#,
)
.unwrap();
std::fs::write(&f_etag, "\"fallback-etag\"").unwrap();
std::fs::write(
&f_body,
r#"{"sha256":"f","size":2,"url":"f","version":20,"signature":"s"}"#,
)
.unwrap();
let pm: Manifest =
serde_json::from_str(&std::fs::read_to_string(&p_body).unwrap()).unwrap();
let fm: Manifest =
serde_json::from_str(&std::fs::read_to_string(&f_body).unwrap()).unwrap();
assert_eq!(pm.version, 10);
assert_eq!(fm.version, 20);
assert_ne!(
std::fs::read_to_string(&p_etag).unwrap(),
std::fs::read_to_string(&f_etag).unwrap()
);
std::fs::remove_file(&p_etag).unwrap();
std::fs::remove_file(&p_body).unwrap();
assert!(
f_etag.exists(),
"fallback ETag should survive primary cleanup"
);
assert!(
f_body.exists(),
"fallback body should survive primary cleanup"
);
}
const VALID_MANIFEST: &str =
r#"{"sha256":"abc","size":1,"url":"https://x.com/db.dat","version":1,"signature":"s"}"#;
#[test]
fn resolve_cache_200_returns_fresh() {
let r = super::resolve_cache(200, Some(VALID_MANIFEST), None).unwrap();
assert_eq!(r, super::CacheResolution::Fresh(VALID_MANIFEST.to_string()));
}
#[test]
fn resolve_cache_200_ignores_cached_body() {
let r = super::resolve_cache(200, Some(VALID_MANIFEST), Some("old")).unwrap();
match r {
super::CacheResolution::Fresh(body) => assert_eq!(body, VALID_MANIFEST),
other => panic!("expected Fresh, got {other:?}"),
}
}
#[test]
fn resolve_cache_304_with_valid_cache_returns_cached() {
let r = super::resolve_cache(304, None, Some(VALID_MANIFEST)).unwrap();
assert_eq!(
r,
super::CacheResolution::Cached(VALID_MANIFEST.to_string())
);
}
#[test]
fn resolve_cache_304_with_no_cache_returns_retry() {
let r = super::resolve_cache(304, None, None).unwrap();
assert_eq!(r, super::CacheResolution::RetryNeeded);
}
#[test]
fn resolve_cache_304_with_corrupt_cache_returns_retry() {
let r = super::resolve_cache(304, None, Some("not json")).unwrap();
assert_eq!(
r,
super::CacheResolution::RetryNeeded,
"corrupt cache should trigger retry, not error"
);
}
#[test]
fn resolve_cache_404_returns_error() {
let r = super::resolve_cache(404, None, None);
assert!(r.is_err());
assert!(r.unwrap_err().contains("404"));
}
#[test]
fn resolve_cache_500_returns_error() {
let r = super::resolve_cache(500, None, None);
assert!(r.is_err());
}
#[test]
fn resolve_cache_200_with_no_body_returns_error() {
let r = super::resolve_cache(200, None, None);
assert!(r.is_err());
assert!(r.unwrap_err().contains("empty"));
}
#[test]
fn resolve_cache_201_accepted_as_success() {
let r = super::resolve_cache(201, Some(VALID_MANIFEST), None).unwrap();
assert_eq!(r, super::CacheResolution::Fresh(VALID_MANIFEST.to_string()));
}
fn fetch_with_state(url: &str, state: &std::path::Path) -> Result<Manifest, String> {
super::fetch_manifest_from_with_state(url, Some(state.to_path_buf()))
}
#[test]
fn transport_200_returns_manifest_and_caches_body() {
let mut server = mockito::Server::new();
let manifest_json = format!(
r#"{{"sha256":"abc","size":1,"url":"{}","version":1,"signature":"sig"}}"#,
server.url()
);
let mock = server
.mock("GET", "/manifest.json")
.with_status(200)
.with_header("etag", "\"etag-from-server\"")
.with_body(&manifest_json)
.create();
let tmp = tempfile::tempdir().unwrap();
let url = format!("{}/manifest.json", server.url());
let result = fetch_with_state(&url, tmp.path());
mock.assert();
let m = result.expect("should succeed on 200");
assert_eq!(m.sha256, "abc");
assert_eq!(m.version, 1);
let key = super::manifest_cache_key(&url);
let state = tmp.path();
let etag_file = state.join(format!("{key}-etag"));
let body_file = state.join(format!("{key}-body"));
assert!(etag_file.exists(), "ETag should be persisted");
assert!(body_file.exists(), "body should be persisted");
assert_eq!(
std::fs::read_to_string(&etag_file).unwrap().trim(),
"\"etag-from-server\""
);
}
#[test]
fn transport_304_with_cached_body_returns_cached_manifest() {
let mut server = mockito::Server::new();
let mock = server
.mock("GET", "/manifest.json")
.match_header("if-none-match", "\"my-etag\"")
.with_status(304)
.create();
let tmp = tempfile::tempdir().unwrap();
let url = format!("{}/manifest.json", server.url());
let key = super::manifest_cache_key(&url);
let state = tmp.path();
std::fs::write(state.join(format!("{key}-etag")), "\"my-etag\"").unwrap();
let cached_json = r#"{"sha256":"cached","size":99,"url":"https://x.com/db.dat","version":42,"signature":"s"}"#;
std::fs::write(state.join(format!("{key}-body")), cached_json).unwrap();
let result = fetch_with_state(&url, tmp.path());
mock.assert();
let m = result.expect("should return cached manifest on 304");
assert_eq!(m.sha256, "cached");
assert_eq!(m.version, 42);
}
#[test]
fn transport_304_without_cache_retries_and_succeeds() {
let mut server = mockito::Server::new();
let mock_304 = server
.mock("GET", "/manifest.json")
.with_status(304)
.expect(1)
.create();
let retry_json = r#"{"sha256":"fresh","size":1,"url":"https://x.com/db.dat","version":7,"signature":"s"}"#;
let mock_200 = server
.mock("GET", "/manifest.json")
.with_status(200)
.with_header("etag", "\"new-etag\"")
.with_body(retry_json)
.expect(1)
.create();
let tmp = tempfile::tempdir().unwrap();
let url = format!("{}/manifest.json", server.url());
let key = super::manifest_cache_key(&url);
let state = tmp.path();
std::fs::write(state.join(format!("{key}-etag")), "\"stale\"").unwrap();
let result = fetch_with_state(&url, tmp.path());
mock_304.assert();
mock_200.assert();
let m = result.expect("retry after 304 should succeed");
assert_eq!(m.sha256, "fresh");
assert_eq!(m.version, 7);
let etag = std::fs::read_to_string(state.join(format!("{key}-etag"))).unwrap();
assert_eq!(etag.trim(), "\"new-etag\"");
}
#[test]
fn transport_404_returns_error() {
let mut server = mockito::Server::new();
let mock = server
.mock("GET", "/manifest.json")
.with_status(404)
.create();
let tmp = tempfile::tempdir().unwrap();
let url = format!("{}/manifest.json", server.url());
let result = fetch_with_state(&url, tmp.path());
mock.assert();
assert!(result.is_err());
assert!(result.unwrap_err().contains("404"));
}
#[test]
fn transport_invalid_json_not_cached() {
let mut server = mockito::Server::new();
let mock = server
.mock("GET", "/manifest.json")
.with_status(200)
.with_header("etag", "\"bad-etag\"")
.with_body("this is not json")
.create();
let tmp = tempfile::tempdir().unwrap();
let url = format!("{}/manifest.json", server.url());
let result = fetch_with_state(&url, tmp.path());
mock.assert();
assert!(result.is_err(), "invalid JSON should fail");
let key = super::manifest_cache_key(&url);
let state = tmp.path();
let body_file = state.join(format!("{key}-body"));
assert!(
!body_file.exists(),
"invalid JSON body should not be cached"
);
}
#[test]
fn transport_sends_user_agent_header() {
let mut server = mockito::Server::new();
let manifest_json = r#"{"sha256":"a","size":1,"url":"u","version":1,"signature":"s"}"#;
let mock = server
.mock("GET", "/manifest.json")
.match_header("user-agent", mockito::Matcher::Regex("tirith/".to_string()))
.with_status(200)
.with_body(manifest_json)
.create();
let tmp = tempfile::tempdir().unwrap();
let url = format!("{}/manifest.json", server.url());
let _ = fetch_with_state(&url, tmp.path());
mock.assert();
}
#[test]
fn read_bounded_bytes_rejects_declared_oversize_body() {
let err = super::read_bounded_bytes(
std::io::Cursor::new(b"abcd".to_vec()),
"https://example.test/feed",
Some(10),
4,
)
.unwrap_err();
assert!(err.contains("too large"));
}
#[test]
fn read_bounded_bytes_rejects_stream_that_exceeds_limit() {
let err = super::read_bounded_bytes(
std::io::Cursor::new(b"abcde".to_vec()),
"https://example.test/feed",
None,
4,
)
.unwrap_err();
assert!(err.contains("exceeded max size"));
}
}