Skip to content
Open
Show file tree
Hide file tree
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
236 changes: 221 additions & 15 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, RwLockWriteGuard},
time::{Duration, SystemTime, UNIX_EPOCH},
};

Expand All @@ -22,7 +23,11 @@ use pccs::{Pccs, PccsError};
use serde::{Deserialize, Serialize};
use thiserror::Error;

use crate::{dcap::DcapVerificationError, gcp::GcpFirmwareCache, measurements::MeasurementPolicy};
use crate::{
dcap::DcapVerificationError,
gcp::GcpFirmwareCache,
measurements::{MeasurementFormatError, MeasurementPolicy},
};

#[cfg(test)]
static TEST_CRYPTO_PROVIDER: OnceLock<()> = OnceLock::new();
Expand Down Expand Up @@ -325,8 +330,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<MeasurementPolicyState>>,
/// Whether to write quotes to files on disk
dump_dcap_quotes: bool,
#[cfg(feature = "azure")]
Expand All @@ -338,12 +344,30 @@ pub struct AttestationVerifier {
internal_pccs: Option<Pccs>,
/// Cached GCP firmware blobs indexed by MRTD
known_gcp_firmware: gcp::GcpFirmwareCache,
/// Dynamic measurement policy to re-fetch from file or URL
dynamic_measurement_policy: Option<String>,
}

/// Measurement policy together with a generation number used to track
/// changes
#[derive(Clone, Debug)]
struct MeasurementPolicyState {
policy: MeasurementPolicy,
generation: u64,
}

impl MeasurementPolicyState {
fn new(policy: MeasurementPolicy) -> Self {
Self { policy, generation: 0 }
}
}

/// Options used to construct an [AttestationVerifier]
pub struct AttestationVerifierBuilder {
/// The measurement policy with accepted values and attestation types
measurement_policy: MeasurementPolicy,
/// A dynamic measurement policy file or URL
dynamic_measurement_policy: Option<String>,
/// A PCCS service to use - defaults to Intel PCS
pccs_url: Option<String>,
dump_dcap_quotes: bool,
Expand Down Expand Up @@ -396,6 +420,17 @@ impl AttestationVerifierBuilder {
self.pccs_url = Some(pccs_url);
self
}

/// Re-fetch the measurement policy from this file or URL after a
/// measurement mismatch.
///
/// Both asynchronous and synchronous verification perform one retry
/// with the refreshed policy. Synchronous URL refreshes block for
/// up to ten seconds.
pub fn with_dynamic_measurements_file_or_url(mut self, file_or_url: String) -> Self {
self.dynamic_measurement_policy = Some(file_or_url);
self
}
}

impl AttestationVerifier {
Expand All @@ -409,12 +444,15 @@ impl AttestationVerifier {
});

Self {
measurement_policy: builder.measurement_policy,
measurement_policy: Arc::new(RwLock::new(MeasurementPolicyState::new(
builder.measurement_policy,
))),
dump_dcap_quotes: builder.dump_dcap_quotes,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: builder.override_azure_outdated_tcb,
internal_pccs,
known_gcp_firmware: GcpFirmwareCache::new(),
dynamic_measurement_policy: builder.dynamic_measurement_policy,
}
}

Expand All @@ -426,45 +464,55 @@ impl AttestationVerifier {
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
internal_pccs_prewarm: Some(true),
dynamic_measurement_policy: None,
}
}

/// Create an [AttestationVerifier] which will only allow no attestation
/// and will reject if one is given
pub fn expect_none() -> Self {
Self {
measurement_policy: MeasurementPolicy::expect_none(),
measurement_policy: Arc::new(RwLock::new(MeasurementPolicyState::new(
MeasurementPolicy::expect_none(),
))),
dump_dcap_quotes: false,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
internal_pccs: None,
known_gcp_firmware: gcp::GcpFirmwareCache::new(),
dynamic_measurement_policy: None,
}
}

