use anyhow::{Context, Result, bail};
use vta_sdk::client::{AutoConnect, ClientIdentity, ConnectedVta, VtaClient};
use zeroize::Zeroize;
use crate::config::{self, SigningConfig, VtaCredentials};
const MAX_AUTH_RETRIES: u32 = 2;
pub async fn authenticate(cfg: &SigningConfig) -> Result<(VtaClient, VtaCredentials)> {
let creds = config::load_vta_credentials(&cfg.did_key_id)?;
validate_credentials(&creds)?;
if creds.mediator_did.is_none()
&& let Some(token) = config::load_cached_token(&cfg.did_key_id)
{
let identity = ClientIdentity::did_key(
&creds.credential_did,
&creds.private_key_multibase,
&creds.vta_did,
);
let client = VtaClient::new(&creds.vta_url).with_identity(identity);
client.set_token(token);
return Ok((client, creds));
}
let connected = connect_with_retry(&creds).await?;
if let Some(token) = &connected.rest_token {
let _ = config::cache_token(
&cfg.did_key_id,
&token.access_token,
token.access_expires_at,
);
}
Ok((connected.client, creds))
}
fn is_loopback_http(url: &str) -> bool {
let Some(rest) = url.strip_prefix("http://") else {
return false;
};
if let Some(after) = rest.strip_prefix('[') {
let Some((host, tail)) = after.split_once(']') else {
return false;
};
return host == "::1" && (tail.is_empty() || tail.starts_with([':', '/', '?']));
}
let host = rest.split(['/', ':', '?']).next().unwrap_or("");
host == "localhost" || host == "127.0.0.1"
}
fn vta_url_is_secure(url: &str) -> bool {
url.starts_with("https://") || is_loopback_http(url)
}
fn validate_credentials(creds: &VtaCredentials) -> Result<()> {
if creds.credential_did.is_empty() {
bail!("credential DID is empty");
}
if creds.key_id.is_empty() {
bail!("signing key ID is empty");
}
if creds.mediator_did.is_some() {
if !creds.vta_url.is_empty() && !vta_url_is_secure(&creds.vta_url) {
bail!(
"VTA URL must use HTTPS (got: {}). Cleartext http:// is allowed only to \
loopback (localhost, 127.0.0.1, [::1]) for local development.",
creds.vta_url
);
}
return Ok(());
}
if creds.vta_url.is_empty() {
bail!("VTA URL is empty");
}
if !vta_url_is_secure(&creds.vta_url) {
bail!(
"VTA URL must use HTTPS (got: {}). Cleartext http:// is allowed only to \
loopback (localhost, 127.0.0.1, [::1]) for local development.",
creds.vta_url
);
}
Ok(())
}
async fn connect_with_retry(creds: &VtaCredentials) -> Result<ConnectedVta> {
let mut last_err = None;
for attempt in 1..=MAX_AUTH_RETRIES {
let result = VtaClient::connect_auto(AutoConnect {
vta_url: &creds.vta_url,
vta_did: &creds.vta_did,
credential_did: &creds.credential_did,
private_key_multibase: &creds.private_key_multibase,
mediator_did: creds.mediator_did.as_deref(),
})
.await;
match result {
Ok(connected) => {
if let Some(token) = &connected.rest_token
&& token.access_token.is_empty()
{
bail!("VTA returned an empty access token");
}
return Ok(connected);
}
Err(e) => {
let err_msg = format!("{e}");
if attempt < MAX_AUTH_RETRIES {
eprintln!(
"VTA connect attempt {attempt}/{MAX_AUTH_RETRIES} failed: {err_msg}, retrying..."
);
}
last_err = Some(err_msg);
}
}
}
bail!(
"VTA connection failed after {MAX_AUTH_RETRIES} attempts: {}",
last_err.unwrap_or_else(|| "unknown error".to_string())
)
}
pub async fn get_signing_key(client: &VtaClient, key_id: &str) -> Result<SeedMaterial> {
let resp = client
.get_key_secret(key_id)
.await
.map_err(|e| anyhow::anyhow!("failed to fetch key secret: {e}"))?;
if resp.key_type != vta_sdk::keys::KeyType::Ed25519 {
bail!(
"signing key {key_id} is {:?}, expected Ed25519",
resp.key_type
);
}
let seed = vta_sdk::did_key::decode_private_key_multibase(&resp.private_key_multibase)
.context("failed to decode signing key")?;
Ok(SeedMaterial(seed))
}
pub struct SeedMaterial([u8; 32]);
impl SeedMaterial {
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
impl Drop for SeedMaterial {
fn drop(&mut self) {
self.0.zeroize();
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_creds() -> VtaCredentials {
VtaCredentials {
vta_url: "https://vta.example.com".to_string(),
vta_did: "did:example:vta".to_string(),
credential_did: "did:key:z6Mk123".to_string(),
private_key_multibase: "z...".to_string(),
key_id: "key-1".to_string(),
mediator_did: None,
}
}
#[test]
fn test_validate_rejects_empty_url() {
let mut creds = test_creds();
creds.vta_url = "".to_string();
assert!(validate_credentials(&creds).is_err());
}
#[test]
fn test_validate_rejects_http() {
let mut creds = test_creds();
creds.vta_url = "http://example.com".to_string();
assert!(validate_credentials(&creds).is_err());
}
#[test]
fn test_validate_allows_https() {
assert!(validate_credentials(&test_creds()).is_ok());
}
#[test]
fn test_validate_allows_localhost() {
let mut creds = test_creds();
creds.vta_url = "http://localhost:3000".to_string();
assert!(validate_credentials(&creds).is_ok());
}
#[test]
fn cleartext_is_allowed_to_every_loopback_form() {
for url in [
"http://localhost",
"http://localhost:3000",
"http://localhost/path",
"http://127.0.0.1:8100",
"http://[::1]:8100",
] {
let mut creds = test_creds();
creds.vta_url = url.to_string();
assert!(
validate_credentials(&creds).is_ok(),
"loopback must stay usable for local dev: {url}"
);
}
}
#[test]
fn cleartext_is_rejected_to_hosts_that_only_look_like_loopback() {
for url in [
"http://localhost.evil.com",
"http://localhost.evil.com/vta",
"http://localhostevil.com",
"http://127.0.0.1.evil.com",
"http://[::1].evil.com",
] {
let mut creds = test_creds();
creds.vta_url = url.to_string();
assert!(
validate_credentials(&creds).is_err(),
"a lookalike host must not pass the HTTPS requirement: {url}"
);
}
}
#[test]
fn the_didcomm_branch_applies_the_same_host_rule() {
let mut creds = test_creds();
creds.mediator_did = Some("did:web:mediator.example".to_string());
creds.vta_url = String::new();
assert!(validate_credentials(&creds).is_ok(), "empty stays allowed");
creds.vta_url = "http://localhost:3000".to_string();
assert!(validate_credentials(&creds).is_ok());
creds.vta_url = "http://localhost.evil.com".to_string();
assert!(
validate_credentials(&creds).is_err(),
"the lookalike must fail on the DIDComm path too"
);
}
#[test]
fn test_validate_rejects_empty_key_id() {
let mut creds = test_creds();
creds.key_id = "".to_string();
assert!(validate_credentials(&creds).is_err());
}
#[test]
fn test_validate_rejects_empty_credential_did() {
let mut creds = test_creds();
creds.credential_did = "".to_string();
assert!(validate_credentials(&creds).is_err());
}
#[test]
fn test_seed_material_zeroizes_on_drop() {
let seed = SeedMaterial([0xAB; 32]);
assert_eq!(seed.as_bytes(), &[0xAB; 32]);
drop(seed);
}
#[test]
fn test_validate_didcomm_only_accepts_empty_url() {
let mut creds = test_creds();
creds.vta_url = "".to_string();
creds.mediator_did = Some("did:peer:0z6Mkmediator".to_string());
assert!(validate_credentials(&creds).is_ok());
}
#[test]
fn test_validate_didcomm_with_url_still_requires_https() {
let mut creds = test_creds();
creds.vta_url = "http://example.com".to_string();
creds.mediator_did = Some("did:peer:0z6Mkmediator".to_string());
assert!(validate_credentials(&creds).is_err());
}
#[test]
fn test_validate_didcomm_still_rejects_empty_credential_did() {
let mut creds = test_creds();
creds.vta_url = "".to_string();
creds.mediator_did = Some("did:peer:0z6Mkmediator".to_string());
creds.credential_did = "".to_string();
assert!(validate_credentials(&creds).is_err());
}
}