Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 43 additions & 12 deletions crates/attestation/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ use std::{
fmt::{self, Display, Formatter},
io::Read,
net::IpAddr,
sync::{Arc, RwLock, RwLockReadGuard},
time::{Duration, SystemTime, UNIX_EPOCH},
};

Expand Down Expand Up @@ -325,8 +326,9 @@ impl AttestationGenerator {
/// Allows remote attestations to be verified
#[derive(Clone, Debug)]
pub struct AttestationVerifier {
/// The measurement policy with accepted values and attestation types
measurement_policy: MeasurementPolicy,
/// The measurement policy with accepted values and attestation types,
/// shared between clones
measurement_policy: Arc<RwLock<MeasurementPolicy>>,
/// Whether to write quotes to files on disk
dump_dcap_quotes: bool,
#[cfg(feature = "azure")]
Expand Down Expand Up @@ -409,7 +411,7 @@ impl AttestationVerifier {
});

Self {
measurement_policy: builder.measurement_policy,
measurement_policy: Arc::new(RwLock::new(builder.measurement_policy)),
dump_dcap_quotes: builder.dump_dcap_quotes,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: builder.override_azure_outdated_tcb,
Expand All @@ -433,7 +435,7 @@ impl AttestationVerifier {
/// and will reject if one is given
pub fn expect_none() -> Self {
Self {
measurement_policy: MeasurementPolicy::expect_none(),
measurement_policy: Arc::new(RwLock::new(MeasurementPolicy::expect_none())),
dump_dcap_quotes: false,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
Expand All @@ -446,7 +448,7 @@ impl AttestationVerifier {
#[cfg(any(test, feature = "mock"))]
pub fn mock() -> Self {
Self {
measurement_policy: MeasurementPolicy::mock(),
measurement_policy: Arc::new(RwLock::new(MeasurementPolicy::mock())),
dump_dcap_quotes: false,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
Expand All @@ -459,7 +461,7 @@ impl AttestationVerifier {
#[cfg(any(test, feature = "mock"))]
pub fn mock_with_pccs(pccs_url: String) -> Self {
Self {
measurement_policy: MeasurementPolicy::mock(),
measurement_policy: Arc::new(RwLock::new(MeasurementPolicy::mock())),
dump_dcap_quotes: false,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
Expand Down Expand Up @@ -551,7 +553,7 @@ impl AttestationVerifier {
.attestation_evidence
.as_ref()
.map(|evidence| evidence.platform.clone());
self.measurement_policy.check_measurement_with_gcp_cache(
self.measurement_policy_read().check_measurement_with_gcp_cache(
&measurements,
platform_metadata.as_ref(),
Some(&self.known_gcp_firmware),
Expand Down Expand Up @@ -628,7 +630,7 @@ impl AttestationVerifier {
.attestation_evidence
.as_ref()
.map(|evidence| evidence.platform.clone());
self.measurement_policy.check_measurement_with_gcp_cache(
self.measurement_policy_read().check_measurement_with_gcp_cache(
&measurements,
platform_metadata.as_ref(),
Some(&self.known_gcp_firmware),
Expand All @@ -640,12 +642,23 @@ impl AttestationVerifier {

/// Whether we allow no remote attestation
pub fn has_remote_attestation(&self) -> bool {
self.measurement_policy.has_remote_attestation()
self.measurement_policy_read().has_remote_attestation()
}

/// Returns a snapshot of the measurement policy currently in use.
pub fn measurement_policy(&self) -> MeasurementPolicy {
self.measurement_policy_read().clone()
}

/// Replaces the measurement policy used by this verifier and all of its
/// clones.
pub fn set_measurement_policy(&self, measurement_policy: MeasurementPolicy) {
*self.measurement_policy.write().unwrap_or_else(|poisoned| poisoned.into_inner()) =
measurement_policy;
}

/// Returns the measurement policy used
pub fn measurement_policy(&self) -> &MeasurementPolicy {
&self.measurement_policy
fn measurement_policy_read(&self) -> RwLockReadGuard<'_, MeasurementPolicy> {
self.measurement_policy.read().unwrap_or_else(|poisoned| poisoned.into_inner())
}
}

Expand Down Expand Up @@ -835,4 +848,22 @@ mod tests {

assert!(result.is_ok(), "expected sync mock verification to succeed: {result:?}");
}

#[test]
fn measurement_policy_can_be_updated_between_verification_attempts() {
let verifier =
AttestationVerifier::builder(MeasurementPolicy::tdx()).with_no_internal_pccs().build();
let verifier_clone = verifier.clone();
let message = AttestationExchangeMessage::without_attestation();
let input_data = [0; 64];

assert!(matches!(
verifier.verify_attestation_sync(message.clone(), input_data),
Err(AttestationError::AttestationTypeNotAccepted)
));

verifier_clone.set_measurement_policy(MeasurementPolicy::expect_none());

assert!(matches!(verifier.verify_attestation_sync(message, input_data), Ok(None)));
}
}
Loading