use super::builder::{ReconnectConfig, ResourceLimits};
use super::errors::{MetricsErrorKind, X509SourceError};
use super::limits::{select_svid, validate_context};
use super::metrics::MetricsRecorder;
use super::supervisor::initial_sync_with_retry;
use crate::bundle::BundleSource;
use crate::cert::Certificate;
use crate::prelude::warn;
use crate::svid::SvidSource;
use crate::workload_api::x509_context::X509Context;
use crate::x509_source::types::{ClientFactory, SvidPicker};
use crate::{TrustDomain, X509Bundle, X509BundleSet, X509SourceBuilder, X509Svid};
use arc_swap::ArcSwap;
use std::cmp::Ordering as CmpOrdering;
use std::fmt::Debug;
use std::future::Future;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::sync::{watch, Mutex};
use tokio::task::JoinHandle;
use tokio_util::sync::{CancellationToken, DropGuard};
#[cfg(test)]
use crate::WorkloadApiError;
#[derive(Clone, Debug)]
pub struct X509SourceUpdates {
rx: watch::Receiver<u64>,
shutdown: CancellationToken,
}
impl X509SourceUpdates {
pub async fn changed(&mut self) -> Result<u64, X509SourceError> {
if self.rx.has_changed().unwrap_or(false) {
self.rx
.changed()
.await
.map_err(|watch::error::RecvError { .. }| X509SourceError::Closed)?;
return Ok(*self.rx.borrow());
}
if self.shutdown.is_cancelled() {
return Err(X509SourceError::Closed);
}
tokio::select! {
biased;
result = self.rx.changed() => {
result.map_err(|watch::error::RecvError { .. }| X509SourceError::Closed)?;
Ok(*self.rx.borrow())
}
() = self.shutdown.cancelled() => Err(X509SourceError::Closed),
}
}
pub fn last(&self) -> u64 {
*self.rx.borrow()
}
pub async fn wait_for<F>(&mut self, mut f: F) -> Result<u64, X509SourceError>
where
F: FnMut(&u64) -> bool,
{
if self.shutdown.is_cancelled() {
return Err(X509SourceError::Closed);
}
let current = self.last();
if f(¤t) {
return Ok(current);
}
loop {
let seq = self.changed().await?;
if f(&seq) {
return Ok(seq);
}
}
}
}
#[derive(Clone, Debug)]
pub struct X509Source {
inner: Arc<Inner>,
_shutdown_guard: Arc<DropGuard>,
}
struct Snapshot {
ctx: Arc<X509Context>,
expiry_unix: i64,
}
pub(super) struct Inner {
snapshot: ArcSwap<Snapshot>,
svid_picker: Option<Box<dyn SvidPicker>>,
limits: ResourceLimits,
reconnect: ReconnectConfig,
make_client: ClientFactory,
metrics: Option<Arc<dyn MetricsRecorder>>,
closed: AtomicBool,
supervisor_running: AtomicBool,
cancel: CancellationToken,
shutdown_timeout: Option<Duration>,
update_seq: AtomicU64,
update_tx: watch::Sender<u64>,
supervisor: Mutex<Option<JoinHandle<()>>>,
}
impl Inner {
pub(super) const fn reconnect(&self) -> ReconnectConfig {
self.reconnect
}
pub(super) fn metrics(&self) -> Option<&dyn MetricsRecorder> {
self.metrics.as_deref()
}
pub(super) fn make_client(&self) -> &ClientFactory {
&self.make_client
}
}
impl Debug for Inner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("X509Source")
.field("snapshot", &"<ArcSwap<Snapshot>>")
.field(
"svid_picker",
&self.svid_picker.as_ref().map(|_| "<SvidPicker>"),
)
.field("reconnect", &self.reconnect)
.field("limits", &self.limits)
.field("make_client", &"<ClientFactory>")
.field(
"metrics",
&self.metrics.as_ref().map(|_| "<MetricsRecorder>"),
)
.field("shutdown_timeout", &self.shutdown_timeout)
.field("closed", &self.closed.load(Ordering::Relaxed))
.field(
"supervisor_running",
&self.supervisor_running.load(Ordering::Relaxed),
)
.field("cancel", &self.cancel)
.field("update_seq", &self.update_seq)
.field("update_tx", &"<watch::Sender<u64>>")
.field("supervisor", &"<Mutex<Option<JoinHandle<()>>>>")
.finish()
}
}
impl X509Source {
pub async fn new() -> Result<Self, X509SourceError> {
X509SourceBuilder::new().build().await
}
pub fn builder() -> X509SourceBuilder {
X509SourceBuilder::new()
}
pub fn updated(&self) -> X509SourceUpdates {
X509SourceUpdates {
rx: self.inner.update_tx.subscribe(),
shutdown: self.inner.cancel.clone(),
}
}
pub fn is_healthy(&self) -> bool {
if self.inner.closed.load(Ordering::Acquire)
|| self.inner.cancel.is_cancelled()
|| !self.inner.supervisor_running.load(Ordering::Acquire)
{
return false;
}
let Some(now) = unix_timestamp_now() else {
return false;
};
self.inner.snapshot.load().expiry_unix > now
}
pub fn x509_context(&self) -> Result<Arc<X509Context>, X509SourceError> {
self.assert_open()?;
Ok(Arc::clone(&self.inner.snapshot.load().ctx))
}
pub fn svid(&self) -> Result<Arc<X509Svid>, X509SourceError> {
self.assert_open()?;
let snapshot = self.inner.snapshot.load();
select_svid(&snapshot.ctx, self.inner.svid_picker.as_deref()).ok_or_else(|| {
self.inner.record_error(MetricsErrorKind::NoSuitableSvid);
X509SourceError::NoSuitableSvid
})
}
pub fn try_svid(&self) -> Option<Arc<X509Svid>> {
self.svid().ok()
}
pub fn bundle_set(&self) -> Result<Arc<X509BundleSet>, X509SourceError> {
self.assert_open()?;
Ok(Arc::clone(self.inner.snapshot.load().ctx.bundle_set()))
}
pub fn try_bundle_for_trust_domain(&self, td: &TrustDomain) -> Option<Arc<X509Bundle>> {
self.bundle_for_trust_domain(td).ok().flatten()
}
pub async fn shutdown(&self) {
if self.inner.closed.swap(true, Ordering::AcqRel) {
return;
}
self.inner.cancel.cancel();
if let Some(handle) = self.inner.supervisor.lock().await.take() {
if let Err(e) = handle.await {
warn!("Error joining supervisor task during shutdown: error={e}");
self.inner
.record_error(MetricsErrorKind::SupervisorJoinFailed);
}
}
}
pub async fn shutdown_with_timeout(&self, timeout: Duration) -> Result<(), X509SourceError> {
if self.inner.closed.swap(true, Ordering::AcqRel) {
return Ok(());
}
self.inner.cancel.cancel();
let Some(mut handle) = self.inner.supervisor.lock().await.take() else {
return Ok(());
};
match tokio::time::timeout(timeout, &mut handle).await {
Ok(Ok(())) => Ok(()),
Ok(Err(e)) => {
warn!("Error joining supervisor task during shutdown: error={e}");
self.inner
.record_error(MetricsErrorKind::SupervisorJoinFailed);
Ok(())
}
Err(_) => {
warn!("Shutdown timeout exceeded; aborting supervisor task");
handle.abort();
let _unused: Result<_, _> = handle.await;
Err(X509SourceError::ShutdownTimeout)
}
}
}
pub async fn shutdown_configured(&self) -> Result<(), X509SourceError> {
if let Some(timeout) = self.inner.shutdown_timeout {
self.shutdown_with_timeout(timeout).await
} else {
self.shutdown().await;
Ok(())
}
}
}
impl X509Source {
pub(super) async fn build_with(
make_client: ClientFactory,
svid_picker: Option<Box<dyn SvidPicker>>,
reconnect: ReconnectConfig,
limits: ResourceLimits,
metrics: Option<Arc<dyn MetricsRecorder>>,
shutdown_timeout: Option<Duration>,
initial_sync_timeout: Option<Duration>,
) -> Result<Self, X509SourceError> {
let reconnect = super::builder::normalize_reconnect(reconnect);
let (update_tx, _update_rx) = watch::channel(0u64);
let cancel = CancellationToken::new();
let shutdown_guard = Arc::new(cancel.clone().drop_guard());
let initial_sync = initial_sync_with_retry(
&make_client,
svid_picker.as_deref(),
&cancel,
reconnect,
limits,
metrics.as_deref(),
);
let (initial_ctx, selected_svid) =
initial_sync_with_timeout(initial_sync, &cancel, initial_sync_timeout).await?;
let snapshot = Arc::new(Snapshot {
ctx: initial_ctx,
expiry_unix: selected_svid.expiry_unix(),
});
let inner = Arc::new(Inner {
snapshot: ArcSwap::from(snapshot),
svid_picker,
reconnect,
make_client,
limits,
metrics,
shutdown_timeout,
closed: AtomicBool::new(false),
supervisor_running: AtomicBool::new(false),
cancel,
update_seq: AtomicU64::new(0),
update_tx,
supervisor: Mutex::new(None),
});
let task_inner = Arc::clone(&inner);
let token = task_inner.cancel.clone();
let guard_inner = Arc::clone(&task_inner);
let handle = tokio::spawn(async move {
let _terminate_on_drop = SupervisorTerminationGuard::new(guard_inner);
task_inner.run_update_supervisor(token).await;
});
*inner.supervisor.lock().await = Some(handle);
Ok(Self {
inner,
_shutdown_guard: shutdown_guard,
})
}
#[cfg(test)]
pub(super) fn new_for_test(
initial_ctx: Arc<X509Context>,
reconnect: ReconnectConfig,
limits: ResourceLimits,
metrics: Option<Arc<dyn MetricsRecorder>>,
svid_picker: Option<Box<dyn SvidPicker>>,
) -> Self {
let reconnect = super::builder::normalize_reconnect(reconnect);
let (update_tx, _update_rx) = watch::channel(0u64);
let cancel = CancellationToken::new();
let shutdown_guard = Arc::new(cancel.clone().drop_guard());
let make_client: ClientFactory =
Arc::new(|| Box::pin(async move { Err(WorkloadApiError::EmptyResponse) }));
let selected_svid_expires_at_unix =
select_svid(&initial_ctx, svid_picker.as_deref()).map_or(0, |svid| svid.expiry_unix());
let snapshot = Arc::new(Snapshot {
ctx: initial_ctx,
expiry_unix: selected_svid_expires_at_unix,
});
let inner = Inner {
snapshot: ArcSwap::from(snapshot),
svid_picker,
reconnect,
make_client,
limits,
metrics,
shutdown_timeout: None,
closed: AtomicBool::new(false),
supervisor_running: AtomicBool::new(false),
cancel,
update_seq: AtomicU64::new(0),
update_tx,
supervisor: Mutex::new(None),
};
Self {
inner: Arc::new(inner),
_shutdown_guard: shutdown_guard,
}
}
fn assert_open(&self) -> Result<(), X509SourceError> {
if self.inner.closed.load(Ordering::Acquire) || self.inner.cancel.is_cancelled() {
return Err(X509SourceError::Closed);
}
Ok(())
}
}
struct SupervisorTerminationGuard {
inner: Arc<Inner>,
}
impl SupervisorTerminationGuard {
fn new(inner: Arc<Inner>) -> Self {
inner.supervisor_running.store(true, Ordering::Release);
Self { inner }
}
}
impl Drop for SupervisorTerminationGuard {
fn drop(&mut self) {
self.inner
.supervisor_running
.store(false, Ordering::Release);
self.inner.cancel.cancel();
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(super) enum ApplyUpdateResult {
Applied,
Unchanged,
}
fn unix_timestamp_now() -> Option<i64> {
let duration = SystemTime::now().duration_since(UNIX_EPOCH).ok()?;
i64::try_from(duration.as_secs()).ok()
}
impl Inner {
pub(super) fn record_error(&self, kind: MetricsErrorKind) {
if let Some(metrics) = self.metrics.as_deref() {
metrics.record_error(kind);
}
}
pub(super) fn record_update(&self) {
if let Some(metrics) = self.metrics.as_deref() {
metrics.record_update();
}
}
pub(super) fn apply_update(
&self,
new_ctx: Arc<X509Context>,
) -> Result<ApplyUpdateResult, X509SourceError> {
match self.validate_and_select(&new_ctx) {
Ok(svid) => {
if same_material_for_update(self.snapshot.load().ctx.as_ref(), new_ctx.as_ref()) {
return Ok(ApplyUpdateResult::Unchanged);
}
self.snapshot.store(Arc::new(Snapshot {
ctx: new_ctx,
expiry_unix: svid.expiry_unix(),
}));
self.notify_update();
self.record_update();
Ok(ApplyUpdateResult::Applied)
}
Err(e) => {
self.record_error(MetricsErrorKind::UpdateRejected);
Err(e)
}
}
}
pub(super) fn notify_update(&self) {
let next = self.update_seq.fetch_add(1, Ordering::Relaxed) + 1;
let _prev = self.update_tx.send_replace(next);
}
pub(super) fn validate_and_select(
&self,
ctx: &X509Context,
) -> Result<Arc<X509Svid>, X509SourceError> {
validate_context(
ctx,
self.svid_picker.as_deref(),
self.limits,
self.metrics.as_deref(),
)
}
}
fn same_material_for_update(current: &X509Context, incoming: &X509Context) -> bool {
if !bundle_set_equal_for_update(current.bundle_set(), incoming.bundle_set()) {
return false;
}
if current.svids().len() != incoming.svids().len() {
return false;
}
let mut left: Vec<&X509Svid> = current.svids().iter().map(AsRef::as_ref).collect();
let mut right: Vec<&X509Svid> = incoming.svids().iter().map(AsRef::as_ref).collect();
left.sort_unstable_by(|a, b| cmp_svid_for_update_dedupe(a, b));
right.sort_unstable_by(|a, b| cmp_svid_for_update_dedupe(a, b));
left == right
}
fn bundle_set_equal_for_update(a: &X509BundleSet, b: &X509BundleSet) -> bool {
if a.len() != b.len() {
return false;
}
a.iter().all(|(trust_domain, bundle)| {
b.get(trust_domain)
.is_some_and(|other| bundle_equal_for_update(bundle.as_ref(), other.as_ref()))
})
}
fn bundle_equal_for_update(a: &X509Bundle, b: &X509Bundle) -> bool {
a.trust_domain() == b.trust_domain()
&& authority_set_equal_for_update(a.authorities(), b.authorities())
}
fn authority_set_equal_for_update(a: &[Certificate], b: &[Certificate]) -> bool {
if a.len() != b.len() {
return false;
}
let mut left: Vec<&[u8]> = a.iter().map(Certificate::as_bytes).collect();
let mut right: Vec<&[u8]> = b.iter().map(Certificate::as_bytes).collect();
left.sort_unstable();
right.sort_unstable();
left == right
}
fn cmp_svid_for_update_dedupe(a: &X509Svid, b: &X509Svid) -> CmpOrdering {
a.spiffe_id()
.cmp(b.spiffe_id())
.then_with(|| a.hint().cmp(&b.hint()))
.then_with(|| cmp_cert_chain_for_update(a.cert_chain(), b.cert_chain()))
.then_with(|| a.private_key().as_ref().cmp(b.private_key().as_ref()))
}
fn cmp_cert_chain_for_update(a: &[Certificate], b: &[Certificate]) -> CmpOrdering {
a.iter()
.map(Certificate::as_bytes)
.cmp(b.iter().map(Certificate::as_bytes))
}
async fn initial_sync_with_timeout<T, F>(
initial_sync: F,
cancel: &CancellationToken,
timeout: Option<Duration>,
) -> Result<T, X509SourceError>
where
F: Future<Output = Result<T, X509SourceError>>,
{
let Some(timeout) = timeout else {
return initial_sync.await;
};
match tokio::time::timeout(timeout, initial_sync).await {
Ok(result) => result,
Err(_elapsed) => {
cancel.cancel();
Err(X509SourceError::InitialSyncTimeout)
}
}
}
impl SvidSource for X509Source {
type Item = X509Svid;
type Error = X509SourceError;
fn svid(&self) -> Result<Arc<Self::Item>, Self::Error> {
Self::svid(self)
}
}
impl BundleSource for X509Source {
type Item = X509Bundle;
type Error = X509SourceError;
fn bundle_for_trust_domain(
&self,
trust_domain: &TrustDomain,
) -> Result<Option<Arc<Self::Item>>, Self::Error> {
self.assert_open()?;
let snapshot = self.inner.snapshot.load();
Ok(snapshot.ctx.bundle_set().get(trust_domain))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use std::sync::Mutex;
fn updates_for_test(rx: watch::Receiver<u64>) -> X509SourceUpdates {
X509SourceUpdates {
rx,
shutdown: CancellationToken::new(),
}
}
fn test_x509_context() -> Arc<X509Context> {
let cert_bytes = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key_bytes = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid = Arc::new(X509Svid::parse_from_der(cert_bytes, key_bytes).unwrap());
Arc::new(X509Context::new([svid], Arc::new(X509BundleSet::new())))
}
fn supervisor_running_guard_for_test(source: &X509Source) -> SupervisorTerminationGuard {
SupervisorTerminationGuard::new(Arc::clone(&source.inner))
}
async fn terminate_supervisor_for_test(terminate_guard: SupervisorTerminationGuard) {
tokio::spawn(async move {
let _terminate_on_drop = terminate_guard;
})
.await
.expect("supervisor termination task should not panic");
}
#[tokio::test]
async fn initial_sync_timeout_returns_timeout_and_cancels_token() {
let cancel = CancellationToken::new();
let result = initial_sync_with_timeout(
std::future::pending::<Result<(), X509SourceError>>(),
&cancel,
Some(Duration::ZERO),
)
.await;
assert!(matches!(result, Err(X509SourceError::InitialSyncTimeout)));
assert!(cancel.is_cancelled());
}
#[tokio::test]
async fn initial_sync_timeout_allows_success_before_timeout() {
let cancel = CancellationToken::new();
let result = initial_sync_with_timeout(
async { Ok::<_, X509SourceError>("synced") },
&cancel,
Some(Duration::from_secs(60)),
)
.await;
assert_eq!(result.unwrap(), "synced");
assert!(!cancel.is_cancelled());
}
#[tokio::test]
async fn initial_sync_without_timeout_waits_for_future() {
let cancel = CancellationToken::new();
let result =
initial_sync_with_timeout(async { Ok::<_, X509SourceError>("synced") }, &cancel, None)
.await;
assert_eq!(result.unwrap(), "synced");
assert!(!cancel.is_cancelled());
}
fn set_selected_svid_expiry_for_test(source: &X509Source, expiry_unix: i64) {
let snapshot = source.inner.snapshot.load();
source.inner.snapshot.store(Arc::new(Snapshot {
ctx: Arc::clone(&snapshot.ctx),
expiry_unix,
}));
}
#[tokio::test]
async fn test_wait_for_immediate_satisfaction() {
let (tx, rx) = watch::channel(5u64);
let mut updates = updates_for_test(rx);
let result = updates.wait_for(|&seq| seq > 3).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), 5);
let _unused: Result<_, _> = tx.send(10);
let result = updates.wait_for(|&seq| seq > 8).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), 10);
}
#[tokio::test]
async fn test_wait_for_waits_when_not_satisfied() {
let (tx, rx) = watch::channel(1u64);
let mut updates = updates_for_test(rx);
let tx_clone = tx.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
let _unused: Result<_, _> = tx_clone.send(5);
});
let result = tokio::time::timeout(Duration::from_secs(1), updates.wait_for(|&seq| seq > 3))
.await
.expect("Should complete within timeout");
assert!(result.is_ok());
assert_eq!(result.unwrap(), 5);
}
#[tokio::test]
async fn test_updated_only_notifies_on_rotations_after_initial_sync() {
let (tx, rx) = watch::channel(0u64);
let mut updates = updates_for_test(rx.clone());
let tx_clone = tx.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
let _unused: Result<_, _> = tx_clone.send(1);
});
let result = tokio::time::timeout(Duration::from_secs(1), updates.changed())
.await
.expect("Should complete within timeout");
assert!(result.is_ok());
assert_eq!(result.unwrap(), 1);
assert_eq!(updates.last(), 1);
}
#[tokio::test]
async fn test_updated_initial_sequence_is_zero() {
let (_tx, rx) = watch::channel(0u64);
let updates = updates_for_test(rx);
assert_eq!(updates.last(), 0);
}
#[tokio::test]
async fn test_updated_subscribes_at_current_sequence() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
source.inner.notify_update();
let mut updates = source.updated();
assert_eq!(updates.last(), 1);
assert!(
tokio::time::timeout(Duration::from_millis(20), updates.changed())
.await
.is_err(),
"new subscribers should wait for updates after subscription"
);
source.inner.notify_update();
assert_eq!(updates.changed().await.unwrap(), 2);
}
#[test]
fn test_notify_update_sequence_before_first_subscriber() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
source.inner.notify_update();
let updates = source.updated();
assert_eq!(updates.last(), 1);
}
#[tokio::test]
async fn test_updates_changed_returns_closed_after_shutdown() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let mut updates = source.updated();
source.shutdown().await;
let result = tokio::time::timeout(Duration::from_secs(1), updates.changed())
.await
.expect("changed should return after shutdown");
assert!(matches!(result, Err(X509SourceError::Closed)));
}
#[tokio::test]
async fn test_updates_wait_for_returns_closed_after_shutdown() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let mut updates = source.updated();
source.shutdown().await;
let result = tokio::time::timeout(Duration::from_secs(1), updates.wait_for(|&seq| seq > 0))
.await
.expect("wait_for should return after shutdown");
assert!(matches!(result, Err(X509SourceError::Closed)));
}
#[tokio::test]
async fn wait_for_returns_closed_after_shutdown_even_when_predicate_matches_current() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let mut updates = source.updated();
let current = updates.last();
source.shutdown().await;
let result = updates.wait_for(|seq| *seq == current).await;
assert!(matches!(result, Err(X509SourceError::Closed)));
}
#[tokio::test]
async fn test_updates_changed_delivers_pending_update_before_closed() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let mut updates = source.updated();
source.inner.notify_update();
source.inner.cancel.cancel();
assert_eq!(updates.changed().await.unwrap(), 1);
let result = tokio::time::timeout(Duration::from_secs(1), updates.changed())
.await
.expect("changed should return closed after pending update is observed");
assert!(matches!(result, Err(X509SourceError::Closed)));
}
#[tokio::test]
async fn test_updates_changed_returns_closed_after_last_source_drop() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let mut updates = source.updated();
drop(source);
let result = tokio::time::timeout(Duration::from_secs(1), updates.changed())
.await
.expect("changed should return after last source handle is dropped");
assert!(matches!(result, Err(X509SourceError::Closed)));
}
#[tokio::test]
async fn test_dropping_last_source_handle_cancels_supervisor_token() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let clone = source.clone();
let task_inner = Arc::clone(&source.inner);
let token = task_inner.cancel.clone();
let (stopped_tx, stopped_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
token.cancelled().await;
let _unused: Result<_, _> = stopped_tx.send(());
drop(task_inner);
});
drop(source);
assert!(!clone.inner.cancel.is_cancelled());
drop(clone);
tokio::time::timeout(Duration::from_secs(1), stopped_rx)
.await
.expect("supervisor token should be cancelled when last source handle is dropped")
.expect("supervisor observer should send stop notification");
}
#[tokio::test]
async fn test_supervisor_termination_marks_unhealthy_and_closes_updates() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let running_guard = supervisor_running_guard_for_test(&source);
let mut changed_updates = source.updated();
let mut wait_updates = source.updated();
assert!(
source.is_healthy(),
"cached SVID should be healthy while supervisor is running"
);
terminate_supervisor_for_test(running_guard).await;
assert!(
!source.is_healthy(),
"source must be unhealthy after supervisor termination"
);
let changed = tokio::time::timeout(Duration::from_secs(1), changed_updates.changed())
.await
.expect("changed should stop waiting after supervisor termination");
assert!(matches!(changed, Err(X509SourceError::Closed)));
let waited = tokio::time::timeout(
Duration::from_secs(1),
wait_updates.wait_for(|&seq| seq > 0),
)
.await
.expect("wait_for should stop waiting after supervisor termination");
assert!(matches!(waited, Err(X509SourceError::Closed)));
}
#[tokio::test]
async fn wait_for_returns_closed_after_supervisor_termination_even_when_predicate_matches_current(
) {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let running_guard = supervisor_running_guard_for_test(&source);
let mut updates = source.updated();
let current = updates.last();
terminate_supervisor_for_test(running_guard).await;
let result = updates.wait_for(|seq| *seq == current).await;
assert!(matches!(result, Err(X509SourceError::Closed)));
}
#[test]
fn test_is_healthy_returns_false_when_selected_svid_is_expired() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let _running_guard = supervisor_running_guard_for_test(&source);
let now = unix_timestamp_now().expect("system time should be after UNIX epoch");
assert!(
source.is_healthy(),
"source should be healthy while selected SVID expiry is in the future"
);
set_selected_svid_expiry_for_test(&source, now);
assert!(
!source.is_healthy(),
"source should be unhealthy once selected SVID expiry is reached"
);
}
#[tokio::test]
async fn is_healthy_false_after_shutdown() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let _running_guard = supervisor_running_guard_for_test(&source);
assert!(source.is_healthy());
source.shutdown().await;
assert!(!source.is_healthy());
}
#[test]
fn is_healthy_false_after_cancel() {
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let _running_guard = supervisor_running_guard_for_test(&source);
assert!(source.is_healthy());
source.inner.cancel.cancel();
assert!(!source.is_healthy());
}
struct TestMetricsRecorder {
counts: Arc<Mutex<HashMap<MetricsErrorKind, u64>>>,
}
impl TestMetricsRecorder {
fn new() -> Self {
Self {
counts: Arc::new(Mutex::new(HashMap::new())),
}
}
fn count(&self, kind: MetricsErrorKind) -> u64 {
*self.counts.lock().unwrap().get(&kind).unwrap_or(&0)
}
}
impl MetricsRecorder for TestMetricsRecorder {
fn record_update(&self) {}
fn record_reconnect(&self) {}
fn record_error(&self, kind: MetricsErrorKind) {
*self.counts.lock().unwrap().entry(kind).or_insert(0) += 1;
}
}
struct OrderingMetricsRecorder {
updates: Mutex<Option<X509SourceUpdates>>,
update_sequences: Mutex<Vec<Option<u64>>>,
}
impl OrderingMetricsRecorder {
fn new() -> Self {
Self {
updates: Mutex::new(None),
update_sequences: Mutex::new(Vec::new()),
}
}
fn set_updates(&self, updates: X509SourceUpdates) {
*self.updates.lock().unwrap() = Some(updates);
}
fn update_sequences(&self) -> Vec<Option<u64>> {
self.update_sequences.lock().unwrap().clone()
}
fn record_observed_update(&self) {
let sequence = self
.updates
.lock()
.unwrap()
.as_ref()
.map(X509SourceUpdates::last);
self.update_sequences.lock().unwrap().push(sequence);
}
}
impl MetricsRecorder for OrderingMetricsRecorder {
fn record_update(&self) {
self.record_observed_update();
}
fn record_reconnect(&self) {}
fn record_error(&self, _kind: MetricsErrorKind) {}
}
struct PanickingOrderingMetricsRecorder {
inner: OrderingMetricsRecorder,
}
impl PanickingOrderingMetricsRecorder {
fn new() -> Self {
Self {
inner: OrderingMetricsRecorder::new(),
}
}
fn set_updates(&self, updates: X509SourceUpdates) {
self.inner.set_updates(updates);
}
fn update_sequences(&self) -> Vec<Option<u64>> {
self.inner.update_sequences()
}
}
impl MetricsRecorder for PanickingOrderingMetricsRecorder {
fn record_update(&self) {
self.inner.record_observed_update();
panic!("intentional metrics panic for update ordering test");
}
fn record_reconnect(&self) {}
fn record_error(&self, _kind: MetricsErrorKind) {}
}
#[test]
fn test_apply_update_notifies_before_recording_success_metrics() {
let metrics = Arc::new(OrderingMetricsRecorder::new());
let source = X509Source::new_for_test(
Arc::new(X509Context::new([], Arc::new(X509BundleSet::new()))),
ReconnectConfig::default(),
ResourceLimits::default(),
Some(Arc::<OrderingMetricsRecorder>::clone(&metrics)),
None,
);
metrics.set_updates(source.updated());
let result = source.inner.apply_update(test_x509_context());
assert_eq!(
result.expect("valid X.509 update should be accepted"),
ApplyUpdateResult::Applied
);
assert_eq!(metrics.update_sequences(), vec![Some(1)]);
assert_eq!(source.updated().last(), 1);
assert_eq!(source.x509_context().unwrap().svids().len(), 1);
}
#[test]
fn test_apply_update_skips_notify_for_unchanged_context() {
let metrics = Arc::new(OrderingMetricsRecorder::new());
let source = X509Source::new_for_test(
test_x509_context(),
ReconnectConfig::default(),
ResourceLimits::default(),
Some(Arc::<OrderingMetricsRecorder>::clone(&metrics)),
None,
);
metrics.set_updates(source.updated());
let result = source.inner.apply_update(test_x509_context());
assert_eq!(
result.expect("re-delivery of an unchanged context should be accepted"),
ApplyUpdateResult::Unchanged
);
assert_eq!(
source.updated().last(),
0,
"no notification should be sent for an unchanged context"
);
assert!(
metrics.update_sequences().is_empty(),
"record_update should not be called for an unchanged context"
);
let mut rotated_bundle_set = X509BundleSet::new();
rotated_bundle_set.add_bundle(X509Bundle::new(TrustDomain::new("example.org").unwrap()));
let rotated_ctx = Arc::new(X509Context::new(
test_x509_context().svids().to_vec(),
Arc::new(rotated_bundle_set),
));
let result = source
.inner
.apply_update(rotated_ctx)
.expect("valid rotation should be accepted");
assert_eq!(result, ApplyUpdateResult::Applied);
assert_eq!(
source.updated().last(),
1,
"a genuine rotation should still notify"
);
assert_eq!(metrics.update_sequences(), vec![Some(1)]);
}
#[test]
fn same_material_for_update_ignores_svid_order() {
let cert = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid_a = Arc::new(
X509Svid::parse_from_der_with_hint(cert, key, Some("internal".into())).unwrap(),
);
let svid_b = Arc::new(
X509Svid::parse_from_der_with_hint(cert, key, Some("external".into())).unwrap(),
);
assert_ne!(
svid_a.as_ref(),
svid_b.as_ref(),
"fixture SVIDs must differ so order-sensitive equality would fail"
);
let mut bundle_set = X509BundleSet::new();
bundle_set.add_bundle(X509Bundle::new(TrustDomain::new("example.org").unwrap()));
let bundle_set = Arc::new(bundle_set);
let ordered = X509Context::new(
[Arc::clone(&svid_a), Arc::clone(&svid_b)],
Arc::clone(&bundle_set),
);
let reversed = X509Context::new([svid_b, svid_a], bundle_set);
assert_ne!(
ordered, reversed,
"public X509Context equality remains order-sensitive"
);
assert!(same_material_for_update(&ordered, &reversed));
}
#[test]
fn test_apply_update_notifies_for_intermediate_chain_change() {
use crate::cert::parsing::to_certificate_vec;
let chain = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let certs = to_certificate_vec(chain).unwrap();
let first_cert = certs
.first()
.expect("fixture must include leaf certificate");
assert!(
certs.len() > 1,
"fixture must include intermediates to exercise chain-length differences"
);
let full = Arc::new(X509Svid::parse_from_der(chain, key).unwrap());
let leaf_only = Arc::new(X509Svid::parse_from_der(first_cert.as_bytes(), key).unwrap());
assert_ne!(
full.as_ref(),
leaf_only.as_ref(),
"full-chain and leaf-only SVIDs must differ under structural equality"
);
let mut bundle_set = X509BundleSet::new();
bundle_set.add_bundle(X509Bundle::new(TrustDomain::new("example.org").unwrap()));
let bundle_set = Arc::new(bundle_set);
let leaf_only_ctx = Arc::new(X509Context::new(
[Arc::clone(&leaf_only)],
Arc::clone(&bundle_set),
));
let with_chain = Arc::new(X509Context::new([full], bundle_set));
assert!(!same_material_for_update(
leaf_only_ctx.as_ref(),
with_chain.as_ref()
));
let metrics = Arc::new(OrderingMetricsRecorder::new());
let source = X509Source::new_for_test(
leaf_only_ctx,
ReconnectConfig::default(),
ResourceLimits::default(),
Some(Arc::<OrderingMetricsRecorder>::clone(&metrics)),
None,
);
metrics.set_updates(source.updated());
source
.inner
.apply_update(with_chain)
.expect("intermediate chain change should be accepted");
assert_eq!(
source.updated().last(),
1,
"intermediate chain changes must notify because TLS uses the full chain"
);
assert_eq!(metrics.update_sequences(), vec![Some(1)]);
assert_eq!(source.svid().unwrap().cert_chain().len(), certs.len());
}
#[test]
fn same_material_for_update_ignores_order_when_chain_differs() {
use crate::cert::parsing::to_certificate_vec;
let chain = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let certs = to_certificate_vec(chain).unwrap();
let first_cert = certs
.first()
.expect("fixture must include leaf certificate");
let full = Arc::new(X509Svid::parse_from_der(chain, key).unwrap());
let leaf_only = Arc::new(X509Svid::parse_from_der(first_cert.as_bytes(), key).unwrap());
assert_ne!(
full.as_ref(),
leaf_only.as_ref(),
"fixtures must differ only by presented chain"
);
let mut bundle_set = X509BundleSet::new();
bundle_set.add_bundle(X509Bundle::new(TrustDomain::new("example.org").unwrap()));
let bundle_set = Arc::new(bundle_set);
let ordered = X509Context::new(
[Arc::clone(&full), Arc::clone(&leaf_only)],
Arc::clone(&bundle_set),
);
let reversed = X509Context::new([leaf_only, full], bundle_set);
assert_ne!(
ordered, reversed,
"public X509Context equality remains order-sensitive"
);
assert!(
same_material_for_update(&ordered, &reversed),
"equivalent SVID multisets should compare equal even when matching requires the chain"
);
}
#[test]
fn same_material_for_update_ignores_bundle_authority_order() {
use crate::cert::parsing::to_certificate_vec;
let cert = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid = Arc::new(X509Svid::parse_from_der(cert, key).unwrap());
let certs = to_certificate_vec(cert).unwrap();
let [first, second, ..] = certs.as_slice() else {
panic!("fixture must include multiple certificates to exercise authority order");
};
let trust_domain = TrustDomain::new("example.org").unwrap();
let mut bundle_a = X509Bundle::new(trust_domain.clone());
bundle_a
.add_authority(first.as_bytes())
.expect("authority should parse");
bundle_a
.add_authority(second.as_bytes())
.expect("second authority should parse");
let mut bundle_b = X509Bundle::new(trust_domain);
bundle_b
.add_authority(second.as_bytes())
.expect("authority should parse");
bundle_b
.add_authority(first.as_bytes())
.expect("second authority should parse");
assert_ne!(bundle_a, bundle_b);
let mut bundle_set_a = X509BundleSet::new();
bundle_set_a.add_bundle(bundle_a);
let mut bundle_set_b = X509BundleSet::new();
bundle_set_b.add_bundle(bundle_b);
let ctx_a = X509Context::new([Arc::clone(&svid)], Arc::new(bundle_set_a));
let ctx_b = X509Context::new([svid], Arc::new(bundle_set_b));
assert!(same_material_for_update(&ctx_a, &ctx_b));
}
#[test]
fn test_apply_update_skips_notify_for_reordered_svids() {
let cert = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid_a = Arc::new(
X509Svid::parse_from_der_with_hint(cert, key, Some("internal".into())).unwrap(),
);
let svid_b = Arc::new(
X509Svid::parse_from_der_with_hint(cert, key, Some("external".into())).unwrap(),
);
let mut bundle_set = X509BundleSet::new();
bundle_set.add_bundle(X509Bundle::new(TrustDomain::new("example.org").unwrap()));
let bundle_set = Arc::new(bundle_set);
let initial = Arc::new(X509Context::new(
[Arc::clone(&svid_a), Arc::clone(&svid_b)],
Arc::clone(&bundle_set),
));
let reordered = Arc::new(X509Context::new([svid_b, svid_a], bundle_set));
let metrics = Arc::new(OrderingMetricsRecorder::new());
let source = X509Source::new_for_test(
initial,
ReconnectConfig::default(),
ResourceLimits::default(),
Some(Arc::<OrderingMetricsRecorder>::clone(&metrics)),
None,
);
metrics.set_updates(source.updated());
let result = source
.inner
.apply_update(reordered)
.expect("reordered but equivalent material should be accepted");
assert_eq!(result, ApplyUpdateResult::Unchanged);
assert_eq!(source.updated().last(), 0);
assert!(metrics.update_sequences().is_empty());
}
#[test]
fn test_apply_update_publishes_and_notifies_before_panicking_success_metrics() {
let metrics = Arc::new(PanickingOrderingMetricsRecorder::new());
let source = X509Source::new_for_test(
Arc::new(X509Context::new([], Arc::new(X509BundleSet::new()))),
ReconnectConfig::default(),
ResourceLimits::default(),
Some(Arc::<PanickingOrderingMetricsRecorder>::clone(&metrics)),
None,
);
metrics.set_updates(source.updated());
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
source.inner.apply_update(test_x509_context())
}));
result.expect_err("metrics recorder should panic after notification");
assert_eq!(metrics.update_sequences(), vec![Some(1)]);
assert_eq!(source.updated().last(), 1);
assert_eq!(source.x509_context().unwrap().svids().len(), 1);
}
#[test]
fn test_metrics_recorded_exactly_once_per_rejected_update() {
use super::super::builder::{ReconnectConfig, ResourceLimits};
use crate::workload_api::x509_context::X509Context;
use crate::{TrustDomain, X509Bundle, X509BundleSet, X509Svid};
use std::sync::Arc;
let cert_bytes = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key_bytes = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid = Arc::new(X509Svid::parse_from_der(cert_bytes, key_bytes).unwrap());
let metrics = Arc::new(TestMetricsRecorder::new());
let limits = ResourceLimits {
max_svids: Some(100),
max_bundles: Some(0), max_bundle_der_bytes: Some(1000),
};
let trust_domain = TrustDomain::new("example.org").unwrap();
let bundle = X509Bundle::new(trust_domain);
let mut bundle_set = X509BundleSet::new();
bundle_set.add_bundle(bundle);
let ctx = X509Context::new([svid], Arc::new(bundle_set));
let source = {
let metrics = Arc::clone(&metrics);
X509Source::new_for_test(
Arc::new(X509Context::new([], Arc::new(X509BundleSet::new()))),
ReconnectConfig::default(),
limits,
Some(metrics),
None,
)
};
let result = source.inner.apply_update(Arc::new(ctx));
assert!(matches!(
result,
Err(X509SourceError::ResourceLimitExceeded {
kind: super::super::errors::LimitKind::MaxBundles,
..
})
));
assert_eq!(metrics.count(MetricsErrorKind::LimitMaxBundles), 1);
assert_eq!(metrics.count(MetricsErrorKind::UpdateRejected), 1);
assert_eq!(metrics.count(MetricsErrorKind::LimitMaxSvids), 0);
assert_eq!(metrics.count(MetricsErrorKind::LimitMaxBundleDerBytes), 0);
}
#[test]
fn test_apply_update_rejects_expired_svid_and_retains_previous_bundles() {
use super::super::builder::{ReconnectConfig, ResourceLimits};
use crate::workload_api::x509_context::X509Context;
use crate::{TrustDomain, X509Bundle, X509BundleSet, X509Svid};
use std::sync::Arc;
let good_trust_domain = TrustDomain::new("good.example.org").unwrap();
let mut good_bundle_set = X509BundleSet::new();
good_bundle_set.add_bundle(X509Bundle::new(good_trust_domain.clone()));
let initial_ctx = Arc::new(X509Context::new(
test_x509_context().svids().to_vec(),
Arc::new(good_bundle_set),
));
let source = X509Source::new_for_test(
Arc::clone(&initial_ctx),
ReconnectConfig::default(),
ResourceLimits::default(),
None,
None,
);
let expired_cert_bytes =
include_bytes!("../../tests/testdata/svid/x509/expired-svid-chain.der");
let expired_key_bytes = include_bytes!("../../tests/testdata/svid/x509/expired-key.der");
let expired_svid = Arc::new(
X509Svid::parse_from_der(expired_cert_bytes, expired_key_bytes)
.expect("expired fixture should parse as an X509-SVID"),
);
let new_trust_domain = TrustDomain::new("new.example.org").unwrap();
let mut new_bundle_set = X509BundleSet::new();
new_bundle_set.add_bundle(X509Bundle::new(new_trust_domain.clone()));
let expired_ctx = Arc::new(X509Context::new([expired_svid], Arc::new(new_bundle_set)));
let result = source.inner.apply_update(expired_ctx);
assert!(matches!(result, Err(X509SourceError::NoSuitableSvid)));
let retained_ctx = source.x509_context().unwrap();
assert_eq!(retained_ctx.svids(), initial_ctx.svids());
let retained_bundles = source.bundle_set().unwrap();
assert!(retained_bundles.get(&good_trust_domain).is_some());
assert!(retained_bundles.get(&new_trust_domain).is_none());
assert_eq!(source.updated().last(), 0);
}
#[test]
fn test_new_with_normalizes_reconnect_config() {
use super::super::builder::{ReconnectConfig, ResourceLimits};
use crate::workload_api::x509_context::X509Context;
use crate::{X509BundleSet, X509Svid};
use std::sync::Arc;
use std::time::Duration;
let cert_bytes = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key_bytes = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid = Arc::new(X509Svid::parse_from_der(cert_bytes, key_bytes).unwrap());
let ctx = X509Context::new([svid], Arc::new(X509BundleSet::new()));
let inverted_reconnect = ReconnectConfig {
min_backoff: Duration::from_secs(10),
max_backoff: Duration::from_secs(1),
};
let source = X509Source::new_for_test(
Arc::new(ctx),
inverted_reconnect,
ResourceLimits::default(),
None,
None,
);
assert_eq!(source.inner.reconnect.min_backoff, Duration::from_secs(1));
assert_eq!(source.inner.reconnect.max_backoff, Duration::from_secs(10));
}
#[test]
fn test_initial_sync_validation_records_correct_metrics() {
use super::super::builder::ResourceLimits;
use super::super::limits::validate_context;
use crate::workload_api::x509_context::X509Context;
use crate::{TrustDomain, X509Bundle, X509BundleSet, X509Svid};
use std::sync::Arc;
let cert_bytes = include_bytes!("../../tests/testdata/svid/x509/1-svid-chain.der");
let key_bytes = include_bytes!("../../tests/testdata/svid/x509/1-key.der");
let svid = Arc::new(X509Svid::parse_from_der(cert_bytes, key_bytes).unwrap());
let metrics = Arc::new(TestMetricsRecorder::new());
let limits = ResourceLimits {
max_svids: Some(100),
max_bundles: Some(0), max_bundle_der_bytes: Some(1000),
};
let trust_domain = TrustDomain::new("example.org").unwrap();
let bundle = X509Bundle::new(trust_domain);
let mut bundle_set = X509BundleSet::new();
bundle_set.add_bundle(bundle);
let ctx = X509Context::new([svid], Arc::new(bundle_set));
let result = validate_context(
&ctx,
None, limits,
Some(metrics.as_ref()),
);
assert!(matches!(
result,
Err(X509SourceError::ResourceLimitExceeded {
kind: super::super::errors::LimitKind::MaxBundles,
..
})
));
assert_eq!(metrics.count(MetricsErrorKind::LimitMaxBundles), 1);
assert_eq!(metrics.count(MetricsErrorKind::UpdateRejected), 0);
assert_eq!(metrics.count(MetricsErrorKind::LimitMaxSvids), 0);
assert_eq!(metrics.count(MetricsErrorKind::LimitMaxBundleDerBytes), 0);
}
#[test]
fn test_resource_limits_unlimited() {
use super::super::builder::ResourceLimits;
let unlimited = ResourceLimits::unlimited();
assert_eq!(unlimited.max_svids, None);
assert_eq!(unlimited.max_bundles, None);
assert_eq!(unlimited.max_bundle_der_bytes, None);
}
#[test]
fn test_resource_limits_default_limits() {
use super::super::builder::ResourceLimits;
let limits = ResourceLimits::default_limits();
assert_eq!(limits.max_svids, Some(100));
assert_eq!(limits.max_bundles, Some(200));
assert_eq!(limits.max_bundle_der_bytes, Some(4 * 1024 * 1024)); }
#[test]
fn test_resource_limits_mixed() {
use super::super::builder::ResourceLimits;
let mixed = ResourceLimits {
max_svids: Some(50),
max_bundles: None, max_bundle_der_bytes: Some(1024 * 1024), };
assert_eq!(mixed.max_svids, Some(50));
assert_eq!(mixed.max_bundles, None);
assert_eq!(mixed.max_bundle_der_bytes, Some(1024 * 1024));
}
}