#[path = "../../tests/utils/certs.rs"]
mod certutils;
use std::{
fs,
io::Write,
sync::{Arc, LazyLock},
};
use anyhow::Result;
use camino::{Utf8Path, Utf8PathBuf};
use tempfile::{NamedTempFile, tempdir};
use test_log::test;
use tracing::info;
use crate::{
RunContext,
certificates::{
HostCertificate,
acme::to_txt_name,
host::load_hostcert,
store::CertStore,
tests::certutils::LocalCert,
watcher::{CertWatcher, RELOAD_GRACE},
},
config::Config,
errors::VicarianError,
};
use certutils::TEST_CERTS;
struct TestHostCerts {
pub snakeoil_1: HostCertificate,
pub snakeoil_2: HostCertificate,
pub www_example: HostCertificate,
pub wildcard_example: HostCertificate,
}
impl TestHostCerts {
fn new() -> Result<Self> {
let snakeoil_1 = from_localcert(&TEST_CERTS.snakeoil_1, true)?;
let snakeoil_2 = from_localcert(&TEST_CERTS.snakeoil_2, true)?;
let www_example = from_localcert(&TEST_CERTS.www_example, false)?;
let wildcard_example = from_localcert(&TEST_CERTS.wildcard_example, false)?;
Ok(Self {
snakeoil_1,
snakeoil_2,
www_example,
wildcard_example,
})
}
}
static TEST_HOST_CERTS: LazyLock<TestHostCerts> = LazyLock::new(|| {
rustls::crypto::aws_lc_rs::default_provider().install_default()
.expect("Failed to install Rustls crypto provider");
TestHostCerts::new().unwrap()
});
fn from_localcert(lc: &LocalCert, watch: bool) -> Result<HostCertificate> {
let hc = futures::executor::block_on(
HostCertificate::new(lc.keyfile.clone(),
lc.certfile.clone(),
watch))?;
Ok(hc)
}
#[tokio::test]
async fn test_load_certs_invalid_pair() -> Result<()> {
let so1 = TEST_HOST_CERTS.snakeoil_1.clone();
let so2 = TEST_HOST_CERTS.snakeoil_2.clone();
let key_path = so1.keyfile();
let other_cert_path = so2.certfile();
let result = HostCertificate::new(key_path.into(), other_cert_path.into(), false).await;
println!("ERR: {result:?}");
assert!(result.is_err());
let err: VicarianError = result.unwrap_err().downcast()?;
assert!(matches!(err, VicarianError::CertificateMismatch(_, _)));
Ok(())
}
#[tokio::test]
async fn test_load_certs_nonexistent_files() {
let key_path = Utf8Path::new("nonexistent.key");
let cert_path = Utf8Path::new("nonexistent.crt");
let result = load_hostcert(key_path, cert_path).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_load_certs_empty_cert_file() -> Result<()> {
let mut empty_cert_file = NamedTempFile::new()?;
empty_cert_file.write_all(b"")?;
let empty_cert_path = Utf8PathBuf::from(empty_cert_file.path().to_str().unwrap());
let so1 = TEST_HOST_CERTS.snakeoil_1.clone();
let result = load_hostcert(so1.keyfile(), &empty_cert_path).await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("No certificates found in TLS .crt file"));
Ok(())
}
#[tokio::test]
#[test_log::test]
async fn test_cert_watcher_file_updates() -> Result<()> {
let temp_dir = tempdir()?;
let key_path = Utf8PathBuf::from_path_buf(temp_dir.path().join("test.key")).unwrap();
let cert_path = Utf8PathBuf::from_path_buf(temp_dir.path().join("test.crt")).unwrap();
let context = Arc::new(RunContext::new(crate::config::Config::default()));
let so1 = TEST_HOST_CERTS.snakeoil_1.clone();
tokio::fs::copy(so1.keyfile(), &key_path).await?;
tokio::fs::copy(so1.certfile(), &cert_path).await?;
let hc = HostCertificate::new(key_path.clone(), cert_path.clone(), true).await?;
let original_host = hc.hostnames()[0].clone();
let store = Arc::new(CertStore::new(context.clone())?);
store.upsert_all(vec![hc])?;
let original_cert = store.by_host(&original_host).unwrap();
let original_expiry = original_cert.expires();
let mut watcher = CertWatcher::new(store.clone(), context.clone());
let watcher_handle = tokio::spawn(async move {
watcher.watch().await
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
println!("Updating cert files");
let so2 = TEST_HOST_CERTS.snakeoil_2.clone();
tokio::fs::copy(so2.keyfile(), &key_path).await?;
tokio::fs::copy(so2.certfile(), &cert_path).await?;
tokio::time::sleep(RELOAD_GRACE + std::time::Duration::from_millis(500)).await;
info!("Checking updated certs");
let updated_cert = store.by_host(&original_host).unwrap();
let updated_expiry = updated_cert.expires();
assert_ne!(original_expiry, updated_expiry);
context.quit()?;
watcher_handle.await??;
Ok(())
}
#[tokio::test]
async fn test_by_host() {
let context = Arc::new(RunContext::new(Config::default()));
let store = CertStore::new(context).unwrap();
let cert = TEST_HOST_CERTS.snakeoil_1.clone();
store.upsert(cert.clone()).unwrap();
let found = store.by_host(&cert.hostnames()[0]).unwrap();
assert_eq!(found, cert);
}
#[tokio::test]
async fn test_by_host_is_case_insensitive() {
let context = Arc::new(RunContext::new(Config::default()));
let store = CertStore::new(context).unwrap();
let cert = TEST_HOST_CERTS.snakeoil_1.clone();
store.upsert(cert.clone()).unwrap();
let mixed_case_host = cert.hostnames()[0].to_ascii_uppercase();
let found = store
.by_host(&mixed_case_host)
.expect("host certificate lookup should ignore DNS hostname case");
assert_eq!(found, cert);
}
#[tokio::test]
async fn test_by_file() {
let context = Arc::new(RunContext::new(Config::default()));
let store = CertStore::new(context).unwrap();
let cert = TEST_HOST_CERTS.snakeoil_1.clone();
store.upsert(cert.clone()).unwrap();
let found = store.by_file(&"target/certs/snakeoil-1.key".into()).unwrap();
assert_eq!(found, cert);
}
#[tokio::test]
async fn test_watchlist() -> Result<()> {
let hc1 = TEST_HOST_CERTS.snakeoil_1.clone();
let hc2 = TEST_HOST_CERTS.www_example.clone();
let context = Arc::new(RunContext::new(Config::default()));
let store = CertStore::new(context)?;
store.upsert(hc1)?;
store.upsert(hc2)?;
let watchlist = store.watchlist();
assert_eq!(watchlist.len(), 2);
assert!(watchlist.contains(&Utf8PathBuf::from("target/certs/snakeoil-1.key")));
assert!(watchlist.contains(&Utf8PathBuf::from("target/certs/snakeoil-1.crt")));
Ok(())
}
#[tokio::test]
async fn test_file_update_success() -> Result<()> {
let temp_dir = tempdir()?;
let key_path = temp_dir.path().join("test.key");
let cert_path = temp_dir.path().join("test.crt");
let cert = TEST_HOST_CERTS.snakeoil_1.clone();
fs::copy(cert.keyfile(), &key_path)?;
fs::copy(cert.certfile(), &cert_path)?;
let context = Arc::new(RunContext::new(Config::default()));
let store = CertStore::new(context)?;
store.upsert(cert.clone())?;
let original_host = cert.hostnames()[0].clone();
let first_cert = store.by_host(&original_host).unwrap();
assert_eq!("snakeoil.example.com", first_cert.hostnames()[0]);
let cert = TEST_HOST_CERTS.snakeoil_2.clone();
fs::copy(cert.keyfile(), &key_path)?;
fs::copy(cert.certfile(), &cert_path)?;
let newcert = HostCertificate::new_from(&first_cert).await?;
store.update(newcert)?;
let updated_cert_from_file = HostCertificate::new(
Utf8PathBuf::from_path_buf(key_path).unwrap(),
Utf8PathBuf::from_path_buf(cert_path).unwrap(),
true
).await?;
let new_host = updated_cert_from_file.hostnames()[0].clone();
let updated_cert_from_store = store.by_host(&new_host).expect("Cert not found for new host");
assert_eq!(updated_cert_from_store.hostnames()[0], new_host);
if original_host != new_host {
assert!(store.by_host(&original_host).is_none(), "Old host entry should be removed");
}
Ok(())
}
#[test]
fn test_to_txt_name() {
let domain = "example.com";
assert_eq!("_acme-challenge.www", to_txt_name(domain, "www.example.com"));
assert_eq!("_acme-challenge.www.dev", to_txt_name(domain, "www.dev.example.com"));
assert_eq!("_acme-challenge", to_txt_name(domain, "example.com"));
assert_eq!("_acme-challenge", to_txt_name(domain, "example.com."));
assert_eq!("_acme-challenge", to_txt_name(domain, ""));
assert_eq!("_acme-challenge", to_txt_name(domain, "*.example.com"));
assert_eq!("_acme-challenge.dev", to_txt_name(domain, "*.dev.example.com"));
}
#[tokio::test]
async fn test_wildcard() -> Result<()> {
let wildcard = TEST_HOST_CERTS.wildcard_example.clone();
assert!(wildcard.hostnames().contains(&"*.example.com".to_string()));
let context = Arc::new(RunContext::new(Config::default()));
let store = CertStore::new(context)?;
store.upsert(wildcard.clone())?;
{
let by_host = store.by_host("otherhost.example.com");
assert!(by_host.is_none());
let by_wildcard = store.by_wildcard("otherhost.example.com").unwrap();
assert_eq!(Some(&"*.example.com".to_string()), by_wildcard.hostnames().first());
}
{
let by_host = store.by_host("*.example.com").unwrap();
assert_eq!(Some(&"*.example.com".to_string()), by_host.hostnames().first());
let by_wildcard = store.by_wildcard("realhost.example.com").unwrap();
assert_eq!(Some(&"*.example.com".to_string()), by_wildcard.hostnames().first());
}
Ok(())
}