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
233 changes: 145 additions & 88 deletions crates/attestation/src/measurements.rs
Original file line number Diff line number Diff line change
Expand Up @@ -83,21 +83,89 @@ fn parse_azure_pcr_index(value: &str) -> Result<u32, MeasurementFormatError> {
Ok(index)
}

#[derive(Clone, PartialEq)]
pub struct DcapMeasurements {
pub mrtd: [u8; 48],
pub rtmr0: [u8; 48],
pub rtmr1: [u8; 48],
pub rtmr2: [u8; 48],
pub rtmr3: [u8; 48],
}

impl DcapMeasurements {
pub fn new(
mrtd: [u8; 48],
rtmr0: [u8; 48],
rtmr1: [u8; 48],
rtmr2: [u8; 48],
rtmr3: [u8; 48],
) -> Self {
Self { mrtd, rtmr0, rtmr1, rtmr2, rtmr3 }
}

fn from_map(
mut measurements: HashMap<DcapMeasurementRegister, [u8; 48]>,
) -> Result<Self, MeasurementFormatError> {
Ok(Self {
mrtd: measurements
.remove(&DcapMeasurementRegister::MRTD)
.ok_or_else(|| MeasurementFormatError::MissingValue("MRTD".to_string()))?,
rtmr0: measurements
.remove(&DcapMeasurementRegister::RTMR0)
.ok_or_else(|| MeasurementFormatError::MissingValue("RTMR0".to_string()))?,
rtmr1: measurements
.remove(&DcapMeasurementRegister::RTMR1)
.ok_or_else(|| MeasurementFormatError::MissingValue("RTMR1".to_string()))?,
rtmr2: measurements
.remove(&DcapMeasurementRegister::RTMR2)
.ok_or_else(|| MeasurementFormatError::MissingValue("RTMR2".to_string()))?,
rtmr3: measurements
.remove(&DcapMeasurementRegister::RTMR3)
.ok_or_else(|| MeasurementFormatError::MissingValue("RTMR3".to_string()))?,
})
}

fn iter(&self) -> impl Iterator<Item = (DcapMeasurementRegister, &[u8; 48])> {
[
(DcapMeasurementRegister::MRTD, &self.mrtd),
(DcapMeasurementRegister::RTMR0, &self.rtmr0),
(DcapMeasurementRegister::RTMR1, &self.rtmr1),
(DcapMeasurementRegister::RTMR2, &self.rtmr2),
(DcapMeasurementRegister::RTMR3, &self.rtmr3),
]
.into_iter()
}

fn get(&self, register: &DcapMeasurementRegister) -> &[u8; 48] {
match register {
DcapMeasurementRegister::MRTD => &self.mrtd,
DcapMeasurementRegister::RTMR0 => &self.rtmr0,
DcapMeasurementRegister::RTMR1 => &self.rtmr1,
DcapMeasurementRegister::RTMR2 => &self.rtmr2,
DcapMeasurementRegister::RTMR3 => &self.rtmr3,
}
}
}

impl fmt::Debug for DcapMeasurements {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
DcapHexDebug(self).fmt(f)
}
}

/// Represents a set of measurements values for one of the supported CVM
/// platforms
#[derive(Clone, PartialEq)]
pub enum MultiMeasurements {
Dcap(HashMap<DcapMeasurementRegister, [u8; 48]>),
Dcap(DcapMeasurements),
Azure(HashMap<u32, [u8; 32]>),
NoAttestation,
}

impl fmt::Debug for MultiMeasurements {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
match self {
Self::Dcap(measurements) => {
f.debug_tuple("DCAP").field(&DcapHexDebug(measurements)).finish()
}
Self::Dcap(measurements) => f.debug_tuple("DCAP").field(measurements).finish(),
Self::Azure(measurements) => {
f.debug_tuple("Azure").field(&AzureHexDebug(measurements)).finish()
}
Expand All @@ -107,17 +175,14 @@ impl fmt::Debug for MultiMeasurements {
}

/// Used to display DCAP measurements as hex
struct DcapHexDebug<'a>(&'a HashMap<DcapMeasurementRegister, [u8; 48]>);
struct DcapHexDebug<'a>(&'a DcapMeasurements);

impl fmt::Debug for DcapHexDebug<'_> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let mut entries: Vec<_> = self.0.iter().collect();
entries.sort_by_key(|(register, _)| (*register).clone() as u8);

