use std::time::Duration;
use serde::Deserialize;
use tokio::io::AsyncReadExt;
use tokio_vsock::{VsockAddr, VsockStream};
use tracing::{info, warn};
use vta_config::AppConfig;
use vta_config::tenant_overlay::{TenantConfigOverlay, TenantOverlayError, apply_tenant_overlay};
const PARENT_CID: u32 = 3;
const CONFIG_PORT: u32 = 5800;
const SUPPORTED_ENVELOPE_VERSION: u32 = 1;
const MAX_ENVELOPE_BYTES: usize = 1024 * 1024;
const CONNECT_MAX_ATTEMPTS: u32 = 30;
const CONNECT_BASE_DELAY: Duration = Duration::from_millis(500);
const CONNECT_MAX_DELAY: Duration = Duration::from_secs(3);
const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const READ_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Debug)]
pub enum OverlayFetchError {
Connect(String),
Read(String),
TooLarge(usize),
Parse(String),
UnsupportedVersion(u32),
Apply(TenantOverlayError),
}
impl std::fmt::Display for OverlayFetchError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Connect(e) => write!(f, "vsock connect to config server failed: {e}"),
Self::Read(e) => write!(f, "reading config envelope failed: {e}"),
Self::TooLarge(n) => write!(
f,
"config envelope exceeded the {MAX_ENVELOPE_BYTES}-byte cap ({n} bytes)"
),
Self::Parse(e) => write!(f, "config envelope did not parse: {e}"),
Self::UnsupportedVersion(v) => write!(
f,
"config envelope version {v} is unsupported (this enclave speaks \
v{SUPPORTED_ENVELOPE_VERSION})"
),
Self::Apply(e) => write!(f, "applying tenant overlay failed: {e}"),
}
}
}
impl std::error::Error for OverlayFetchError {}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct ConfigEnvelope {
#[allow(dead_code)]
version: u32,
overlay: TenantConfigOverlay,
#[serde(default)]
#[allow(dead_code)]
integrity: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
struct EnvelopeVersionProbe {
version: u32,
}
fn parse_envelope(bytes: &[u8]) -> Result<TenantConfigOverlay, OverlayFetchError> {
let probe: EnvelopeVersionProbe =
serde_json::from_slice(bytes).map_err(|e| OverlayFetchError::Parse(e.to_string()))?;
if probe.version != SUPPORTED_ENVELOPE_VERSION {
return Err(OverlayFetchError::UnsupportedVersion(probe.version));
}
let envelope: ConfigEnvelope =
serde_json::from_slice(bytes).map_err(|e| OverlayFetchError::Parse(e.to_string()))?;
Ok(envelope.overlay)
}
async fn connect_with_retry() -> Result<VsockStream, OverlayFetchError> {
let addr = VsockAddr::new(PARENT_CID, CONFIG_PORT);
let mut delay = CONNECT_BASE_DELAY;
for attempt in 1..=CONNECT_MAX_ATTEMPTS {
match tokio::time::timeout(CONNECT_TIMEOUT, VsockStream::connect(addr)).await {
Ok(Ok(stream)) => {
if attempt > 1 {
info!("connected to config server on attempt {attempt}");
}
return Ok(stream);
}
Ok(Err(e)) if attempt == CONNECT_MAX_ATTEMPTS => {
return Err(OverlayFetchError::Connect(e.to_string()));
}
Err(_) if attempt == CONNECT_MAX_ATTEMPTS => {
return Err(OverlayFetchError::Connect("connect timed out".into()));
}
Ok(Err(e)) => warn!(
"config server not ready on vsock:{CONFIG_PORT} \
(attempt {attempt}/{CONNECT_MAX_ATTEMPTS}): {e}; retrying in {delay:?}"
),
Err(_) => warn!(
"config server connect timed out \
(attempt {attempt}/{CONNECT_MAX_ATTEMPTS}); retrying in {delay:?}"
),
}
tokio::time::sleep(delay).await;
delay = (delay * 2).min(CONNECT_MAX_DELAY);
}
unreachable!("the loop returns on the final attempt")
}
async fn read_envelope(stream: &mut VsockStream) -> Result<Vec<u8>, OverlayFetchError> {
match tokio::time::timeout(READ_TIMEOUT, read_envelope_to_eof(stream)).await {
Ok(result) => result,
Err(_) => Err(OverlayFetchError::Read(format!(
"envelope not fully received within {READ_TIMEOUT:?}"
))),
}
}
async fn read_envelope_to_eof(stream: &mut VsockStream) -> Result<Vec<u8>, OverlayFetchError> {
let mut buf = Vec::new();
let mut chunk = [0u8; 8192];
loop {
let n = stream
.read(&mut chunk)
.await
.map_err(|e| OverlayFetchError::Read(e.to_string()))?;
if n == 0 {
break; }
if buf.len() + n > MAX_ENVELOPE_BYTES {
return Err(OverlayFetchError::TooLarge(buf.len() + n));
}
buf.extend_from_slice(&chunk[..n]);
}
Ok(buf)
}
pub async fn fetch_and_apply_overlay(config: &mut AppConfig) -> Result<(), OverlayFetchError> {
info!("fetching tenant-config overlay over vsock:{CONFIG_PORT}");
let mut stream = connect_with_retry().await?;
let bytes = read_envelope(&mut stream).await?;
if bytes.is_empty() {
return Err(OverlayFetchError::Read(
"parent served an empty envelope".into(),
));
}
let overlay = parse_envelope(&bytes)?;
apply_tenant_overlay(config, overlay).map_err(OverlayFetchError::Apply)?;
info!("tenant-config overlay applied");
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
const GOOD_ARN: &str = "arn:aws:kms:us-east-1:111122223333:key/abcd-ef01";
#[test]
fn parses_a_well_formed_v1_envelope() {
let env = format!(
r#"{{"version":1,"overlay":{{"vta_name":"acme","tee_kms":{{"key_arn":"{GOOD_ARN}"}}}},"integrity":null}}"#
);
let overlay = parse_envelope(env.as_bytes()).expect("valid envelope");
assert_eq!(overlay.vta_name.as_deref(), Some("acme"));
assert_eq!(overlay.tee_kms.unwrap().key_arn, GOOD_ARN);
}
#[test]
fn rejects_wrong_version() {
let env = r#"{"version":2,"overlay":{},"integrity":null}"#;
assert!(matches!(
parse_envelope(env.as_bytes()),
Err(OverlayFetchError::UnsupportedVersion(2))
));
}
#[test]
fn rejects_overlay_with_forbidden_field() {
let env =
r#"{"version":1,"overlay":{"tee_kms":{"key_arn":"x","admin_did":"did:key:zEvil"}}}"#;
assert!(matches!(
parse_envelope(env.as_bytes()),
Err(OverlayFetchError::Parse(_))
));
}
#[test]
fn rejects_envelope_with_unknown_field() {
let env = r#"{"version":1,"overlay":{},"unexpected":"value"}"#;
assert!(matches!(
parse_envelope(env.as_bytes()),
Err(OverlayFetchError::Parse(_))
));
}
#[test]
fn rejects_malformed_json() {
assert!(matches!(
parse_envelope(b"not json"),
Err(OverlayFetchError::Parse(_))
));
}
#[test]
fn a_future_version_reports_the_version_not_a_parse_error() {
let v2 =
r#"{"version":2,"overlay":{"vta_name":"acme","future_field":"x"},"integrity":null}"#;
assert!(
matches!(
parse_envelope(v2.as_bytes()),
Err(OverlayFetchError::UnsupportedVersion(2))
),
"a v2 envelope must be reported as an unsupported version"
);
let v3 = r#"{"version":3,"overlay":{},"signature":"..."}"#;
assert!(matches!(
parse_envelope(v3.as_bytes()),
Err(OverlayFetchError::UnsupportedVersion(3))
));
let v1_unknown = r#"{"version":1,"overlay":{"future_field":"x"}}"#;
assert!(matches!(
parse_envelope(v1_unknown.as_bytes()),
Err(OverlayFetchError::Parse(_))
));
assert!(matches!(
parse_envelope(br#"{"overlay":{}}"#),
Err(OverlayFetchError::Parse(_))
));
}
}