use serde::{Deserialize, Serialize};
use std::net::{SocketAddr, IpAddr};
use std::sync::Arc;
use rustls::pki_types::ServerName;
use tokio::net::TcpStream;
use tokio::time::{timeout, Duration};
use tokio_rustls::TlsConnector;
use crate::error::Result;
use crate::api::tls::SkipVerification;
pub async fn check_tls_vulns(req: &TlsVulnCheckRequest) -> Result<TlsVulnCheckResult> {
let hostname_str = &req.hostname;
let port = req.port;
let ip = match crate::api::helpers::resolve_hostname_to_ip(hostname_str, req.timeout_secs).await {
Ok(ip) => ip,
Err(e) => return Ok(TlsVulnCheckResult {
hostname: hostname_str.clone(),
port,
tls10_accepted: false,
tls11_accepted: false,
tls12_accepted: false,
tls13_accepted: false,
forward_secrecy: false,
error: Some(format!("DNS resolution failed: {}", e)),
}),
};
let (tls10_accepted, tls11_accepted, tls12_accepted, tls13_accepted) = tokio::join!(
test_tls_version(&ip, port, hostname_str, "1.0", req.timeout_secs),
test_tls_version(&ip, port, hostname_str, "1.1", req.timeout_secs),
test_tls_version(&ip, port, hostname_str, "1.2", req.timeout_secs),
test_tls_version(&ip, port, hostname_str, "1.3", req.timeout_secs),
);
let forward_secrecy = tls12_accepted || tls13_accepted;
Ok(TlsVulnCheckResult {
hostname: hostname_str.clone(),
port,
tls10_accepted,
tls11_accepted,
tls12_accepted,
tls13_accepted,
forward_secrecy,
error: None,
})
}
async fn test_tls_version(
ip: &IpAddr,
port: u16,
hostname: &str,
version: &str,
timeout_secs: u64,
) -> bool {
match timeout(
Duration::from_secs(timeout_secs),
connect_with_tls(ip, port, hostname),
)
.await
{
Ok(Ok((tls_ver, _))) => {
match version {
"1.3" => tls_ver.as_deref() == Some("TLSv1.3"),
"1.2" => {
matches!(tls_ver.as_deref(), Some("TLSv1.2") | Some("TLSv1.3"))
}
"1.1" => {
false
}
"1.0" => {
false
}
_ => false,
}
}
_ => false,
}
}
async fn connect_with_tls(
ip: &IpAddr,
port: u16,
hostname: &str,
) -> std::result::Result<(Option<String>, Option<String>), Box<dyn std::error::Error>> {
let addr = SocketAddr::new(*ip, port);
let socket = TcpStream::connect(&addr).await?;
let config = rustls::ClientConfig::builder_with_provider(
Arc::new(rustls::crypto::ring::default_provider()),
)
.with_safe_default_protocol_versions()
.map_err(|e| format!("TLS config error: {}", e))?
.dangerous()
.with_custom_certificate_verifier(Arc::new(SkipVerification))
.with_no_client_auth();
let connector = TlsConnector::from(Arc::new(config));
let server_name = ServerName::try_from(hostname.to_string())?;
let tls_stream = connector.connect(server_name, socket).await?;
let (_, conn) = tls_stream.get_ref();
let tls_version = conn.protocol_version().map(|v| match v {
rustls::ProtocolVersion::TLSv1_2 => "TLSv1.2".to_string(),
rustls::ProtocolVersion::TLSv1_3 => "TLSv1.3".to_string(),
_ => "Unknown".to_string(),
});
let cipher_suite = conn.negotiated_cipher_suite().map(|cs| format!("{:?}", cs));
Ok((tls_version, cipher_suite))
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TlsVulnCheckRequest {
pub hostname: String,
#[serde(default = "default_port")]
pub port: u16,
#[serde(default = "default_timeout")]
pub timeout_secs: u64,
}
fn default_port() -> u16 { 443 }
fn default_timeout() -> u64 { 10 }
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TlsVulnCheckResult {
pub hostname: String,
pub port: u16,
pub tls10_accepted: bool,
pub tls11_accepted: bool,
pub tls12_accepted: bool,
pub tls13_accepted: bool,
pub forward_secrecy: bool,
#[serde(default)]
pub error: Option<String>,
}