let mut map = f.debug_map();
for (register, value) in entries {
for (register, value) in self.0.iter() {
let hex_value = hex::encode(value);
map.entry(register, &hex_value);
map.entry(&register, &hex_value);
}
map.finish()
}
Expand Down Expand Up @@ -155,7 +220,7 @@ impl MultiMeasurements {
let measurements_map = match self {
MultiMeasurements::Dcap(dcap_measurements) => dcap_measurements
.iter()
.map(|(register, value)| ((register.clone() as u8).to_string(), hex::encode(value)))
.map(|(register, value)| ((register as u8).to_string(), hex::encode(value)))
.collect(),
MultiMeasurements::Azure(azure_measurements) => azure_measurements
.iter()
Expand Down Expand Up @@ -201,7 +266,7 @@ impl MultiMeasurements {
))
})
.collect::<Result<_, MeasurementFormatError>>()?;
Self::Dcap(measurements_map)
Self::Dcap(DcapMeasurements::from_map(measurements_map)?)
}
})
}
Expand All @@ -217,13 +282,13 @@ impl MultiMeasurements {
return Err(DcapVerificationError::SgxNotSupported);
}
};
Ok(Self::Dcap(HashMap::from([
(DcapMeasurementRegister::MRTD, report.mr_td),
(DcapMeasurementRegister::RTMR0, report.rt_mr0),
(DcapMeasurementRegister::RTMR1, report.rt_mr1),
(DcapMeasurementRegister::RTMR2, report.rt_mr2),
(DcapMeasurementRegister::RTMR3, report.rt_mr3),
])))
Ok(Self::Dcap(DcapMeasurements::new(
report.mr_td,
report.rt_mr0,
report.rt_mr1,
report.rt_mr2,
report.rt_mr3,
)))
}