/// Expect mock measurements used in tests
#[cfg(any(test, feature = "mock"))]
pub fn mock() -> Self {
Self {
measurement_policy: MeasurementPolicy::mock(),
measurement_policy: Arc::new(RwLock::new(MeasurementPolicyState::new(
MeasurementPolicy::mock(),
))),
dump_dcap_quotes: false,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
internal_pccs: None,
known_gcp_firmware: gcp::GcpFirmwareCache::new(),
dynamic_measurement_policy: None,
}
}

/// Expect mock measurements used in tests, and use a PCCS
#[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(MeasurementPolicyState::new(
MeasurementPolicy::mock(),
))),
dump_dcap_quotes: false,
#[cfg(feature = "azure")]
override_azure_outdated_tcb: false,
internal_pccs: Some(Pccs::new(Some(pccs_url))),
known_gcp_firmware: gcp::GcpFirmwareCache::new(),
dynamic_measurement_policy: None,
}
}

Expand Down Expand Up @@ -551,11 +599,31 @@ impl AttestationVerifier {
.attestation_evidence
.as_ref()
.map(|evidence| evidence.platform.clone());
self.measurement_policy.check_measurement_with_gcp_cache(

let policy_state = self.measurement_policy_read().clone();
let policy_check = policy_state.policy.check_measurement_with_gcp_cache(
&measurements,
platform_metadata.as_ref(),
Some(&self.known_gcp_firmware),
)?;
);

if let Err(err) = policy_check {
// If this fails, and we have dynamic measurement policy, re-retrieve our
// measurement policy, then check the policy a second time
if let Some(file_or_url) = &self.dynamic_measurement_policy {
let new_measurement_policy =
MeasurementPolicy::from_file_or_url(file_or_url.to_string()).await?;
let measurement_policy =
self.set_measurement_policy(new_measurement_policy, policy_state.generation);
measurement_policy.check_measurement_with_gcp_cache(
&measurements,
platform_metadata.as_ref(),
Some(&self.known_gcp_firmware),
)?;
} else {
return Err(err);
}
}

tracing::debug!("Verification successful");
Ok(Some(measurements))
Expand Down Expand Up @@ -628,24 +696,71 @@ impl AttestationVerifier {
.attestation_evidence
.as_ref()
.map(|evidence| evidence.platform.clone());
self.measurement_policy.check_measurement_with_gcp_cache(
let policy_state = self.measurement_policy_read().clone();
let policy_check = policy_state.policy.check_measurement_with_gcp_cache(
&measurements,
platform_metadata.as_ref(),
Some(&self.known_gcp_firmware),
)?;
);

if let Err(err) = policy_check {
if let Some(file_or_url) = &self.dynamic_measurement_policy {
let new_measurement_policy =
MeasurementPolicy::from_file_or_url_sync(file_or_url.to_string())?;
let measurement_policy =
self.set_measurement_policy(new_measurement_policy, policy_state.generation);
measurement_policy.check_measurement_with_gcp_cache(
&measurements,
platform_metadata.as_ref(),
Some(&self.known_gcp_firmware),
)?;
} else {
return Err(err);
}
}

tracing::debug!("Verification successful");
Ok(Some(measurements))
}

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

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

/// Whether this verifier automatically refreshes its measurement policy
/// after a mismatch.
pub fn has_dynamic_measurement_policy(&self) -> bool {
self.dynamic_measurement_policy.is_some()
}

/// Replaces the measurement policy used by this verifier and all of its
/// clones if it has not changed since `expected_generation` was
/// observed.
pub(crate) fn set_measurement_policy(
&self,
measurement_policy: MeasurementPolicy,
expected_generation: u64,
) -> MeasurementPolicy {
let mut state = self.measurement_policy_write();
if state.generation == expected_generation {
state.policy = measurement_policy;
state.generation = state.generation.wrapping_add(1);
}
state.policy.clone()
}

fn measurement_policy_read(&self) -> RwLockReadGuard<'_, MeasurementPolicyState> {
self.measurement_policy.read().unwrap_or_else(|poisoned| poisoned.into_inner())
}

