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::selfupdate::SemVer;
use tirith_core::threatdb::{ThreatDb, ThreatDbWriter, ThreatSource, MAX_FORMAT_VERSION};
use tirith_core::threatdb_feeds::{
parse_domain_blocklist_reader, parse_phishtank_csv, parse_threatfox_zip,
parse_tor_exit_list_reader, parse_urlhaus_csv, MAX_FEED_ENTRIES, MAX_FEED_INPUT_BYTES,
};
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/download/threatdb-current/threatdb-manifest.json";
const INDEX_V2_URL_PRIMARY: &str =
"https://raw.githubusercontent.com/sheeki03/tirith/main/threatdb-index-v2.json";
const INDEX_V2_URL_FALLBACK: &str =
"https://github.com/sheeki03/tirith/releases/download/threatdb-current/threatdb-index-v2.json";
const MAX_MANIFEST_SIZE: u64 = 64 * 1024;
const MAX_DB_SIZE: u64 = 256 * 1024 * 1024;
const MAX_INDEX_ASSET_SIZE: u64 = MAX_DB_SIZE;
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 = MAX_FEED_INPUT_BYTES;
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";
fn guarded_http_client(timeout_secs: u64) -> Result<reqwest::blocking::Client, String> {
reqwest::blocking::Client::builder()
.no_proxy()
.dns_resolver(tirith_core::ssrf_guard::ssrf_guard_resolver())
.timeout(std::time::Duration::from_secs(timeout_secs))
.redirect(tirith_core::ssrf_guard::server_redirect_policy())
.build()
.map_err(|e| format!("HTTP client error: {e}"))
}
fn validate_remote_url(url: &str, purpose: &str) -> Result<(), String> {
tirith_core::url_validate::validate_server_url(url)
.map_err(|reason| format!("refusing unsafe {purpose} URL: {reason}"))
}
#[derive(Debug, Clone, 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 verify_key = VerifyingKey::from_bytes(VERIFY_KEY_BYTES)
.map_err(|e| format!("invalid embedded public key: {e}"))?;
self.verify_signature_with_key(&verify_key)
}
fn verify_signature_with_key(&self, verify_key: &VerifyingKey) -> 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 payload = self.canonical_payload();
use ed25519_dalek::Verifier;
verify_key
.verify(payload.as_bytes(), &signature)
.map_err(|_| "manifest signature verification failed".to_string())
}
}
#[derive(Debug, Clone, serde::Deserialize)]
struct IndexAsset {
format: u32,
filename: String,
url: String,
sha256: String,
size: u64,
#[serde(default)]
min_tirith_version: Option<String>,
}
#[derive(Debug, Clone, serde::Deserialize)]
struct IndexV2 {
manifest_version: u64,
sequence: u64,
assets: Vec<IndexAsset>,
signature: String,
}
const SIGNED_MANIFEST_VERSION: u64 = 2;
impl IndexV2 {
fn validate_generation(&self) -> Result<(), String> {
if self.assets.len() != 2 {
return Err(format!(
"v2 index must contain exactly one v1 and one v2 asset, got {}",
self.assets.len()
));
}
for format in [1u32, 2u32] {
let count = self
.assets
.iter()
.filter(|asset| asset.format == format)
.count();
if count != 1 {
return Err(format!(
"v2 index must contain exactly one format-v{format} asset, got {count}"
));
}
}
for (index, asset) in self.assets.iter().enumerate() {
for other in self.assets.iter().skip(index + 1) {
if asset.filename == other.filename {
return Err(format!(
"v2 index assets must have distinct filenames, duplicate {:?}",
asset.filename
));
}
if asset.url == other.url {
return Err(format!(
"v2 index assets must have distinct URLs, duplicate {:?}",
asset.url
));
}
}
}
for asset in &self.assets {
if asset.size == 0 || asset.size > MAX_INDEX_ASSET_SIZE {
return Err(format!(
"format-v{} asset has invalid size {}",
asset.format, asset.size
));
}
if asset.sha256.len() != 64
|| !asset.sha256.bytes().all(|byte| byte.is_ascii_hexdigit())
{
return Err(format!(
"format-v{} asset has an invalid SHA-256",
asset.format
));
}
let parsed = url::Url::parse(&asset.url).map_err(|error| {
format!("format-v{} asset URL is invalid: {error}", asset.format)
})?;
let url_filename = parsed
.path_segments()
.and_then(|mut segments| segments.next_back());
if url_filename != Some(asset.filename.as_str()) {
return Err(format!(
"format-v{} asset filename does not match its signed URL",
asset.format
));
}
}
Ok(())
}
fn canonical_payload(&self) -> String {
let assets: Vec<serde_json::Value> = self
.assets
.iter()
.map(|a| {
let mut m = serde_json::Map::new();
m.insert("filename".to_string(), serde_json::json!(a.filename));
m.insert("format".to_string(), serde_json::json!(a.format));
if let Some(ref v) = a.min_tirith_version {
m.insert("min_tirith_version".to_string(), serde_json::json!(v));
}
m.insert("sha256".to_string(), serde_json::json!(a.sha256));
m.insert("size".to_string(), serde_json::json!(a.size));
m.insert("url".to_string(), serde_json::json!(a.url));
serde_json::Value::Object(m)
})
.collect();
let mut top = serde_json::Map::new();
top.insert("assets".to_string(), serde_json::Value::Array(assets));
top.insert(
"manifest_version".to_string(),
serde_json::json!(self.manifest_version),
);
top.insert("sequence".to_string(), serde_json::json!(self.sequence));
serde_json::Value::Object(top).to_string()
}
fn verify_signature_with_key(&self, verify_key: &VerifyingKey) -> Result<(), String> {
let sig_bytes =
base64::Engine::decode(&base64::engine::general_purpose::STANDARD, &self.signature)
.map_err(|e| format!("invalid v2 index signature encoding: {e}"))?;
if sig_bytes.len() != SIGNATURE_LENGTH {
return Err(format!(
"v2 index signature wrong length: {} (expected {})",
sig_bytes.len(),
SIGNATURE_LENGTH
));
}
let signature = Signature::from_slice(&sig_bytes)
.map_err(|e| format!("invalid v2 index signature: {e}"))?;
let payload = self.canonical_payload();
use ed25519_dalek::Verifier;
verify_key
.verify(payload.as_bytes(), &signature)
.map_err(|_| "v2 index signature verification failed".to_string())?;
if self.manifest_version != SIGNED_MANIFEST_VERSION {
return Err(format!(
"v2 index manifest_version {} is unsupported (expected signed schema {}); falling back to v1",
self.manifest_version, SIGNED_MANIFEST_VERSION
));
}
self.validate_generation()
}
fn select_asset(&self, current_version: &str) -> Option<&IndexAsset> {
let current = SemVer::parse(current_version);
let mut compatible = self
.assets
.iter()
.filter(|a| a.format <= MAX_FORMAT_VERSION)
.filter(|a| a.size <= MAX_INDEX_ASSET_SIZE)
.filter(|a| match &a.min_tirith_version {
None => true,
Some(min) => match (SemVer::parse(min), current) {
(Some(min_v), Some(cur_v)) => cur_v >= min_v,
_ => false,
},
});
let best = compatible.next()?;
let (top, top_count) = compatible.fold((best, 1usize), |(top, count), a| {
use std::cmp::Ordering;
match a.format.cmp(&top.format) {
Ordering::Greater => (a, 1),
Ordering::Equal => (top, count + 1),
Ordering::Less => (top, count),
}
});
if top_count > 1 {
return None;
}
Some(top)
}
}
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
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum UpdateOutcome {
Installed,
AlreadyCurrent,
NoCompatibleAsset,
}
fn do_update(force: bool) -> Result<(), String> {
let outcome = match try_v2_index_update(force) {
Ok(UpdateOutcome::NoCompatibleAsset) => {
do_update_legacy(force)?
}
Ok(other) => other,
Err(e) => {
eprintln!("tirith: v2 index unavailable ({e}), falling back to legacy manifest...");
do_update_legacy(force)?
}
};
ThreatDb::refresh_cache();
if let Err(e) = reconcile_supplemental_after_primary(outcome, || {
update_supplemental_db(&policy::Policy::discover(None))
}) {
eprintln!("tirith: warning: supplemental threat DB update failed: {e}");
}
Ok(())
}
fn reconcile_supplemental_after_primary<F>(
outcome: UpdateOutcome,
reconcile: F,
) -> Result<(), String>
where
F: FnOnce() -> Result<(), String>,
{
match outcome {
UpdateOutcome::Installed | UpdateOutcome::AlreadyCurrent => reconcile(),
UpdateOutcome::NoCompatibleAsset => Err(
"internal error: supplemental reconciliation reached before primary selection"
.to_string(),
),
}
}
fn try_v2_index_update(force: bool) -> Result<UpdateOutcome, String> {
let index = fetch_index_v2()?;
let current_tirith_version = env!("CARGO_PKG_VERSION");
let asset = match index.select_asset(current_tirith_version) {
Some(a) => a,
None => {
eprintln!(
"tirith: v2 index has no asset compatible with this build (max format {}, tirith {}); using legacy manifest",
MAX_FORMAT_VERSION, current_tirith_version
);
return Ok(UpdateOutcome::NoCompatibleAsset);
}
};
let current = ThreatDb::cached().map(|db| (db.build_sequence(), db.stats().format_version));
if !index_install_needed(index.sequence, asset.format, current, force)? {
eprintln!(
"tirith: threat DB is already up to date (v2 index sequence {}, format v{})",
index.sequence, asset.format
);
return Ok(UpdateOutcome::AlreadyCurrent);
}
if asset.size > MAX_DB_SIZE {
return Err(format!(
"v2 asset too large: {} bytes (max {})",
asset.size, MAX_DB_SIZE
));
}
eprintln!(
"tirith: downloading threat DB (format v{}, seq {}) from v2 index...",
asset.format, index.sequence
);
let data = download_url(&asset.url, asset.size)?;
let computed_hash = hex::encode(Sha256::digest(&data));
if computed_hash != asset.sha256 {
return Err(format!(
"v2 asset SHA-256 mismatch: expected {}, got {}",
asset.sha256, computed_hash
));
}
let equal_sequence_format_switch = !force
&& current
.is_some_and(|(sequence, format)| sequence == index.sequence && format != asset.format);
install_primary_db(
data,
asset.format,
index.sequence,
force || equal_sequence_format_switch,
)?;
if asset.format == 1 {
retire_primary_v2()?;
}
Ok(UpdateOutcome::Installed)
}
fn index_install_needed(
index_sequence: u64,
selected_format: u32,
current: Option<(u64, u32)>,
force: bool,
) -> Result<bool, String> {
if force {
return Ok(true);
}
let Some((current_sequence, current_format)) = current else {
return Ok(true);
};
if index_sequence < current_sequence {
return Err(format!(
"rollback protection: v2 index sequence {index_sequence} < current {current_sequence}"
));
}
Ok(index_sequence > current_sequence || selected_format != current_format)
}
fn do_update_legacy(force: bool) -> Result<UpdateOutcome, String> {
let manifest = fetch_manifest()?;
manifest.verify_signature()?;
let current = ThreatDb::cached().map(|db| (db.build_sequence(), db.stats().format_version));
let install_needed = legacy_install_needed(manifest.version, current, force)?;
if !install_needed {
eprintln!(
"tirith: threat DB is already up to date (version {})",
manifest.version
);
return Ok(UpdateOutcome::AlreadyCurrent);
}
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 equal_sequence_format_switch = !force
&& current.is_some_and(|(sequence, format)| sequence == manifest.version && format == 2);
install_primary_db(
data,
1,
manifest.version,
force || equal_sequence_format_switch,
)?;
retire_primary_v2()?;
Ok(UpdateOutcome::Installed)
}
fn legacy_install_needed(
manifest_version: u64,
current: Option<(u64, u32)>,
force: bool,
) -> Result<bool, String> {
if force {
return Ok(true);
}
let Some((current_sequence, current_format)) = current else {
return Ok(true);
};
if manifest_version < current_sequence {
return Err(format!(
"rollback protection: manifest version {manifest_version} < current {current_sequence}"
));
}
Ok(manifest_version > current_sequence || current_format == 2)
}
fn primary_db_dest(format: u32) -> Result<PathBuf, String> {
match format {
2 => ThreatDb::default_path_v2().ok_or_else(|| "cannot determine v2 data path".to_string()),
_ => ThreatDb::default_path().ok_or_else(|| "cannot determine data directory".to_string()),
}
}
fn install_primary_db(data: Vec<u8>, format: u32, version: u64, force: bool) -> Result<(), String> {
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 stamped = db.stats().format_version;
if stamped != format {
return Err(format!(
"DB format mismatch: index/manifest declared format {format} but the downloaded file is format {stamped}"
));
}
if db.build_sequence() != version {
return Err(format!(
"DB sequence mismatch: index/manifest declared {version} but the downloaded file is sequence {}",
db.build_sequence()
));
}
let dest = primary_db_dest(format)?;
atomic_write(&dest, &data)?;
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{version} (format v{format}, {total_entries} entries)"
);
Ok(())
}
fn retire_primary_v2() -> Result<(), String> {
let v2_path = ThreatDb::default_path_v2()
.ok_or_else(|| "cannot determine v2 data path for retirement".to_string())?;
if ThreatDb::default_path().as_ref() == Some(&v2_path) {
return Err("refusing to retire v2 because it aliases the v1 data path".to_string());
}
let parent = v2_path
.parent()
.ok_or_else(|| "cannot determine v2 data directory".to_string())?;
let removed = match std::fs::remove_file(&v2_path) {
Ok(()) => true,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => false,
Err(error) => {
return Err(format!(
"failed to retire local v2 threat DB {}: {error}",
v2_path.display()
));
}
};
sync_parent_directory(parent)?;
if removed {
eprintln!(
"tirith: retired local v2 threat DB after verified legacy install ({})",
v2_path.display()
);
}
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,
) -> Result<usize, String> {
self.ingest_with_limit(entries, source, MAX_FEED_ENTRIES)
}
fn ingest_with_limit(
&mut self,
entries: tirith_core::threatdb_feeds::FeedEntries,
source: ThreatSource,
limit: usize,
) -> Result<usize, String> {
let count = entries
.hostnames
.len()
.checked_add(entries.ips.len())
.ok_or_else(|| "supplemental feed entry count overflow".to_string())?;
let current = self
.hostnames
.len()
.checked_add(self.ips.len())
.ok_or_else(|| "supplemental aggregate entry count overflow".to_string())?;
let projected = current
.checked_add(count)
.ok_or_else(|| "supplemental aggregate entry count overflow".to_string())?;
if projected > limit {
return Err(format!(
"supplemental feeds exceed the aggregate indicator limit of {limit}"
));
}
self.hostnames
.extend(entries.hostnames.into_iter().map(|h| (h, source)));
self.ips
.extend(entries.ips.into_iter().map(|ip| (ip, source)));
Ok(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 {
remove_disabled_supplemental(&supplemental_path)?;
ThreatDb::refresh_cache();
return Ok(());
}
let client = guarded_http_client(SUPPLEMENTAL_DOWNLOAD_TIMEOUT_SECS)
.map_err(|e| format!("supplemental feed {e}"))?;
let mut supplemental = SupplementalEntries::default();
let mut attempted_feeds = 0usize;
let mut failed_feeds: Vec<&str> = Vec::new();
if let Some(auth_key) = policy.threat_intel.abusech_auth_key.as_deref() {
if !auth_key.trim().is_empty() {
attempted_feeds += 1;
if !log_feed_result(
"URLhaus",
fetch_urlhaus_feed(&client, auth_key.trim(), &mut supplemental),
) {
failed_feeds.push("URLhaus");
}
attempted_feeds += 1;
if !log_feed_result(
"ThreatFox",
fetch_threatfox_feed(&client, auth_key.trim(), &mut supplemental),
) {
failed_feeds.push("ThreatFox");
}
}
}
if policy.threat_intel.phishing_army_enabled {
attempted_feeds += 1;
if !log_feed_result(
"Phishing Army",
fetch_phishing_army_feed(&client, &mut supplemental),
) {
failed_feeds.push("Phishing Army");
}
attempted_feeds += 1;
if !log_feed_result(
"PhishTank",
fetch_phishtank_feed(&client, &mut supplemental),
) {
failed_feeds.push("PhishTank");
}
}
attempted_feeds += 1;
if !log_feed_result("Tor exit", fetch_tor_exit_feed(&client, &mut supplemental)) {
failed_feeds.push("Tor exit");
}
if !failed_feeds.is_empty() {
eprintln!(
"tirith: warning: supplemental feed(s) failed ({}); keeping the existing supplemental threat DB unchanged",
failed_feeds.join(", ")
);
return Ok(());
}
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 remove_disabled_supplemental(path: &std::path::Path) -> Result<(), String> {
match std::fs::remove_file(path) {
Ok(()) => Ok(()),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(format!(
"failed to remove disabled supplemental threat DB {}: {error}",
path.display()
)),
}
}
fn log_feed_result(feed_name: &str, result: Result<usize, String>) -> bool {
match result {
Ok(0) => {
eprintln!("tirith: warning: {feed_name} feed returned no entries");
false
}
Ok(_) => true,
Err(e) => {
eprintln!("tirith: warning: {feed_name} feed failed: {e}");
false
}
}
}
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 response = fetch_feed_response(client, &url)?;
let entries = parse_urlhaus_csv(response).map_err(|e| format!("URLhaus parse failed: {e}"))?;
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))?;
supplemental.ingest(entries, ThreatSource::ThreatFoxIoc)
}
fn fetch_phishing_army_feed(
client: &reqwest::blocking::Client,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let response = fetch_feed_response(client, PHISHING_ARMY_URL)?;
let entries = parse_domain_blocklist_reader(response)
.map_err(|e| format!("Phishing Army parse failed: {e}"))?;
supplemental.ingest(entries, ThreatSource::PhishingArmy)
}
fn fetch_phishtank_feed(
client: &reqwest::blocking::Client,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let response = fetch_feed_response(client, PHISHTANK_URL)?;
let entries =
parse_phishtank_csv(response).map_err(|e| format!("PhishTank parse failed: {e}"))?;
supplemental.ingest(entries, ThreatSource::PhishTank)
}
fn fetch_tor_exit_feed(
client: &reqwest::blocking::Client,
supplemental: &mut SupplementalEntries,
) -> Result<usize, String> {
let response = fetch_feed_response(client, TOR_EXIT_URL)?;
let entries =
parse_tor_exit_list_reader(response).map_err(|e| format!("Tor exit parse failed: {e}"))?;
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_feed_response(
client: &reqwest::blocking::Client,
url: &str,
) -> Result<reqwest::blocking::Response, String> {
let safe = redact_url(url);
validate_remote_url(url, "supplemental feed")?;
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| {
let reason = if e.is_timeout() {
"timed out"
} else if e.is_connect() {
"connection failed"
} else if e.is_status() {
"unexpected HTTP status"
} else {
"request failed"
};
format!("fetch failed for {safe}: {reason}")
})?;
if response
.content_length()
.is_some_and(|length| length > MAX_SUPPLEMENTAL_FEED_SIZE)
{
return Err(format!(
"response body for {safe} exceeds {} bytes",
MAX_SUPPLEMENTAL_FEED_SIZE
));
}
Ok(response)
}
fn fetch_bytes(client: &reqwest::blocking::Client, url: &str) -> Result<Vec<u8>, String> {
let safe = redact_url(url);
let response = fetch_feed_response(client, url)?;
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)]
pub(crate) struct ThreatDbStatus {
pub(crate) installed: bool,
pub(crate) path: Option<String>,
pub(crate) age_hours: Option<f64>,
pub(crate) build_timestamp: Option<u64>,
pub(crate) build_sequence: Option<u64>,
pub(crate) package_count: Option<u32>,
pub(crate) hostname_count: Option<u32>,
pub(crate) ip_count: Option<u32>,
pub(crate) typosquat_count: Option<u32>,
pub(crate) popular_count: Option<u32>,
pub(crate) total_entries: Option<u32>,
pub(crate) skipped_range_only: Option<u32>,
pub(crate) signature_valid: Option<bool>,
pub(crate) stale: bool,
pub(crate) error: Option<String>,
}
pub(crate) fn gather_status() -> ThreatDbStatus {
let db_path = ThreatDb::resolve_primary_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(offline_flag: bool) {
if offline_flag || super::offline_env_active() {
return;
}
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> {
let verify_key = VerifyingKey::from_bytes(VERIFY_KEY_BYTES)
.map_err(|error| format!("invalid embedded public key: {error}"))?;
fetch_manifest_with(fetch_manifest_from, &verify_key)
}
fn fetch_verified_manifest_candidate<F>(
fetch: &mut F,
url: &str,
verify_key: &VerifyingKey,
) -> Result<Manifest, String>
where
F: FnMut(&str) -> Result<Manifest, String>,
{
let manifest = fetch(url)?;
manifest.verify_signature_with_key(verify_key)?;
Ok(manifest)
}
fn fetch_manifest_with<F>(mut fetch: F, verify_key: &VerifyingKey) -> Result<Manifest, String>
where
F: FnMut(&str) -> Result<Manifest, String>,
{
let primary = fetch_verified_manifest_candidate(&mut fetch, MANIFEST_URL_PRIMARY, verify_key);
let fallback = fetch_verified_manifest_candidate(&mut fetch, MANIFEST_URL_FALLBACK, verify_key);
match (primary, fallback) {
(Ok(primary), Ok(fallback)) => match primary.version.cmp(&fallback.version) {
std::cmp::Ordering::Greater => Ok(primary),
std::cmp::Ordering::Less => Ok(fallback),
std::cmp::Ordering::Equal => {
if primary.canonical_payload() != fallback.canonical_payload() {
return Err(format!(
"legacy manifest equivocation: primary and fallback both claim version {} with different signed payloads",
primary.version
));
}
Ok(primary)
}
},
(Ok(primary), Err(fallback_error)) => {
eprintln!(
"tirith: legacy manifest fallback unavailable or invalid ({fallback_error}); using verified primary"
);
Ok(primary)
}
(Err(primary_error), Ok(fallback)) => {
eprintln!(
"tirith: legacy manifest primary unavailable or invalid ({primary_error}); using verified fallback"
);
Ok(fallback)
}
(Err(primary_error), Err(fallback_error)) => Err(format!(
"legacy manifest fetch/verification failed: primary: {primary_error}; fallback: {fallback_error}"
)),
}
}
#[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 = hex::encode(&hash[..8]);
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> {
validate_remote_url(url, "threat DB manifest")?;
let client = guarded_http_client(MANIFEST_TIMEOUT_SECS)?;
fetch_manifest_from_with_state_and_client(url, state, &client)
}
fn fetch_manifest_from_with_state_and_client(
url: &str,
state: Option<std::path::PathBuf>,
client: &reqwest::blocking::Client,
) -> Result<Manifest, String> {
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_bytes = read_bounded_bytes(resp, "manifest", None, MAX_MANIFEST_SIZE)?;
let body = String::from_utf8(body_bytes)
.map_err(|e| format!("manifest body is not valid UTF-8: {e}"))?;
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_bytes =
read_bounded_bytes(retry_resp, "manifest-retry", None, MAX_MANIFEST_SIZE)?;
let retry_body = String::from_utf8(retry_body_bytes)
.map_err(|e| format!("retry body is not valid UTF-8: {e}"))?;
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
));
}
download_url(&manifest.url, manifest.size)
}
fn download_url(url: &str, declared_size: u64) -> Result<Vec<u8>, String> {
if declared_size > MAX_DB_SIZE {
return Err(format!(
"DB file too large: {} bytes (max {})",
declared_size, MAX_DB_SIZE
));
}
validate_remote_url(url, "threat DB asset")?;
let client = guarded_http_client(DB_DOWNLOAD_TIMEOUT_SECS)?;
let resp = client
.get(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 = read_bounded_bytes(resp, "threatdb", None, MAX_DB_SIZE)?;
Ok(bytes)
}
fn fetch_index_v2() -> Result<IndexV2, String> {
let verify_key = VerifyingKey::from_bytes(VERIFY_KEY_BYTES)
.map_err(|error| format!("invalid embedded public key: {error}"))?;
fetch_index_v2_with(fetch_index_v2_from, &verify_key)
}
fn fetch_verified_index_candidate<F>(
fetch: &mut F,
url: &str,
verify_key: &VerifyingKey,
) -> Result<IndexV2, String>
where
F: FnMut(&str) -> Result<IndexV2, String>,
{
let index = fetch(url)?;
index.verify_signature_with_key(verify_key)?;
Ok(index)
}
fn fetch_index_v2_with<F>(mut fetch: F, verify_key: &VerifyingKey) -> Result<IndexV2, String>
where
F: FnMut(&str) -> Result<IndexV2, String>,
{
let primary = fetch_verified_index_candidate(&mut fetch, INDEX_V2_URL_PRIMARY, verify_key);
let fallback = fetch_verified_index_candidate(&mut fetch, INDEX_V2_URL_FALLBACK, verify_key);
match (primary, fallback) {
(Ok(primary), Ok(fallback)) => match primary.sequence.cmp(&fallback.sequence) {
std::cmp::Ordering::Greater => Ok(primary),
std::cmp::Ordering::Less => Ok(fallback),
std::cmp::Ordering::Equal => {
if primary.canonical_payload() != fallback.canonical_payload() {
return Err(format!(
"v2 index equivocation: primary and fallback both claim sequence {} with different signed generations",
primary.sequence
));
}
Ok(primary)
}
},
(Ok(primary), Err(fallback_err)) => {
eprintln!(
"tirith: v2 index fallback unavailable or invalid ({fallback_err}); using verified primary"
);
Ok(primary)
}
(Err(primary_err), Ok(fallback)) => {
eprintln!(
"tirith: v2 index primary unavailable or invalid ({primary_err}); using verified fallback"
);
Ok(fallback)
}
(Err(primary_err), Err(fallback_err)) => Err(format!(
"v2 index fetch/verification failed: primary: {primary_err}; fallback: {fallback_err}"
)),
}
}
fn fetch_index_v2_from(url: &str) -> Result<IndexV2, String> {
validate_remote_url(url, "threat DB index")?;
let client = guarded_http_client(MANIFEST_TIMEOUT_SECS)?;
let resp = client
.get(url)
.header(
"User-Agent",
format!("tirith/{}", env!("CARGO_PKG_VERSION")),
)
.send()
.map_err(|e| format!("v2 index fetch failed: {e}"))?;
if !resp.status().is_success() {
return Err(format!("v2 index HTTP {}", resp.status()));
}
let content_len = resp.content_length().unwrap_or(0);
if content_len > MAX_MANIFEST_SIZE {
return Err(format!(
"v2 index too large: {} bytes (max {})",
content_len, MAX_MANIFEST_SIZE
));
}
let body_bytes = read_bounded_bytes(resp, "v2-index", None, MAX_MANIFEST_SIZE)?;
let body = String::from_utf8(body_bytes)
.map_err(|e| format!("v2 index body is not valid UTF-8: {e}"))?;
serde_json::from_str::<IndexV2>(&body).map_err(|e| format!("invalid v2 index JSON: {e}"))
}
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.as_file()
.sync_all()
.map_err(|e| format!("failed to sync temp file: {e}"))?;
let persisted = tmp
.persist(dest)
.map_err(|e| format!("failed to rename temp file: {e}"))?;
persisted
.sync_all()
.map_err(|e| format!("failed to sync installed file: {e}"))?;
sync_parent_directory(parent)?;
Ok(())
}
fn sync_parent_directory(parent: &std::path::Path) -> Result<(), String> {
#[cfg(unix)]
{
std::fs::File::open(parent)
.and_then(|directory| directory.sync_all())
.map_err(|error| {
format!(
"failed to sync containing directory {}: {error}",
parent.display()
)
})?;
}
#[cfg(not(unix))]
let _ = parent;
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 {
use std::fmt::Write as _;
pub fn encode(data: impl AsRef<[u8]>) -> String {
let bytes = data.as_ref();
bytes
.iter()
.fold(String::with_capacity(bytes.len() * 2), |mut s, b| {
let _ = write!(s, "{b:02x}");
s
})
}
}
use std::net::Ipv4Addr;
use tirith_core::threatdb::{Confidence, Ecosystem, SourceTier};
const HISTORY_FILE: &str = "threatdb-history.jsonl";
const HISTORY_MAX_LINES: usize = 64;
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
struct CategoryCounts {
packages: u64,
hostnames: u64,
ips: u64,
typosquats: u64,
popular: u64,
}
impl CategoryCounts {
fn total(&self) -> u64 {
self.packages + self.hostnames + self.ips + self.typosquats + self.popular
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
struct DbSnapshot {
recorded_at: u64,
build_sequence: u64,
build_timestamp: u64,
signature_valid: bool,
counts: CategoryCounts,
#[serde(default)]
sources: std::collections::BTreeMap<String, u64>,
}
fn history_path() -> Option<PathBuf> {
policy::state_dir().map(|d| d.join(HISTORY_FILE))
}
fn current_snapshot() -> Option<DbSnapshot> {
let db = ThreatDb::cached()?;
let stats = db.stats();
let breakdown = db.source_breakdown();
let mut sources = std::collections::BTreeMap::new();
for (src, count) in breakdown.per_source() {
sources.insert(src.as_str().to_string(), *count);
}
Some(DbSnapshot {
recorded_at: unix_now(),
build_sequence: stats.build_sequence,
build_timestamp: stats.build_timestamp,
signature_valid: db.verify_signature().is_ok(),
counts: CategoryCounts {
packages: stats.package_count as u64,
hostnames: stats.hostname_count as u64,
ips: stats.ip_count as u64,
typosquats: stats.typosquat_count as u64,
popular: stats.popular_count as u64,
},
sources,
})
}
fn load_history() -> (Vec<DbSnapshot>, Option<String>) {
let Some(path) = history_path() else {
return (Vec::new(), None);
};
let content = match std::fs::read_to_string(&path) {
Ok(c) => c,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return (Vec::new(), None),
Err(e) => {
return (
Vec::new(),
Some(format!(
"could not read snapshot history at {} ({e}) — check file permissions; \
the diff below cannot use any earlier snapshot",
path.display()
)),
);
}
};
let snapshots = content
.lines()
.filter(|l| !l.trim().is_empty())
.filter_map(|l| serde_json::from_str::<DbSnapshot>(l).ok())
.collect();
(snapshots, None)
}
fn record_snapshot(snapshot: &DbSnapshot) {
let Some(path) = history_path() else {
return;
};
let (mut history, _) = load_history();
if history.iter().any(|s| {
s.build_sequence == snapshot.build_sequence
&& s.build_timestamp == snapshot.build_timestamp
&& s.signature_valid == snapshot.signature_valid
&& s.counts == snapshot.counts
&& s.sources == snapshot.sources
}) {
return;
}
history.push(snapshot.clone());
if history.len() > HISTORY_MAX_LINES {
let drop = history.len() - HISTORY_MAX_LINES;
history.drain(0..drop);
}
if let Some(parent) = path.parent() {
if std::fs::create_dir_all(parent).is_err() {
return;
}
}
let mut body = String::new();
for s in &history {
if let Ok(line) = serde_json::to_string(s) {
body.push_str(&line);
body.push('\n');
}
}
let _ = atomic_write(&path, body.as_bytes());
}
fn snapshot_current_db() {
if let Some(snapshot) = current_snapshot() {
record_snapshot(&snapshot);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "lowercase")]
enum IndicatorKind {
Ip,
Package,
Domain,
}
struct ParsedIndicator {
kind: IndicatorKind,
ecosystem: Option<Ecosystem>,
version: Option<String>,
value: String,
}
fn parse_indicator(raw: &str) -> ParsedIndicator {
let trimmed = raw.trim();
if let Ok(ip) = trimmed.parse::<Ipv4Addr>() {
return ParsedIndicator {
kind: IndicatorKind::Ip,
ecosystem: None,
version: None,
value: ip.to_string(),
};
}
if let Some((prefix, rest)) = trimmed.split_once(':') {
if let Some(eco) = Ecosystem::from_name(prefix) {
let (name, version) = split_name_version(rest);
return ParsedIndicator {
kind: IndicatorKind::Package,
ecosystem: Some(eco),
version,
value: name,
};
}
}
if let Some((name, version)) = split_at_version(trimmed) {
return ParsedIndicator {
kind: IndicatorKind::Package,
ecosystem: None,
version: Some(version),
value: name,
};
}
if trimmed.contains('.') && !trimmed.contains('/') && !trimmed.contains(char::is_whitespace) {
return ParsedIndicator {
kind: IndicatorKind::Domain,
ecosystem: None,
version: None,
value: trimmed.to_ascii_lowercase(),
};
}
ParsedIndicator {
kind: IndicatorKind::Package,
ecosystem: None,
version: None,
value: trimmed.to_string(),
}
}
fn split_at_version(s: &str) -> Option<(String, String)> {
let search_from = if s.starts_with('@') { 1 } else { 0 };
let idx = s[search_from..].find('@')? + search_from;
let name = &s[..idx];
let version = &s[idx + 1..];
if name.is_empty() || version.is_empty() {
return None;
}
Some((name.to_string(), version.to_string()))
}
fn split_name_version(rest: &str) -> (String, Option<String>) {
match split_at_version(rest) {
Some((name, version)) => (name, Some(version)),
None => (rest.to_string(), None),
}
}
#[derive(Debug, serde::Serialize)]
struct ExplainResult {
indicator: String,
kind: IndicatorKind,
#[serde(skip_serializing_if = "Option::is_none")]
ecosystem: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
version: Option<String>,
present: bool,
db_missing: bool,
findings: Vec<ExplainFinding>,
}
#[derive(Debug, serde::Serialize)]
struct ExplainFinding {
classification: String,
#[serde(skip_serializing_if = "Option::is_none")]
source: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
source_label: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
confidence: Option<Confidence>,
detail: String,
#[serde(skip_serializing_if = "Option::is_none")]
reference_url: Option<String>,
}
pub fn explain(indicator: &str, json: bool) -> i32 {
let parsed = parse_indicator(indicator);
let db = ThreatDb::cached();
let mut findings: Vec<ExplainFinding> = Vec::new();
let db_missing = db.is_none();
if let Some(ref db) = db {
match parsed.kind {
IndicatorKind::Ip => {
if let Ok(ip) = parsed.value.parse::<Ipv4Addr>() {
if let Some(m) = db.check_ip(ip) {
findings.push(ExplainFinding {
classification: "malicious_ip".to_string(),
source: Some(m.source.as_str().to_string()),
source_label: Some(m.source.label().to_string()),
confidence: Some(m.confidence),
detail: format!(
"IP address is listed as malicious infrastructure by {}.",
m.source.label()
),
reference_url: m.reference_url,
});
}
}
}
IndicatorKind::Domain => {
if let Some(m) = db.check_hostname(&parsed.value) {
findings.push(ExplainFinding {
classification: "malicious_hostname".to_string(),
source: Some(m.source.as_str().to_string()),
source_label: Some(m.source.label().to_string()),
confidence: Some(m.confidence),
detail: format!(
"Hostname is listed as malicious infrastructure by {}.",
m.source.label()
),
reference_url: m.reference_url,
});
}
}
IndicatorKind::Package => {
let ecosystems: Vec<Ecosystem> = match parsed.ecosystem {
Some(e) => vec![e],
None => ALL_ECOSYSTEMS.to_vec(),
};
for eco in ecosystems {
explain_package(
db,
eco,
&parsed.value,
parsed.version.as_deref(),
&mut findings,
);
}
}
}
}
let result = ExplainResult {
indicator: indicator.trim().to_string(),
kind: parsed.kind,
ecosystem: parsed.ecosystem.map(|e| e.to_string()),
version: parsed.version.clone(),
present: !findings.is_empty(),
db_missing,
findings,
};
snapshot_current_db();
if json {
return print_json_value(&result);
}
print_explain_human(&result);
0
}
const ALL_ECOSYSTEMS: [Ecosystem; 8] = [
Ecosystem::Npm,
Ecosystem::PyPI,
Ecosystem::RubyGems,
Ecosystem::Crates,
Ecosystem::Go,
Ecosystem::Maven,
Ecosystem::NuGet,
Ecosystem::Packagist,
];
fn explain_package(
db: &ThreatDb,
eco: Ecosystem,
name: &str,
version: Option<&str>,
findings: &mut Vec<ExplainFinding>,
) {
if let Some(m) = db.check_package(eco, name, version) {
let versions = if m.all_versions_malicious {
"all versions".to_string()
} else {
"specific affected versions".to_string()
};
findings.push(ExplainFinding {
classification: "malicious_package".to_string(),
source: Some(m.source.as_str().to_string()),
source_label: Some(m.source.label().to_string()),
confidence: Some(m.confidence),
detail: format!(
"{} package '{}' is listed as malicious by {} ({}).",
eco,
name,
m.source.label(),
versions
),
reference_url: m.reference_url,
});
}
if let Some(ts) = db.check_typosquat(eco, name) {
findings.push(ExplainFinding {
classification: "typosquat".to_string(),
source: Some(ThreatSource::EcosystemsTyposquat.as_str().to_string()),
source_label: Some(ThreatSource::EcosystemsTyposquat.label().to_string()),
confidence: None,
detail: format!(
"{} package '{}' is a known typosquat of '{}'.",
eco, ts.malicious_name, ts.target_name
),
reference_url: None,
});
}
if let Some((popular, distance)) = db.check_popular_distance(eco, name) {
findings.push(ExplainFinding {
classification: "popular_lookalike".to_string(),
source: None,
source_label: None,
confidence: None,
detail: format!(
"{} package '{}' is edit-distance {} from the popular package '{}' \
— a possible slopsquat/typo. Not itself listed as malicious.",
eco, name, distance, popular
),
reference_url: None,
});
}
}
fn print_explain_human(r: &ExplainResult) {
println!("threat-db explain: {}", r.indicator);
let kind_label = match r.kind {
IndicatorKind::Ip => "IPv4 address",
IndicatorKind::Package => "package",
IndicatorKind::Domain => "domain / hostname",
};
print!(" type: {kind_label}");
if let Some(ref eco) = r.ecosystem {
print!(" ({eco})");
}
if let Some(ref v) = r.version {
print!(" @ {v}");
}
println!();
if r.db_missing {
println!(" result: threat DB not installed");
println!(" Hint: run 'tirith threat-db update' to install the signed DB.");
return;
}
if !r.present {
println!(" result: not present");
match r.kind {
IndicatorKind::Package => println!(
" The threat DB has no malicious-package, typosquat, or \
popular-lookalike record for this name."
),
IndicatorKind::Domain => {
println!(" The threat DB has no malicious-hostname record for this domain.")
}
IndicatorKind::Ip => {
println!(" The threat DB has no malicious-infrastructure record for this IP.")
}
}
println!(" Absence is not a guarantee of safety — the DB only covers known threats.");
return;
}
println!(" result: PRESENT — {} finding(s)", r.findings.len());
for (i, f) in r.findings.iter().enumerate() {
println!();
println!(" [{}] {}", i + 1, f.classification);
if let Some(ref label) = f.source_label {
println!(" source: {label}");
}
if let Some(c) = f.confidence {
println!(" confidence: {}", c.as_str());
}
println!(" {}", f.detail);
if let Some(ref url) = f.reference_url {
println!(" reference: {url}");
}
}
}
#[derive(Debug, serde::Serialize)]
struct SourcesReport {
db_installed: bool,
#[serde(skip_serializing_if = "Option::is_none")]
build_sequence: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
build_timestamp: Option<u64>,
sources: Vec<SourceInfo>,
}
#[derive(Debug, serde::Serialize)]
struct SourceInfo {
id: String,
name: String,
tier: SourceTier,
upstream_url: String,
record_count: Option<u64>,
}
pub fn sources(json: bool) -> i32 {
let db = ThreatDb::cached();
let breakdown = db.as_ref().map(|d| d.source_breakdown());
let stats = db.as_ref().map(|d| d.stats());
let mut source_infos = Vec::new();
for src in ThreatSource::ALL {
let record_count = breakdown.as_ref().map(|b| b.count_for(src));
source_infos.push(SourceInfo {
id: src.as_str().to_string(),
name: src.label().to_string(),
tier: src.tier(),
upstream_url: src.upstream_url().to_string(),
record_count,
});
}
let report = SourcesReport {
db_installed: db.is_some(),
build_sequence: stats.as_ref().map(|s| s.build_sequence),
build_timestamp: stats.as_ref().map(|s| s.build_timestamp),
sources: source_infos,
};
snapshot_current_db();
if json {
return print_json_value(&report);
}
print_sources_human(&report, breakdown.as_ref().map(|b| b.popular_count));
0
}
fn print_sources_human(r: &SourcesReport, popular_count: Option<u64>) {
println!("threat-db sources");
if r.db_installed {
if let (Some(seq), Some(ts)) = (r.build_sequence, r.build_timestamp) {
println!(" DB version {seq}, built {}", format_epoch(ts));
}
} else {
println!(" threat DB not installed — counts unavailable");
println!(" (run 'tirith threat-db update' to install the signed DB)");
}
for tier in [SourceTier::Primary, SourceTier::Supplemental] {
let heading = match tier {
SourceTier::Primary => "Primary feeds (signed CI database)",
SourceTier::Supplemental => "Supplemental feeds (optional user-local overlay)",
};
println!();
println!(" {heading}");
for s in r.sources.iter().filter(|s| s.tier == tier) {
let count = match s.record_count {
Some(c) => format!("{c} records"),
None => "count unavailable".to_string(),
};
println!(" {:<26} {}", s.name, count);
println!(" {}", s.upstream_url);
}
}
if r.db_installed {
println!();
println!(
" Note: typosquat counts are reported under the ecosyste.ms Typosquats feed; \
popular-package baselines ({} entries) are not a threat feed.",
popular_count.unwrap_or(0)
);
}
}
#[derive(Debug, serde::Serialize)]
struct HealthReport {
installed: bool,
path: Option<String>,
signature_valid: Option<bool>,
age_hours: Option<f64>,
refresh_interval_hours: u64,
stale: bool,
build_sequence: Option<u64>,
build_timestamp: Option<u64>,
counts: Option<CategoryCounts>,
supplemental: SupplementalHealth,
error: Option<String>,
status: String,
}
#[derive(Debug, serde::Serialize)]
struct SupplementalHealth {
present: bool,
path: Option<String>,
}
pub fn health(json: bool) -> i32 {
let report = gather_health();
snapshot_current_db();
let exit = if report.error.is_some() { 1 } else { 0 };
if json {
return print_json_value(&report).max(exit);
}
print_health_human(&report);
exit
}
fn gather_health() -> HealthReport {
let db_path = ThreatDb::resolve_primary_path();
let path_str = db_path.as_ref().map(|p| p.display().to_string());
let policy = policy::Policy::discover(None);
let refresh_interval_hours = policy.threat_intel.auto_update_hours;
let supplemental_path = ThreatDb::supplemental_path();
let supplemental = SupplementalHealth {
present: supplemental_path
.as_ref()
.map(|p| p.exists())
.unwrap_or(false),
path: supplemental_path.map(|p| p.display().to_string()),
};
let exists = db_path.as_ref().map(|p| p.exists()).unwrap_or(false);
if !exists {
return HealthReport {
installed: false,
path: path_str,
signature_valid: None,
age_hours: None,
refresh_interval_hours,
stale: false,
build_sequence: None,
build_timestamp: None,
counts: None,
supplemental,
error: None,
status: "not_installed".to_string(),
};
}
let db_path_ref = db_path.as_ref().expect("path exists when exists==true");
match ThreatDb::load_from_path(db_path_ref, 0) {
Ok(db) => {
let sig_valid = db.verify_signature().is_ok();
let stats = db.stats();
let age_secs = unix_now().saturating_sub(stats.build_timestamp);
let age_hours = age_secs as f64 / 3600.0;
let stale =
refresh_interval_hours != 0 && age_hours > (refresh_interval_hours as f64 * 2.0);
let counts = CategoryCounts {
packages: stats.package_count as u64,
hostnames: stats.hostname_count as u64,
ips: stats.ip_count as u64,
typosquats: stats.typosquat_count as u64,
popular: stats.popular_count as u64,
};
let status = if !sig_valid {
"error"
} else if stale {
"stale"
} else {
"ok"
};
HealthReport {
installed: true,
path: path_str,
signature_valid: Some(sig_valid),
age_hours: Some(age_hours),
refresh_interval_hours,
stale,
build_sequence: Some(stats.build_sequence),
build_timestamp: Some(stats.build_timestamp),
counts: Some(counts),
supplemental,
error: if sig_valid {
None
} else {
Some("Ed25519 signature verification failed".to_string())
},
status: status.to_string(),
}
}
Err(e) => HealthReport {
installed: true,
path: path_str,
signature_valid: None,
age_hours: None,
refresh_interval_hours,
stale: false,
build_sequence: None,
build_timestamp: None,
counts: None,
supplemental,
error: Some(format!("{e}")),
status: "error".to_string(),
},
}
}
fn print_health_human(r: &HealthReport) {
println!("threat-db health");
if !r.installed {
println!(" status: NOT INSTALLED");
if let Some(ref p) = r.path {
println!(" expected at: {p}");
}
println!(" Hint: run 'tirith threat-db update' to install the signed DB.");
print_supplemental_health(&r.supplemental);
return;
}
if let Some(ref err) = r.error {
println!(" status: ERROR — {err}");
if let Some(ref p) = r.path {
println!(" path: {p}");
}
println!(" Hint: re-download with 'tirith threat-db update --force'.");
print_supplemental_health(&r.supplemental);
return;
}
let status_label = match r.status.as_str() {
"ok" => "OK",
"stale" => "STALE",
other => other,
};
println!(" status: {status_label}");
if let Some(ref p) = r.path {
println!(" path: {p}");
}
match r.signature_valid {
Some(true) => println!(" signature: valid (Ed25519)"),
Some(false) => println!(" signature: INVALID"),
None => println!(" signature: unknown"),
}
if let Some(seq) = r.build_sequence {
println!(" version: {seq}");
}
if let Some(ts) = r.build_timestamp {
println!(" built: {}", format_epoch(ts));
}
if let Some(age) = r.age_hours {
println!(" age: {}", format_age(age));
}
if r.refresh_interval_hours == 0 {
println!(" refresh: auto-update disabled (auto_update_hours = 0)");
} else {
println!(
" refresh: every {}h (stale after {}h)",
r.refresh_interval_hours,
r.refresh_interval_hours * 2
);
if r.stale {
println!(" -> DB is stale; run 'tirith threat-db update'.");
}
}
if let Some(ref c) = r.counts {
println!(
" entries: {} total — {} packages, {} hostnames, {} IPs, {} typosquats, {} popular",
c.total(),
c.packages,
c.hostnames,
c.ips,
c.typosquats,
c.popular
);
}
print_supplemental_health(&r.supplemental);
}
fn print_supplemental_health(s: &SupplementalHealth) {
if s.present {
println!(" supplemental: present (user-local opt-in feed overlay)");
} else {
println!(" supplemental: none (no opt-in feeds configured)");
}
}
#[derive(Debug, serde::Serialize)]
struct DiffReport {
since: String,
since_kind: String,
baseline: Option<SnapshotSummary>,
current: Option<SnapshotSummary>,
delta: Option<CountDelta>,
#[serde(skip_serializing_if = "std::collections::BTreeMap::is_empty")]
source_delta: std::collections::BTreeMap<String, i64>,
limitation: String,
note: Option<String>,
}
#[derive(Debug, serde::Serialize)]
struct SnapshotSummary {
build_sequence: u64,
build_timestamp: u64,
recorded_at: u64,
counts: CategoryCounts,
}
#[derive(Debug, serde::Serialize)]
struct CountDelta {
packages: i64,
hostnames: i64,
ips: i64,
typosquats: i64,
popular: i64,
total: i64,
}
fn delta_of(current: &CategoryCounts, baseline: &CategoryCounts) -> CountDelta {
let d = |c: u64, b: u64| c as i64 - b as i64;
CountDelta {
packages: d(current.packages, baseline.packages),
hostnames: d(current.hostnames, baseline.hostnames),
ips: d(current.ips, baseline.ips),
typosquats: d(current.typosquats, baseline.typosquats),
popular: d(current.popular, baseline.popular),
total: d(current.total(), baseline.total()),
}
}
fn parse_since(since: &str) -> Result<(String, Option<u64>, Option<u64>), String> {
let s = since.trim();
if let Ok(version) = s.parse::<u64>() {
return Ok(("version".to_string(), Some(version), None));
}
if let Some(epoch) = parse_iso_date(s) {
return Ok(("date".to_string(), None, Some(epoch)));
}
Err(format!(
"could not parse --since value '{since}' — expected a DB version number \
(e.g. 42) or an ISO date (e.g. 2026-01-15)"
))
}
const MONTH_DAYS: [i64; 12] = [31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31];
fn is_leap_year(y: i64) -> bool {
(y % 4 == 0 && y % 100 != 0) || y % 400 == 0
}
fn parse_iso_date(s: &str) -> Option<u64> {
let date_part = s.split(['T', ' ']).next().unwrap_or(s);
let mut it = date_part.split('-');
let year: i64 = it.next()?.parse().ok()?;
let month: i64 = it.next()?.parse().ok()?;
let day: i64 = it.next()?.parse().ok()?;
if it.next().is_some() {
return None;
}
if !(1970..=9999).contains(&year) || !(1..=12).contains(&month) {
return None;
}
let max_day = if month == 2 && is_leap_year(year) {
29
} else {
MONTH_DAYS[(month - 1) as usize]
};
if !(1..=max_day).contains(&day) {
return None;
}
let mut days: i64 = 0;
for y in 1970..year {
days += if is_leap_year(y) { 366 } else { 365 };
}
for (m, md) in MONTH_DAYS.iter().enumerate() {
if (m as i64) + 1 >= month {
break;
}
days += md;
if (m as i64) + 1 == 2 && is_leap_year(year) {
days += 1;
}
}
days += day - 1;
Some((days * 86400) as u64)
}
pub fn diff(since: &str, json: bool) -> i32 {
snapshot_current_db();
let limitation = "The threat DB format retains no per-entry history, so this diff reports \
category and per-source COUNT deltas between recorded snapshots — not the \
exact entries added or removed. Snapshots accrue each time a transparency \
command runs."
.to_string();
let (since_kind, want_version, want_epoch) = match parse_since(since) {
Ok(v) => v,
Err(e) => {
if json {
let _ = print_json_value(&DiffReport {
since: since.to_string(),
since_kind: "invalid".to_string(),
baseline: None,
current: None,
delta: None,
source_delta: Default::default(),
limitation,
note: Some(e.clone()),
});
} else {
eprintln!("tirith: {e}");
}
return 1;
}
};
let (history, history_read_error) = load_history();
let current = current_snapshot();
let baseline = history
.iter()
.filter(|s| match (want_version, want_epoch) {
(Some(v), _) => s.build_sequence <= v,
(_, Some(e)) => s.recorded_at <= e,
_ => false,
})
.max_by_key(|s| (s.recorded_at, s.build_sequence))
.cloned();
let summarize = |s: &DbSnapshot| SnapshotSummary {
build_sequence: s.build_sequence,
build_timestamp: s.build_timestamp,
recorded_at: s.recorded_at,
counts: s.counts.clone(),
};
let (delta, source_delta, note) = match (&baseline, ¤t) {
(Some(b), Some(c)) => {
let d = delta_of(&c.counts, &b.counts);
let mut sd: std::collections::BTreeMap<String, i64> = std::collections::BTreeMap::new();
for (src, cur_count) in &c.sources {
let base_count = b.sources.get(src).copied().unwrap_or(0);
let diff = *cur_count as i64 - base_count as i64;
if diff != 0 {
sd.insert(src.clone(), diff);
}
}
let note = if b.build_sequence == c.build_sequence {
Some(
"Baseline and current snapshot are the same DB version — no \
change since the requested point."
.to_string(),
)
} else {
None
};
(Some(d), sd, note)
}
(None, Some(_)) => (
None,
Default::default(),
Some(history_read_error.clone().unwrap_or_else(|| {
format!(
"No snapshot was recorded at or before '{since}'. tirith only began \
retaining snapshots from the first transparency command after this \
feature was installed; a diff needs at least one earlier snapshot. \
Run 'tirith threat-db health' periodically to build up history."
)
})),
),
(_, None) => (
None,
Default::default(),
Some(
"Threat DB is not installed — nothing to diff. Run \
'tirith threat-db update' first."
.to_string(),
),
),
};
let report = DiffReport {
since: since.to_string(),
since_kind,
baseline: baseline.as_ref().map(summarize),
current: current.as_ref().map(summarize),
delta,
source_delta,
limitation,
note,
};
if json {
return print_json_value(&report);
}
print_diff_human(&report);
0
}
fn print_diff_human(r: &DiffReport) {
println!("threat-db diff (since {} = {})", r.since, r.since_kind);
println!(" note: {}", r.limitation);
if let (Some(b), Some(c)) = (&r.baseline, &r.current) {
println!();
println!(
" baseline: DB v{} built {} (snapshot recorded {})",
b.build_sequence,
format_epoch(b.build_timestamp),
format_epoch(b.recorded_at)
);
println!(
" current: DB v{} built {}",
c.build_sequence,
format_epoch(c.build_timestamp)
);
if let Some(ref d) = r.delta {
println!();
println!(" count change (current - baseline):");
print_delta_line("packages", d.packages);
print_delta_line("hostnames", d.hostnames);
print_delta_line("IPs", d.ips);
print_delta_line("typosquats", d.typosquats);
print_delta_line("popular", d.popular);
print_delta_line("TOTAL", d.total);
}
if !r.source_delta.is_empty() {
println!();
println!(" per-source count change:");
for (src, delta) in &r.source_delta {
print_delta_line(src, *delta);
}
}
}
if let Some(ref note) = r.note {
println!();
println!(" {note}");
}
}
fn print_delta_line(label: &str, delta: i64) {
let sign = if delta > 0 {
format!("+{delta}")
} else {
delta.to_string()
};
println!(" {label:<14} {sign}");
}
#[must_use]
fn print_json_value(value: &impl serde::Serialize) -> i32 {
match serde_json::to_string_pretty(value) {
Ok(s) => {
println!("{s}");
0
}
Err(e) => {
eprintln!("tirith: JSON serialization failed: {e}");
1
}
}
}
fn format_epoch(epoch: u64) -> String {
let days = epoch / 86400;
let secs_of_day = epoch % 86400;
let (hh, mm, ss) = (
secs_of_day / 3600,
(secs_of_day % 3600) / 60,
secs_of_day % 60,
);
let mut year: i64 = 1970;
let mut remaining = days as i64;
loop {
let year_len = if is_leap_year(year) { 366 } else { 365 };
if remaining < year_len {
break;
}
remaining -= year_len;
year += 1;
}
let mut month = 1;
for (m, md) in MONTH_DAYS.iter().enumerate() {
let mut len = *md;
if m == 1 && is_leap_year(year) {
len += 1;
}
if remaining < len {
break;
}
remaining -= len;
month += 1;
}
let day = remaining + 1;
format!("{year:04}-{month:02}-{day:02} {hh:02}:{mm:02}:{ss:02} UTC")
}
fn format_age(hours: f64) -> String {
if hours < 1.0 {
format!("{:.0} minutes", hours * 60.0)
} else if hours < 48.0 {
format!("{hours:.0} hours")
} else {
format!("{:.1} days", hours / 24.0)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cli::test_harness::{EnvGuard, ENV_LOCK};
use std::path::Path;
use std::sync::atomic::Ordering;
use tirith_core::threatdb::ThreatDbFormat;
fn signed_index_v2(sequence: u64, assets: Vec<IndexAsset>, key: &SigningKey) -> IndexV2 {
let mut idx = IndexV2 {
manifest_version: SIGNED_MANIFEST_VERSION,
sequence,
assets,
signature: String::new(),
};
let payload = idx.canonical_payload();
use ed25519_dalek::Signer;
let sig = key.sign(payload.as_bytes());
idx.signature =
base64::Engine::encode(&base64::engine::general_purpose::STANDARD, sig.to_bytes());
idx
}
fn asset(format: u32, min: Option<&str>) -> IndexAsset {
let filename = format!("tirith-threatdb-v{format}.dat");
IndexAsset {
format,
url: format!("https://example.com/{filename}"),
filename,
sha256: "00".repeat(32),
size: 1024,
min_tirith_version: min.map(str::to_string),
}
}
fn signed_manifest(version: u64, url: &str, sha256: &str, key: &SigningKey) -> Manifest {
let mut manifest = Manifest {
sha256: sha256.to_string(),
size: 1024,
url: url.to_string(),
version,
signature: String::new(),
};
use ed25519_dalek::Signer;
let signature = key.sign(manifest.canonical_payload().as_bytes());
manifest.signature = base64::Engine::encode(
&base64::engine::general_purpose::STANDARD,
signature.to_bytes(),
);
manifest
}
#[test]
fn already_current_primary_still_reconciles_supplemental_state() {
for outcome in [UpdateOutcome::Installed, UpdateOutcome::AlreadyCurrent] {
let mut calls = 0usize;
reconcile_supplemental_after_primary(outcome, || {
calls += 1;
Ok(())
})
.unwrap();
assert_eq!(calls, 1, "{outcome:?} must reconcile exactly once");
}
let mut calls = 0usize;
assert!(
reconcile_supplemental_after_primary(UpdateOutcome::NoCompatibleAsset, || {
calls += 1;
Ok(())
})
.is_err()
);
assert_eq!(
calls, 0,
"an unresolved primary must not publish an overlay"
);
}
#[test]
fn disabled_supplemental_removal_is_idempotent_and_reports_real_failures() {
let root = tempfile::tempdir().unwrap();
let path = root.path().join("supplemental.dat");
std::fs::write(&path, b"stale overlay").unwrap();
remove_disabled_supplemental(&path).unwrap();
assert!(!path.exists());
remove_disabled_supplemental(&path).unwrap();
let directory = root.path().join("not-a-file");
std::fs::create_dir(&directory).unwrap();
let error = remove_disabled_supplemental(&directory).unwrap_err();
assert!(error.contains("failed to remove"), "{error}");
assert!(directory.is_dir());
}
#[test]
fn supplemental_aggregate_limit_is_atomic() {
let mut supplemental = SupplementalEntries::default();
let first = tirith_core::threatdb_feeds::FeedEntries {
hostnames: vec!["one.example".to_string()],
ips: vec![],
};
assert_eq!(
supplemental
.ingest_with_limit(first, ThreatSource::Urlhaus, 1)
.unwrap(),
1
);
let second = tirith_core::threatdb_feeds::FeedEntries {
hostnames: vec!["two.example".to_string()],
ips: vec![],
};
let error = supplemental
.ingest_with_limit(second, ThreatSource::PhishingArmy, 1)
.unwrap_err();
assert!(error.contains("aggregate indicator limit"), "{error}");
assert_eq!(supplemental.hostnames.len(), 1);
assert_eq!(supplemental.hostnames[0].0, "one.example");
}
#[test]
fn legacy_invalid_primary_uses_valid_signed_fallback() {
let key = SigningKey::from_bytes(&[0x31; 32]);
let fallback = signed_manifest(
12,
"https://example.com/fallback.dat",
&"b".repeat(64),
&key,
);
let mut invalid_primary =
signed_manifest(13, "https://example.com/primary.dat", &"a".repeat(64), &key);
invalid_primary.version = 14;
let selected = fetch_manifest_with(
|url| {
if url == MANIFEST_URL_PRIMARY {
Ok(invalid_primary.clone())
} else {
Ok(fallback.clone())
}
},
&key.verifying_key(),
)
.expect("valid fallback must survive an unauthenticated primary");
assert_eq!(selected.version, 12);
assert!(selected
.verify_signature_with_key(&key.verifying_key())
.is_ok());
}
#[test]
fn legacy_selection_chooses_newest_only_after_both_verify() {
let key = SigningKey::from_bytes(&[0x32; 32]);
let primary = signed_manifest(11, "https://example.com/primary.dat", &"a".repeat(64), &key);
let fallback = signed_manifest(
12,
"https://example.com/fallback.dat",
&"b".repeat(64),
&key,
);
let selected = fetch_manifest_with(
|url| {
if url == MANIFEST_URL_PRIMARY {
Ok(primary.clone())
} else {
Ok(fallback.clone())
}
},
&key.verifying_key(),
)
.unwrap();
assert_eq!(selected.version, 12);
assert_eq!(selected.url, fallback.url);
}
#[test]
fn legacy_equal_version_signed_equivocation_fails_closed() {
let key = SigningKey::from_bytes(&[0x33; 32]);
let primary = signed_manifest(12, "https://example.com/primary.dat", &"a".repeat(64), &key);
let fallback = signed_manifest(
12,
"https://example.com/fallback.dat",
&"b".repeat(64),
&key,
);
let error = fetch_manifest_with(
|url| {
if url == MANIFEST_URL_PRIMARY {
Ok(primary.clone())
} else {
Ok(fallback.clone())
}
},
&key.verifying_key(),
)
.unwrap_err();
assert!(error.contains("equivocation"), "{error}");
}
#[test]
fn generation_index_requires_one_matching_asset_per_format() {
let key = SigningKey::from_bytes(&[2u8; 32]);
let valid = signed_index_v2(1, vec![asset(1, None), asset(2, Some("0.3.4"))], &key);
assert!(valid.validate_generation().is_ok());
let partial = signed_index_v2(1, vec![asset(2, Some("0.3.4"))], &key);
assert!(partial.validate_generation().is_err());
let mut mismatched = asset(2, Some("0.3.4"));
mismatched.filename = "different.dat".to_string();
let mismatched = signed_index_v2(1, vec![asset(1, None), mismatched], &key);
assert!(mismatched.validate_generation().is_err());
let mut duplicate_name_v1 = asset(1, None);
duplicate_name_v1.filename = "shared.dat".to_string();
duplicate_name_v1.url = "https://example.com/shared.dat?format=1".to_string();
let mut duplicate_name_v2 = asset(2, Some("0.3.4"));
duplicate_name_v2.filename = "shared.dat".to_string();
duplicate_name_v2.url = "https://example.com/shared.dat?format=2".to_string();
let duplicate_names = signed_index_v2(1, vec![duplicate_name_v1, duplicate_name_v2], &key);
assert!(duplicate_names
.validate_generation()
.unwrap_err()
.contains("distinct filenames"));
let mut duplicate_url_v2 = asset(2, Some("0.3.4"));
duplicate_url_v2.url = asset(1, None).url;
let duplicate_urls = signed_index_v2(1, vec![asset(1, None), duplicate_url_v2], &key);
assert!(duplicate_urls
.validate_generation()
.unwrap_err()
.contains("distinct URLs"));
}
#[test]
fn index_v2_canonical_payload_is_sorted_and_excludes_signature() {
let key = SigningKey::from_bytes(&[3u8; 32]);
let idx = signed_index_v2(42, vec![asset(2, Some("0.3.4")), asset(1, None)], &key);
let payload = idx.canonical_payload();
assert!(!payload.contains(' '));
assert!(!payload.contains("signature"));
let pos_assets = payload.find("\"assets\"").unwrap();
let pos_version = payload.find("\"manifest_version\"").unwrap();
let pos_sequence = payload.find("\"sequence\"").unwrap();
assert!(
pos_assets < pos_version && pos_version < pos_sequence,
"top-level keys are alphabetical"
);
let a = payload.find("\"filename\"").unwrap();
let b = payload.find("\"format\"").unwrap();
let c = payload.find("\"sha256\"").unwrap();
assert!(a < b && b < c, "asset keys must be alphabetical");
}
#[test]
fn index_asset_filename_is_in_signed_canonical_payload() {
let key = SigningKey::from_bytes(&[4u8; 32]);
let mut a = asset(2, None);
a.filename = "tirith-threatdb-known-name.dat".to_string();
let idx = signed_index_v2(1, vec![a], &key);
let payload = idx.canonical_payload();
assert!(
payload.contains(r#""filename":"tirith-threatdb-known-name.dat""#),
"filename must appear in the signed canonical payload: {payload}"
);
}
#[test]
fn canonical_payload_keys_sorted_and_signature_excluded() {
let key = SigningKey::from_bytes(&[7u8; 32]);
let idx = signed_index_v2(42, vec![asset(2, Some("0.3.4")), asset(1, None)], &key);
let assets: Vec<serde_json::Value> = idx
.assets
.iter()
.map(|a| {
let mut m = serde_json::json!({
"url": a.url,
"size": a.size,
"sha256": a.sha256,
"format": a.format,
"filename": a.filename,
});
if let Some(ref v) = a.min_tirith_version {
m.as_object_mut()
.unwrap()
.insert("min_tirith_version".to_string(), serde_json::json!(v));
}
m
})
.collect();
let value = serde_json::json!({
"sequence": idx.sequence,
"manifest_version": idx.manifest_version,
"assets": assets,
});
let sorted = sort_json_keys(&value);
let independent = serde_json::to_string(&sorted).unwrap();
assert_eq!(
idx.canonical_payload(),
independent,
"canonical_payload must equal an independently sorted, signature-free serialization"
);
assert_json_object_keys_sorted(&sorted);
assert!(!independent.contains("signature"), "signature excluded");
let reparsed: serde_json::Value = serde_json::from_str(&idx.canonical_payload()).unwrap();
assert_eq!(reparsed["sequence"], serde_json::json!(idx.sequence));
assert_eq!(
reparsed["manifest_version"],
serde_json::json!(idx.manifest_version)
);
assert_eq!(
reparsed["assets"].as_array().unwrap().len(),
idx.assets.len()
);
}
fn sort_json_keys(v: &serde_json::Value) -> serde_json::Value {
match v {
serde_json::Value::Object(map) => {
let mut sorted = serde_json::Map::new();
let mut keys: Vec<&String> = map.keys().collect();
keys.sort();
for k in keys {
sorted.insert(k.clone(), sort_json_keys(&map[k]));
}
serde_json::Value::Object(sorted)
}
serde_json::Value::Array(arr) => {
serde_json::Value::Array(arr.iter().map(sort_json_keys).collect())
}
other => other.clone(),
}
}
fn assert_json_object_keys_sorted(v: &serde_json::Value) {
match v {
serde_json::Value::Object(map) => {
let keys: Vec<&String> = map.keys().collect();
let mut expected = keys.clone();
expected.sort();
assert_eq!(keys, expected, "object keys must be sorted: {map:?}");
for val in map.values() {
assert_json_object_keys_sorted(val);
}
}
serde_json::Value::Array(arr) => {
for val in arr {
assert_json_object_keys_sorted(val);
}
}
_ => {}
}
}
#[test]
fn index_v2_signature_roundtrips_against_signer() {
let key = SigningKey::from_bytes(&[9u8; 32]);
let idx = signed_index_v2(7, vec![asset(2, None)], &key);
let sig_bytes =
base64::Engine::decode(&base64::engine::general_purpose::STANDARD, &idx.signature)
.unwrap();
let signature = Signature::from_slice(&sig_bytes).unwrap();
use ed25519_dalek::Verifier;
assert!(key
.verifying_key()
.verify(idx.canonical_payload().as_bytes(), &signature)
.is_ok());
let mut tampered = idx.clone();
tampered.sequence = 8;
assert!(key
.verifying_key()
.verify(tampered.canonical_payload().as_bytes(), &signature)
.is_err());
let mut tampered = idx.clone();
tampered.manifest_version += 1;
assert!(key
.verifying_key()
.verify(tampered.canonical_payload().as_bytes(), &signature)
.is_err());
}
#[test]
fn verify_signature_rejects_unknown_manifest_version() {
let key = SigningKey::from_bytes(&[6u8; 32]);
let mut idx = signed_index_v2(7, vec![asset(1, None), asset(2, None)], &key);
idx.manifest_version = SIGNED_MANIFEST_VERSION + 1;
use ed25519_dalek::Signer;
idx.signature = base64::Engine::encode(
&base64::engine::general_purpose::STANDARD,
key.sign(idx.canonical_payload().as_bytes()).to_bytes(),
);
let err = idx
.verify_signature_with_key(&key.verifying_key())
.expect_err("an unknown manifest_version must be rejected");
assert!(
err.contains("manifest_version"),
"error must name manifest_version, got: {err}"
);
let mut tampered = signed_index_v2(7, vec![asset(1, None), asset(2, None)], &key);
tampered.manifest_version += 1;
let err = tampered
.verify_signature_with_key(&key.verifying_key())
.unwrap_err();
assert!(
err.contains("signature verification failed"),
"unsigned version mutation must fail authenticity first, got: {err}"
);
}
#[test]
fn index_v2_verify_signature_rejects_garbage() {
let key = SigningKey::from_bytes(&[13u8; 32]);
let idx = IndexV2 {
manifest_version: SIGNED_MANIFEST_VERSION,
sequence: 1,
assets: vec![asset(1, None)],
signature: "not-base64-or-too-short".to_string(),
};
assert!(idx.verify_signature_with_key(&key.verifying_key()).is_err());
}
#[test]
fn invalid_primary_index_uses_valid_release_fallback() {
let key = SigningKey::from_bytes(&[12u8; 32]);
let fallback = signed_index_v2(9, vec![asset(1, None), asset(2, None)], &key);
let mut primary = fallback.clone();
primary.sequence = 8;
let selected = fetch_index_v2_with(
|url| {
if url == INDEX_V2_URL_PRIMARY {
Ok(primary.clone())
} else if url == INDEX_V2_URL_FALLBACK {
Ok(fallback.clone())
} else {
Err(format!("unexpected URL: {url}"))
}
},
&key.verifying_key(),
)
.expect("a valid release index must survive an invalid primary");
assert_eq!(selected.sequence, 9);
assert!(selected
.verify_signature_with_key(&key.verifying_key())
.is_ok());
}
#[test]
fn fresh_client_chooses_newest_verified_discovery_surface() {
let key = SigningKey::from_bytes(&[14u8; 32]);
let primary = signed_index_v2(8, vec![asset(1, None), asset(2, None)], &key);
let fallback = signed_index_v2(9, vec![asset(1, None), asset(2, None)], &key);
let selected = fetch_index_v2_with(
|url| {
if url == INDEX_V2_URL_PRIMARY {
Ok(primary.clone())
} else {
Ok(fallback.clone())
}
},
&key.verifying_key(),
)
.expect("a replayed older primary must not outrank the release pointer");
assert_eq!(selected.sequence, 9);
}
#[test]
fn equal_sequence_discovery_equivocation_fails_closed() {
let key = SigningKey::from_bytes(&[15u8; 32]);
let primary = signed_index_v2(9, vec![asset(1, None), asset(2, None)], &key);
let mut alternate_v2 = asset(2, None);
alternate_v2.filename = "tirith-threatdb-v2-alternate.dat".to_string();
alternate_v2.url = "https://example.com/tirith-threatdb-v2-alternate.dat".to_string();
let fallback = signed_index_v2(9, vec![asset(1, None), alternate_v2], &key);
let error = fetch_index_v2_with(
|url| {
if url == INDEX_V2_URL_PRIMARY {
Ok(primary.clone())
} else {
Ok(fallback.clone())
}
},
&key.verifying_key(),
)
.expect_err("one sequence cannot identify two signed generations");
assert!(error.contains("equivocation"), "{error}");
}
#[test]
fn index_v2_select_prefers_highest_compatible_format() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let idx = signed_index_v2(1, vec![asset(1, None), asset(2, Some("0.3.4"))], &key);
let chosen = idx.select_asset("0.3.4").expect("an asset is compatible");
assert_eq!(chosen.format, 2, "highest compatible format wins");
}
#[test]
fn post_r3_signed_index_selection_contract_is_frozen() {
let key = SigningKey::from_bytes(&[0xc0; 32]);
let idx = signed_index_v2(181, vec![asset(2, Some("0.3.4")), asset(1, None)], &key);
for (client, expected_format, expected_filename) in [
("0.3.3", 1, "tirith-threatdb-v1.dat"),
("0.3.4", 2, "tirith-threatdb-v2.dat"),
("0.4.0", 2, "tirith-threatdb-v2.dat"),
] {
let selected = idx
.select_asset(client)
.unwrap_or_else(|| panic!("post-r3 client {client} must select an asset"));
assert_eq!(selected.format, expected_format, "client {client}");
assert_eq!(selected.filename, expected_filename, "client {client}");
}
let reversed = signed_index_v2(181, vec![asset(1, None), asset(2, Some("0.3.4"))], &key);
assert_eq!(
reversed
.select_asset("0.3.4")
.map(|asset| (asset.format, asset.filename.as_str())),
Some((2, "tirith-threatdb-v2.dat")),
"signed-index asset order must not affect the selected channel"
);
}
#[test]
fn index_v2_select_skips_format_above_ceiling() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let idx = signed_index_v2(1, vec![asset(1, None), asset(99, None)], &key);
let chosen = idx.select_asset("9.9.9").expect("v1 still compatible");
assert_eq!(chosen.format, 1);
assert!(chosen.format <= MAX_FORMAT_VERSION);
}
#[test]
fn index_v2_select_honors_min_tirith_version() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let idx = signed_index_v2(1, vec![asset(1, None), asset(2, Some("0.4.0"))], &key);
let chosen = idx.select_asset("0.3.3").expect("v1 compatible");
assert_eq!(
chosen.format, 1,
"too-old client must not pick the v2 asset"
);
assert_eq!(idx.select_asset("0.4.0").unwrap().format, 2);
}
#[test]
fn index_v2_select_none_when_nothing_compatible() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let idx = signed_index_v2(1, vec![asset(99, None)], &key);
assert!(idx.select_asset("0.3.3").is_none());
}
#[test]
fn index_v2_select_skips_oversized_asset() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let mut huge = asset(2, None);
huge.size = MAX_INDEX_ASSET_SIZE + 1;
let idx = signed_index_v2(1, vec![asset(1, None), huge], &key);
let chosen = idx.select_asset("9.9.9").expect("v1 compatible");
assert_eq!(chosen.format, 1, "oversized v2 asset is skipped");
}
#[test]
fn index_v2_select_unparseable_min_version_is_incompatible() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let idx = signed_index_v2(
1,
vec![asset(1, None), asset(2, Some("not.a.version"))],
&key,
);
let chosen = idx.select_asset("0.3.3").expect("v1 compatible");
assert_eq!(chosen.format, 1);
}
#[test]
fn select_asset_rejects_duplicate_format() {
let key = SigningKey::from_bytes(&[1u8; 32]);
let mut second = asset(2, None);
second.url = "https://example.com/db-v2-alt.dat".to_string();
let idx = signed_index_v2(1, vec![asset(2, None), second], &key);
assert!(
idx.select_asset("9.9.9").is_none(),
"an ambiguous index with two top-format assets must select nothing"
);
let idx = signed_index_v2(
1,
vec![asset(1, None), asset(1, None), asset(2, None)],
&key,
);
assert_eq!(
idx.select_asset("9.9.9").map(|a| a.format),
Some(2),
"a duplicate below the top format must not block the unambiguous top"
);
}
#[test]
fn primary_db_dest_routes_v2_to_distinct_path_never_v1() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let v1_path = tmp.path().join("tirith-threatdb.dat");
let _path_guard = EnvGuard::set("TIRITH_THREATDB_PATH", &v1_path);
let v1_dest = primary_db_dest(1).unwrap();
let v2_dest = primary_db_dest(2).unwrap();
assert_eq!(v1_dest, v1_path);
assert_eq!(v2_dest, tmp.path().join("tirith-threatdb-v2.dat"));
assert_ne!(v1_dest, v2_dest, "v2 must never resolve to the v1 path");
}
#[test]
fn legacy_equal_sequence_v2_requires_install_and_retirement() {
assert!(legacy_install_needed(8, Some((8, 2)), false).unwrap());
assert!(!legacy_install_needed(8, Some((8, 1)), false).unwrap());
assert!(legacy_install_needed(9, Some((8, 2)), false).unwrap());
assert!(legacy_install_needed(7, Some((8, 2)), false).is_err());
assert!(legacy_install_needed(7, Some((8, 2)), true).unwrap());
}
#[test]
fn index_equal_sequence_format_changes_are_not_already_current() {
assert!(index_install_needed(8, 2, Some((8, 1)), false).unwrap());
assert!(index_install_needed(8, 1, Some((8, 2)), false).unwrap());
assert!(!index_install_needed(8, 2, Some((8, 2)), false).unwrap());
assert!(!index_install_needed(8, 1, Some((8, 1)), false).unwrap());
assert!(index_install_needed(9, 2, Some((8, 2)), false).unwrap());
assert!(index_install_needed(7, 2, Some((8, 2)), false).is_err());
}
#[test]
fn retiring_v2_cache_preserves_v1_and_is_idempotent() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let v1_path = tmp.path().join("tirith-threatdb.dat");
let v2_path = tmp.path().join("tirith-threatdb-v2.dat");
std::fs::write(&v1_path, b"verified-v1-placeholder").unwrap();
std::fs::write(&v2_path, b"stale-v2-placeholder").unwrap();
let _path_guard = EnvGuard::set("TIRITH_THREATDB_PATH", &v1_path);
retire_primary_v2().expect("v2 retirement succeeds");
assert_eq!(std::fs::read(&v1_path).unwrap(), b"verified-v1-placeholder");
assert!(!v2_path.exists());
retire_primary_v2().expect("already-retired v2 is a successful no-op");
}
#[test]
fn build_format_stamps_distinct_version_per_format() {
let key = SigningKey::from_bytes(&[5u8; 32]);
let v1 = ThreatDbWriter::new(1, 1)
.build_format(ThreatDbFormat::V1, &key)
.unwrap();
let v2 = ThreatDbWriter::new(2, 2)
.build_format(ThreatDbFormat::V2, &key)
.unwrap();
assert_eq!(u32::from_le_bytes(v1[8..12].try_into().unwrap()), 1);
assert_eq!(u32::from_le_bytes(v2[8..12].try_into().unwrap()), 2);
}
const WORKFLOW_V2_INDEX_PAYLOAD_WITH_MIN: &str = concat!(
"{\"assets\":[",
"{\"filename\":\"tirith-threatdb-7-1.dat\",\"format\":1,",
"\"sha256\":\"1111111111111111111111111111111111111111111111111111111111111111\",",
"\"size\":4096,",
"\"url\":\"https://github.com/sheeki03/tirith/releases/download/threatdb-latest/tirith-threatdb-7-1.dat\"},",
"{\"filename\":\"tirith-threatdb-v2-7-1.dat\",\"format\":2,",
"\"min_tirith_version\":\"0.3.4\",",
"\"sha256\":\"2222222222222222222222222222222222222222222222222222222222222222\",",
"\"size\":8192,",
"\"url\":\"https://github.com/sheeki03/tirith/releases/download/threatdb-latest/tirith-threatdb-v2-7-1.dat\"}",
"],\"manifest_version\":2,\"sequence\":7}"
);
const WORKFLOW_V2_INDEX_PAYLOAD_NO_MIN: &str = concat!(
"{\"assets\":[",
"{\"filename\":\"tirith-threatdb-7-1.dat\",\"format\":1,",
"\"sha256\":\"1111111111111111111111111111111111111111111111111111111111111111\",",
"\"size\":4096,",
"\"url\":\"https://github.com/sheeki03/tirith/releases/download/threatdb-latest/tirith-threatdb-7-1.dat\"},",
"{\"filename\":\"tirith-threatdb-v2-7-1.dat\",\"format\":2,",
"\"sha256\":\"2222222222222222222222222222222222222222222222222222222222222222\",",
"\"size\":8192,",
"\"url\":\"https://github.com/sheeki03/tirith/releases/download/threatdb-latest/tirith-threatdb-v2-7-1.dat\"}",
"],\"manifest_version\":2,\"sequence\":7}"
);
fn published_index_json(canonical_payload: &str, signature: &str) -> String {
let mut value: serde_json::Value = serde_json::from_str(canonical_payload).unwrap();
let obj = value.as_object_mut().unwrap();
obj.insert(
"signature".to_string(),
serde_json::Value::String(signature.to_string()),
);
value.to_string()
}
#[test]
fn workflow_v2_index_payload_matches_client_canonical_with_min() {
let published = published_index_json(WORKFLOW_V2_INDEX_PAYLOAD_WITH_MIN, "AA==");
let index: IndexV2 = serde_json::from_str(&published).unwrap();
assert_eq!(
index.canonical_payload(),
WORKFLOW_V2_INDEX_PAYLOAD_WITH_MIN,
"client canonical_payload() must be byte-identical to the workflow's jq -cS output"
);
assert!(index.validate_generation().is_ok());
let key = SigningKey::from_bytes(&[7u8; 32]);
use ed25519_dalek::{Signer, Verifier};
let sig = key.sign(index.canonical_payload().as_bytes());
assert!(key
.verifying_key()
.verify(index.canonical_payload().as_bytes(), &sig)
.is_ok());
let mut tampered = WORKFLOW_V2_INDEX_PAYLOAD_WITH_MIN.as_bytes().to_vec();
tampered[0] ^= 0x01;
assert!(
key.verifying_key().verify(&tampered, &sig).is_err(),
"signature must not verify over a mutated payload"
);
}
#[test]
fn workflow_v2_index_payload_matches_client_canonical_no_min() {
let published = published_index_json(WORKFLOW_V2_INDEX_PAYLOAD_NO_MIN, "AA==");
let index: IndexV2 = serde_json::from_str(&published).unwrap();
assert_eq!(index.assets.len(), 2);
assert!(index.assets[1].min_tirith_version.is_none());
assert!(index.validate_generation().is_ok());
assert_eq!(
index.canonical_payload(),
WORKFLOW_V2_INDEX_PAYLOAD_NO_MIN,
"absent min_tirith_version must yield the same byte shape on both sides"
);
let key = SigningKey::from_bytes(&[8u8; 32]);
use ed25519_dalek::{Signer, Verifier};
let sig = key.sign(index.canonical_payload().as_bytes());
assert!(key
.verifying_key()
.verify(index.canonical_payload().as_bytes(), &sig)
.is_ok());
}
#[test]
fn canonical_payload_large_sequence_round_trips_losslessly() {
let big: u64 = 9_007_199_254_740_993; let key = SigningKey::from_bytes(&[11u8; 32]);
let idx = signed_index_v2(big, vec![asset(2, None)], &key);
let payload = idx.canonical_payload();
assert!(
payload.contains(&format!("\"sequence\":{big}")),
"sequence must serialize as the exact u64 literal, got: {payload}"
);
let reparsed: serde_json::Value = serde_json::from_str(&payload).unwrap();
assert_eq!(
reparsed["sequence"].as_u64(),
Some(big),
"sequence must round-trip losslessly as u64"
);
let published = published_index_json(&payload, &idx.signature);
let parsed: IndexV2 = serde_json::from_str(&published).unwrap();
assert_eq!(parsed.sequence, big, "wire round-trip must preserve u64");
}
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 = ENV_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();
let _policy_guard = EnvGuard::set("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"
);
}
#[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 checked_in_legacy_manifest_has_valid_pinned_signature() {
let manifest: Manifest =
serde_json::from_str(include_str!("../../../../threatdb-manifest.json"))
.expect("checked-in legacy manifest must be valid JSON");
manifest
.verify_signature()
.expect("checked-in legacy manifest must verify with the pinned public key");
}
#[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 = super::MANIFEST_URL_PRIMARY;
let fallback = super::MANIFEST_URL_FALLBACK;
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> {
let client = reqwest::blocking::Client::builder()
.no_proxy()
.timeout(std::time::Duration::from_secs(5))
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test HTTP client");
super::fetch_manifest_from_with_state_and_client(url, Some(state.to_path_buf()), &client)
}
#[test]
fn production_fetch_paths_reject_private_initial_destinations() {
let private = "https://127.0.0.1/threatdb";
for error in [
super::fetch_manifest_from_with_state(private, None).unwrap_err(),
super::download_url(private, 1).unwrap_err(),
super::fetch_index_v2_from(private).unwrap_err(),
] {
assert!(
error.contains("refusing unsafe"),
"unexpected error: {error}"
);
}
let client = super::guarded_http_client(1).expect("guarded client builds");
let error = super::fetch_bytes(&client, private).unwrap_err();
assert!(error.contains("refusing unsafe supplemental feed URL"));
}
#[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"));
}
#[test]
fn offline_env_active_recognizes_truthy_values() {
let mut guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
for v in ["1", "true", "TRUE", "yes", "On", " on "] {
guard.set_env("TIRITH_OFFLINE", v);
assert!(
crate::cli::offline_env_active(),
"TIRITH_OFFLINE={v:?} should be treated as offline"
);
}
}
#[test]
fn offline_env_active_rejects_falsey_and_unset() {
let mut guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
for v in ["0", "false", "no", "", "off", "garbage"] {
guard.set_env("TIRITH_OFFLINE", v);
assert!(
!crate::cli::offline_env_active(),
"TIRITH_OFFLINE={v:?} should NOT be treated as offline"
);
}
guard.remove_env("TIRITH_OFFLINE");
assert!(
!crate::cli::offline_env_active(),
"unset TIRITH_OFFLINE should not be offline"
);
}
#[test]
fn offline_flag_skips_background_update_no_network_attempt() {
let mut guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
guard.remove_env("TIRITH_OFFLINE");
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
super::maybe_background_update(true);
let spawned_at = tmp.path().join("tirith").join(SPAWNED_AT_FILE);
assert!(
!spawned_at.exists(),
"--offline must skip the background update before any state write"
);
}
#[test]
fn offline_env_skips_background_update_no_network_attempt() {
let mut guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
guard.set_env("TIRITH_OFFLINE", "1");
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
super::maybe_background_update(false);
let spawned_at = tmp.path().join("tirith").join(SPAWNED_AT_FILE);
assert!(
!spawned_at.exists(),
"TIRITH_OFFLINE=1 must skip the background update before any state write"
);
}
#[test]
fn offline_short_circuits_before_update_attempted_latch() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let latch = AtomicBool::new(false);
let offline = true;
if !offline {
latch.swap(true, Ordering::Relaxed);
}
assert!(
!latch.load(Ordering::Relaxed),
"an offline call must not consume the once-per-process latch"
);
}
#[test]
fn parse_indicator_recognizes_ipv4() {
let p = parse_indicator("203.0.113.50");
assert_eq!(p.kind, IndicatorKind::Ip);
assert_eq!(p.value, "203.0.113.50");
assert!(p.ecosystem.is_none());
assert!(p.version.is_none());
}
#[test]
fn parse_indicator_recognizes_ecosystem_prefix() {
let p = parse_indicator("npm:left-pad");
assert_eq!(p.kind, IndicatorKind::Package);
assert_eq!(p.ecosystem, Some(Ecosystem::Npm));
assert_eq!(p.value, "left-pad");
assert!(p.version.is_none());
}
#[test]
fn parse_indicator_recognizes_ecosystem_prefix_with_version() {
let p = parse_indicator("pypi:requests@2.0.0");
assert_eq!(p.kind, IndicatorKind::Package);
assert_eq!(p.ecosystem, Some(Ecosystem::PyPI));
assert_eq!(p.value, "requests");
assert_eq!(p.version.as_deref(), Some("2.0.0"));
}
#[test]
fn parse_indicator_host_colon_port_is_not_a_package() {
let p = parse_indicator("example.com:8080");
assert_eq!(p.kind, IndicatorKind::Domain);
}
#[test]
fn parse_indicator_recognizes_name_at_version() {
let p = parse_indicator("lodash@4.17.21");
assert_eq!(p.kind, IndicatorKind::Package);
assert!(p.ecosystem.is_none());
assert_eq!(p.value, "lodash");
assert_eq!(p.version.as_deref(), Some("4.17.21"));
}
#[test]
fn parse_indicator_scoped_npm_package_is_not_split_on_leading_at() {
let p = parse_indicator("@angular/core");
assert_eq!(p.kind, IndicatorKind::Package);
assert_eq!(p.value, "@angular/core");
assert!(p.version.is_none());
}
#[test]
fn parse_indicator_scoped_npm_package_with_version() {
let p = parse_indicator("@angular/core@17.0.0");
assert_eq!(p.kind, IndicatorKind::Package);
assert_eq!(p.value, "@angular/core");
assert_eq!(p.version.as_deref(), Some("17.0.0"));
}
#[test]
fn parse_indicator_dotted_token_is_domain() {
let p = parse_indicator("evil.example.com");
assert_eq!(p.kind, IndicatorKind::Domain);
assert_eq!(p.value, "evil.example.com");
}
#[test]
fn parse_indicator_domain_is_lowercased() {
let p = parse_indicator("EVIL.Example.COM");
assert_eq!(p.kind, IndicatorKind::Domain);
assert_eq!(p.value, "evil.example.com");
}
#[test]
fn parse_indicator_bare_name_is_package() {
let p = parse_indicator("react");
assert_eq!(p.kind, IndicatorKind::Package);
assert_eq!(p.value, "react");
}
#[test]
fn split_at_version_rejects_missing_parts() {
assert!(split_at_version("react").is_none());
assert!(split_at_version("react@").is_none());
assert!(split_at_version("@1.0.0").is_none());
assert_eq!(
split_at_version("react@1.0.0"),
Some(("react".to_string(), "1.0.0".to_string()))
);
}
#[test]
fn parse_since_accepts_version_number() {
let (kind, version, epoch) = parse_since("42").unwrap();
assert_eq!(kind, "version");
assert_eq!(version, Some(42));
assert_eq!(epoch, None);
}
#[test]
fn parse_since_accepts_iso_date() {
let (kind, version, epoch) = parse_since("2026-01-15").unwrap();
assert_eq!(kind, "date");
assert_eq!(version, None);
assert_eq!(epoch, Some(1768435200));
}
#[test]
fn parse_since_rejects_garbage() {
assert!(parse_since("not-a-date").is_err());
assert!(parse_since("2026-13-01").is_err());
assert!(parse_since("2026-01-99").is_err());
}
#[test]
fn parse_iso_date_epoch_zero_is_unix_epoch() {
assert_eq!(parse_iso_date("1970-01-01"), Some(0));
}
#[test]
fn parse_iso_date_handles_leap_year() {
let feb29 = parse_iso_date("2024-02-29").unwrap();
let mar01 = parse_iso_date("2024-03-01").unwrap();
assert_eq!(mar01 - feb29, 86400);
}
#[test]
fn parse_iso_date_rejects_day_past_month_length() {
assert_eq!(parse_iso_date("2026-02-30"), None);
assert_eq!(parse_iso_date("2026-04-31"), None);
assert_eq!(parse_iso_date("2026-06-31"), None);
assert_eq!(parse_iso_date("2025-02-29"), None);
assert!(parse_iso_date("2026-02-28").is_some());
assert!(parse_iso_date("2026-04-30").is_some());
assert!(parse_iso_date("2026-01-31").is_some());
assert!(parse_since("2026-02-30").is_err());
}
#[test]
fn parse_iso_date_accepts_datetime_suffix() {
assert_eq!(
parse_iso_date("2026-01-15T12:30:00"),
parse_iso_date("2026-01-15")
);
}
#[test]
fn format_epoch_round_trips_with_parse_iso_date() {
let epoch = parse_iso_date("2026-05-21").unwrap();
assert!(format_epoch(epoch).starts_with("2026-05-21 00:00:00"));
}
#[test]
fn format_epoch_known_timestamp() {
assert_eq!(format_epoch(1700000000), "2023-11-14 22:13:20 UTC");
}
#[test]
fn delta_of_computes_signed_category_changes() {
let baseline = CategoryCounts {
packages: 10,
hostnames: 5,
ips: 3,
typosquats: 2,
popular: 100,
};
let current = CategoryCounts {
packages: 12,
hostnames: 5,
ips: 1,
typosquats: 4,
popular: 100,
};
let d = delta_of(¤t, &baseline);
assert_eq!(d.packages, 2);
assert_eq!(d.hostnames, 0);
assert_eq!(d.ips, -2);
assert_eq!(d.typosquats, 2);
assert_eq!(d.popular, 0);
assert_eq!(d.total, 2);
}
#[test]
fn category_counts_total_sums_all_sections() {
let c = CategoryCounts {
packages: 1,
hostnames: 2,
ips: 4,
typosquats: 8,
popular: 16,
};
assert_eq!(c.total(), 31);
}
#[test]
fn record_snapshot_dedups_on_build_sequence() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
let snap = |seq: u64, recorded: u64| DbSnapshot {
recorded_at: recorded,
build_sequence: seq,
build_timestamp: 1_700_000_000,
signature_valid: true,
counts: CategoryCounts::default(),
sources: Default::default(),
};
record_snapshot(&snap(42, 1000));
record_snapshot(&snap(42, 2000));
record_snapshot(&snap(43, 3000));
let (history, _) = load_history();
assert_eq!(
history.len(),
2,
"duplicate build_sequence should be skipped"
);
assert_eq!(history[0].build_sequence, 42);
assert_eq!(history[1].build_sequence, 43);
}
#[test]
fn record_snapshot_caps_history_length() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
for seq in 0..(HISTORY_MAX_LINES as u64 + 20) {
record_snapshot(&DbSnapshot {
recorded_at: 1000 + seq,
build_sequence: seq,
build_timestamp: 1_700_000_000,
signature_valid: true,
counts: CategoryCounts::default(),
sources: Default::default(),
});
}
let (history, _) = load_history();
assert_eq!(
history.len(),
HISTORY_MAX_LINES,
"history must be capped at HISTORY_MAX_LINES"
);
assert_eq!(
history.last().unwrap().build_sequence,
HISTORY_MAX_LINES as u64 + 19
);
}
#[test]
fn load_history_skips_corrupt_lines() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
let state = tmp.path().join("tirith");
std::fs::create_dir_all(&state).unwrap();
let valid = r#"{"recorded_at":1000,"build_sequence":1,"build_timestamp":1700000000,"signature_valid":true,"counts":{"packages":0,"hostnames":0,"ips":0,"typosquats":0,"popular":0},"sources":{}}"#;
std::fs::write(
state.join(HISTORY_FILE),
format!("not json\n{valid}\n\nalso not json\n"),
)
.unwrap();
let (history, _) = load_history();
assert_eq!(history.len(), 1, "only the one valid line should parse");
assert_eq!(history[0].build_sequence, 1);
}
#[test]
fn load_history_missing_file_is_not_an_error() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
let _appdata_guard = EnvGuard::set("APPDATA", tmp.path());
let (history, read_error) = load_history();
assert!(history.is_empty());
assert!(
read_error.is_none(),
"a missing history file must not surface as a read error"
);
}
#[test]
fn load_history_unreadable_file_surfaces_a_read_error() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::tempdir().unwrap();
let _state_guard = EnvGuard::set("XDG_STATE_HOME", tmp.path());
let _appdata_guard = EnvGuard::set("APPDATA", tmp.path());
let state = tmp.path().join("tirith");
std::fs::create_dir_all(state.join(HISTORY_FILE)).unwrap();
let (history, read_error) = load_history();
assert!(history.is_empty());
assert!(
read_error.is_some(),
"an existing-but-unreadable history file must surface a read error, \
not be silently treated as 'no snapshots'"
);
}
#[test]
fn confidence_as_str_covers_all_levels() {
assert_eq!(Confidence::Low.as_str(), "low");
assert_eq!(Confidence::Medium.as_str(), "medium");
assert_eq!(Confidence::Confirmed.as_str(), "confirmed");
}
}