pub fn from_pcrs<'a>(pcrs: impl Iterator<Item = &'a [u8; 32]>) -> Self {
Expand All @@ -234,13 +299,13 @@ impl MultiMeasurements {
/// Mock TDX measurement values used in tests
#[cfg(any(test, feature = "mock"))]
pub fn mock_dcap_measurements() -> MultiMeasurements {
MultiMeasurements::Dcap(HashMap::from([
(DcapMeasurementRegister::MRTD, mock_tdx::MOCK_MRTD),
(DcapMeasurementRegister::RTMR0, mock_tdx::MOCK_RTMR0),
(DcapMeasurementRegister::RTMR1, mock_tdx::MOCK_RTMR1),
(DcapMeasurementRegister::RTMR2, mock_tdx::MOCK_RTMR2),
(DcapMeasurementRegister::RTMR3, mock_tdx::MOCK_RTMR3),
]))
MultiMeasurements::Dcap(DcapMeasurements::new(
mock_tdx::MOCK_MRTD,
mock_tdx::MOCK_RTMR0,
mock_tdx::MOCK_RTMR1,
mock_tdx::MOCK_RTMR2,
mock_tdx::MOCK_RTMR3,
))
}

/// An accepted measurement value given in the measurements file
Expand Down Expand Up @@ -364,27 +429,25 @@ impl MeasurementPolicy {
known_gcp_firmware: Option<&GcpFirmwareCache>,
) -> Result<(), AttestationError> {
if self.accepted_measurements.iter().any(|measurement_record| match measurements {
MultiMeasurements::Dcap(dcap_measurements) => {
match &measurement_record.measurements {
ExpectedMeasurements::Dcap(expected) => {
// All measurements in our policy must be given and must match
for (k, v) in expected.iter() {
match dcap_measurements.get(k) {
Some(actual_value) if v.iter().any(|v| actual_value == v) => {}
_ => return false,
}
MultiMeasurements::Dcap(dcap_measurements) => match &measurement_record.measurements {
ExpectedMeasurements::Dcap(expected) => {
// All measurements in our policy must be given and must match
for (k, v) in expected.iter() {
let actual_value = dcap_measurements.get(k);
if !v.iter().any(|v| actual_value == v) {
return false;
}
true
}
ExpectedMeasurements::Image(image_hashes) => compare_portable_dcap_measurement(
image_hashes,
dcap_measurements,
platform_metadata,
known_gcp_firmware,
),
ExpectedMeasurements::Azure(_) | ExpectedMeasurements::NoAttestation => false,
true
}
}
ExpectedMeasurements::Image(image_hashes) => compare_portable_dcap_measurement(
image_hashes,
dcap_measurements,
platform_metadata,
known_gcp_firmware,
),
ExpectedMeasurements::Azure(_) | ExpectedMeasurements::NoAttestation => false,
},
MultiMeasurements::Azure(azure_measurements) => {
if let ExpectedMeasurements::Azure(expected) = &measurement_record.measurements {
for (k, v) in expected.iter() {
Expand Down Expand Up @@ -598,7 +661,7 @@ impl MeasurementPolicy {
/// otherwise logs a warning and returns false.
pub(crate) fn compare_portable_dcap_measurement(
image_hashes: &DcapImageHashes,
dcap_measurements: &HashMap<DcapMeasurementRegister, [u8; 48]>,
dcap_measurements: &DcapMeasurements,
platform_metadata: Option<&PlatformMetadata>,
known_gcp_firmware: Option<&GcpFirmwareCache>,
) -> bool {
Expand All @@ -609,10 +672,7 @@ pub(crate) fn compare_portable_dcap_measurement(
// On GCP, fetch the firmware associated with the MRTD
let firmware = match platform_metadata.attestation_type {
ImageAttestationType::GcpTdx => {
let Some(mrtd) = dcap_measurements.get(&DcapMeasurementRegister::MRTD) else {
warn!("Could not match image hash measurement due to missing MRTD");
return false;
};
let mrtd = dcap_measurements.get(&DcapMeasurementRegister::MRTD);

let result = if let Some(cache) = known_gcp_firmware {
cache.get_or_fetch(*mrtd)
Expand Down Expand Up @@ -656,9 +716,8 @@ pub(crate) fn compare_portable_dcap_measurement(
};

if let Some(expected_mrtd) = expected_measurements.mrtd {
match dcap_measurements.get(&DcapMeasurementRegister::MRTD) {
Some(mrtd) if mrtd == &expected_mrtd => {}
_ => return false,
if dcap_measurements.get(&DcapMeasurementRegister::MRTD) != &expected_mrtd {
return false;
}
} else {
// This will only be the case with SelfHostedTdx which currently would
Expand All @@ -667,23 +726,20 @@ pub(crate) fn compare_portable_dcap_measurement(
}

if let Some(expected_rtmr0) = expected_measurements.rtmr0 {
match dcap_measurements.get(&DcapMeasurementRegister::RTMR0) {
Some(rtmr0) if rtmr0 == &expected_rtmr0 => {}
_ => return false,
if dcap_measurements.get(&DcapMeasurementRegister::RTMR0) != &expected_rtmr0 {
return false;
}
} else {
// This will only be the case with SelfHostedTdx which currently would
// already bail with the check above
return false;
}

if dcap_measurements.get(&DcapMeasurementRegister::RTMR1) != Some(&expected_measurements.rtmr1)
{
if dcap_measurements.get(&DcapMeasurementRegister::RTMR1) != &expected_measurements.rtmr1 {
return false;
}

if dcap_measurements.get(&DcapMeasurementRegister::RTMR2) != Some(&expected_measurements.rtmr2)
{
if dcap_measurements.get(&DcapMeasurementRegister::RTMR2) != &expected_measurements.rtmr2 {
return false;
}

Expand Down Expand Up @@ -739,6 +795,10 @@ mod tests {

use super::*;

fn test_dcap_measurements(mrtd: [u8; 48], rtmr0: [u8; 48]) -> MultiMeasurements {
MultiMeasurements::Dcap(DcapMeasurements::new(mrtd, rtmr0, [0u8; 48], [0u8; 48], [0u8; 48]))
}

/// MRTD from the pinned GCP firmware snapshot test asset
const GCP_FIRMWARE_MRTD: &str = "feb7486608382c1ff0e15b4648ddc0acea6ca974eb53e3529f4c4bd5ffbaa20bf335cb75965cea65fe473aed9647c162";
/// CFV from the same pinned GCP firmware snapshot test asset
Expand Down Expand Up @@ -877,13 +937,13 @@ mod tests {
let expected_measurements =
expected_dcap_registers(&image_hashes, &platform_metadata, Some(&firmware)).unwrap();

let measurements = MultiMeasurements::Dcap(HashMap::from([
(DcapMeasurementRegister::MRTD, expected_measurements.mrtd.unwrap()),
(DcapMeasurementRegister::RTMR0, expected_measurements.rtmr0.unwrap()),
(DcapMeasurementRegister::RTMR1, expected_measurements.rtmr1),
(DcapMeasurementRegister::RTMR2, expected_measurements.rtmr2),
(DcapMeasurementRegister::RTMR3, mock_tdx::MOCK_RTMR3),
]));
let measurements = MultiMeasurements::Dcap(DcapMeasurements::new(
expected_measurements.mrtd.unwrap(),
expected_measurements.rtmr0.unwrap(),
expected_measurements.rtmr1,
expected_measurements.rtmr2,
mock_tdx::MOCK_RTMR3,
));

policy.check_measurement(&measurements, Some(&platform_metadata)).unwrap();
}
Expand Down Expand Up @@ -941,6 +1001,18 @@ mod tests {
}
}

#[test]
fn test_dcap_header_format_rejects_incomplete_measurements() {
let input = serde_json::to_string(&HashMap::from([("0", hex::encode([0u8; 48]))])).unwrap();

let result = MultiMeasurements::from_header_format(&input, AttestationType::DcapTdx);

assert!(matches!(
result,
Err(MeasurementFormatError::MissingValue(register)) if register == "RTMR0"
));
}

/// A JSON policy that pins image-component hashes rather than raw
/// register values must parse into [`ExpectedMeasurements::Image`]
/// so the verifier can reconstruct the expected RTMRs from platform
Expand Down Expand Up @@ -1041,18 +1113,15 @@ mod tests {
let policy = MeasurementPolicy::from_json_bytes(json.as_bytes().to_vec()).unwrap();

// First value should match
let measurements1 =
MultiMeasurements::Dcap(HashMap::from([(DcapMeasurementRegister::MRTD, [0u8; 48])]));
let measurements1 = test_dcap_measurements([0u8; 48], [0u8; 48]);
assert!(policy.check_measurement(&measurements1, None).is_ok());

// Second value should also match
let measurements2 =
MultiMeasurements::Dcap(HashMap::from([(DcapMeasurementRegister::MRTD, [0x11u8; 48])]));
let measurements2 = test_dcap_measurements([0x11u8; 48], [0u8; 48]);
assert!(policy.check_measurement(&measurements2, None).is_ok());

// Different value should not match
let measurements3 =
MultiMeasurements::Dcap(HashMap::from([(DcapMeasurementRegister::MRTD, [0x22u8; 48])]));
let measurements3 = test_dcap_measurements([0x22u8; 48], [0u8; 48]);
assert!(policy.check_measurement(&measurements3, None).is_err());
}

Expand Down Expand Up @@ -1129,24 +1198,15 @@ mod tests {
let policy = MeasurementPolicy::from_json_bytes(json.as_bytes().to_vec()).unwrap();

// Both match (single + first of any)
let measurements1 = MultiMeasurements::Dcap(HashMap::from([
(DcapMeasurementRegister::MRTD, [0u8; 48]),
(DcapMeasurementRegister::RTMR0, [0x11u8; 48]),
]));
let measurements1 = test_dcap_measurements([0u8; 48], [0x11u8; 48]);
assert!(policy.check_measurement(&measurements1, None).is_ok());

// Both match (single + second of any)
let measurements2 = MultiMeasurements::Dcap(HashMap::from([
(DcapMeasurementRegister::MRTD, [0u8; 48]),
(DcapMeasurementRegister::RTMR0, [0x22u8; 48]),
]));
let measurements2 = test_dcap_measurements([0u8; 48], [0x22u8; 48]);
assert!(policy.check_measurement(&measurements2, None).is_ok());

// Single matches but any doesn't
let measurements3 = MultiMeasurements::Dcap(HashMap::from([
(DcapMeasurementRegister::MRTD, [0u8; 48]),
(DcapMeasurementRegister::RTMR0, [0x33u8; 48]),
]));
let measurements3 = test_dcap_measurements([0u8; 48], [0x33u8; 48]);
assert!(policy.check_measurement(&measurements3, None).is_err());
}

Expand Down Expand Up @@ -1342,10 +1402,7 @@ mod tests {
#[test]
fn test_multi_measurements_debug_prints_hex() {
let register_value = [0xabu8; 48];
let dcap = MultiMeasurements::Dcap(HashMap::from([(
DcapMeasurementRegister::MRTD,
register_value,
)]));
let dcap = test_dcap_measurements(register_value, [0u8; 48]);
let dcap_debug = format!("{dcap:?}");
assert!(dcap_debug.contains("DCAP"));
assert!(dcap_debug.contains(&hex::encode(register_value)));
Expand Down
Loading