use std::net::IpAddr;
use std::str::FromStr;
use crate::error::Result;
pub const DEFAULT_REQUEST_TIMEOUT_SECS: u64 = 30;
pub struct VerdictEngine;
impl VerdictEngine {
pub fn determine_threat_verdict(flagged_count: u8, risk_score: u8) -> String {
if flagged_count >= 2 {
"malicious".to_string()
} else if flagged_count == 1 || risk_score > 70 {
"suspicious".to_string()
} else {
"clean".to_string()
}
}
}
pub async fn resolve_hostname_to_ip(hostname: &str, timeout_secs: u64) -> Result<IpAddr> {
let dns_req = crate::api::DnsCheckRequest {
domain: hostname.to_string(),
record_types: vec!["A".to_string()],
timeout_secs,
..Default::default()
};
let results = crate::api::check_dns(&dns_req).await?;
if results.is_empty() || results[0].answers.is_empty() {
return Err(crate::error::ShoheError::DnsResolution(format!(
"No DNS records for {}",
hostname
)));
}
for record in &results[0].answers {
if let crate::api::RecordData::A(ip_str) = &record.data {
return Ok(std::net::IpAddr::from_str(ip_str)
.map_err(|_| crate::error::ShoheError::Parse(format!(
"Invalid IP address: {}",
ip_str
)))?);
}
}
Err(crate::error::ShoheError::DnsResolution(format!(
"No A records found for {}",
hostname
)))
}
pub fn validate_url_safety(url: &str) -> std::result::Result<(), String> {
let parsed = url::Url::parse(url).map_err(|e| format!("Invalid URL: {}", e))?;
match parsed.scheme() {
"http" | "https" => {}
scheme => return Err(format!("Disallowed URL scheme '{}' — only http/https allowed", scheme)),
}
let host = parsed.host_str().ok_or("URL has no host")?;
if let Ok(ip) = host.parse::<std::net::IpAddr>() {
if is_private_or_special_ip(&ip) {
return Err(format!("Host IP '{}' is a private/reserved address — blocked for security", ip));
}
} else {
let lower = host.to_lowercase();
if lower == "localhost"
|| lower.ends_with(".local")
|| lower.ends_with(".internal")
|| lower.ends_with(".intranet")
{
return Err(format!("Host '{}' appears to be an internal hostname — blocked for security", host));
}
}
Ok(())
}
fn is_private_or_special_ip(ip: &std::net::IpAddr) -> bool {
match ip {
std::net::IpAddr::V4(v4) => {
let o = v4.octets();
o[0] == 127
|| o[0] == 10
|| (o[0] == 172 && o[1] >= 16 && o[1] <= 31)
|| (o[0] == 192 && o[1] == 168)
|| (o[0] == 169 && o[1] == 254)
|| o[0] == 0
|| (o[0] == 100 && o[1] >= 64 && o[1] <= 127)
}
std::net::IpAddr::V6(v6) => {
v6.is_loopback()
|| v6.is_unspecified()
|| (v6.segments()[0] & 0xfe00 == 0xfc00)
|| (v6.segments()[0] & 0xffc0 == 0xfe80)
}
}
}
pub fn dns_request_for_record_type(
domain: String,
record_type: &str,
timeout_secs: u64,
) -> crate::api::DnsCheckRequest {
crate::api::DnsCheckRequest {
domain,
record_types: vec![record_type.to_string()],
timeout_secs,
..Default::default()
}
}