fn measurement_policy_write(&self) -> RwLockWriteGuard<'_, MeasurementPolicyState> {
self.measurement_policy.write().unwrap_or_else(|poisoned| poisoned.into_inner())
}
}

Expand Down Expand Up @@ -773,6 +888,8 @@ pub enum AttestationError {
AttestationTypeNotAccepted,
#[error("Measurements not accepted")]
MeasurementsNotAccepted,
#[error("Failed to refresh measurement policy: {0}")]
MeasurementPolicyRefresh(#[from] MeasurementFormatError),
#[cfg(feature = "azure")]
#[error("MAA: {0}")]
Maa(#[from] azure::MaaError),
Expand Down Expand Up @@ -835,4 +952,93 @@ 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)
));

let generation = verifier_clone.measurement_policy_read().generation;
verifier_clone.set_measurement_policy(MeasurementPolicy::expect_none(), generation);

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

#[test]
fn stale_measurement_policy_refresh_does_not_overwrite_newer_policy() {
let verifier = AttestationVerifier::expect_none();
let stale_generation = verifier.measurement_policy_read().generation;

verifier.set_measurement_policy(MeasurementPolicy::tdx(), stale_generation);
let installed_policy =
verifier.set_measurement_policy(MeasurementPolicy::expect_none(), stale_generation);

assert!(installed_policy.has_remote_attestation());
assert!(verifier.has_remote_attestation());
}

#[tokio::test]
async fn dynamic_measurement_policy_refetches_on_mismatch() {
let temp_dir = tempfile::tempdir().unwrap();
let policy_path = temp_dir.path().join("measurements.json");
tokio::fs::write(&policy_path, br#"[{"attestation_type":"none"}]"#).await.unwrap();

let initial_policy = MeasurementPolicy::from_file(policy_path.clone()).await.unwrap();
let verifier = AttestationVerifier::builder(initial_policy)
.with_no_internal_pccs()
.with_dynamic_measurements_file_or_url(policy_path.to_string_lossy().into_owned())
.build();

let input_data = [7u8; 64];
let quote = dcap::create_dcap_attestation(input_data).unwrap();
let attestation = AttestationEvidence {
quote,
platform: mock_platform_metadata(AttestationType::DcapTdx).unwrap(),
};
let measurements = measurements::mock_dcap_measurements();

assert!(verifier.measurement_policy().check_measurement(&measurements, None).is_err());

tokio::fs::write(&policy_path, br#"[{"attestation_type":"dcap-tdx"}]"#).await.unwrap();

verifier.verify_attestation(attestation.into(), input_data).await.unwrap();

assert!(verifier.measurement_policy().check_measurement(&measurements, None).is_ok());
}

#[tokio::test]
async fn sync_verification_refetches_dynamic_measurement_policy_on_mismatch() {
let temp_dir = tempfile::tempdir().unwrap();
let policy_path = temp_dir.path().join("measurements.json");
std::fs::write(&policy_path, br#"[{"attestation_type":"none"}]"#).unwrap();

let policy_source = policy_path.to_string_lossy().into_owned();
let initial_policy =
MeasurementPolicy::from_file_or_url_sync(policy_source.clone()).unwrap();
let mock_pcs_server = spawn_mock_pcs_server(MockPcsConfig::default()).await.unwrap();
let verifier = AttestationVerifier::builder(initial_policy)
.pccs_url(mock_pcs_server.base_url.clone())
.with_dynamic_measurements_file_or_url(policy_source)
.build();
verifier.ready().await.unwrap();

let input_data = [7u8; 64];
let quote = dcap::create_dcap_attestation(input_data).unwrap();
let attestation = AttestationEvidence {
quote,
platform: mock_platform_metadata(AttestationType::DcapTdx).unwrap(),
};

std::fs::write(&policy_path, br#"[{"attestation_type":"dcap-tdx"}]"#).unwrap();

verifier.verify_attestation_sync(attestation.into(), input_data).unwrap();
}
}
Loading
Loading