diff --git a/crates/attestation/src/measurements.rs b/crates/attestation/src/measurements.rs index fbdba06..a0b3fed 100644 --- a/crates/attestation/src/measurements.rs +++ b/crates/attestation/src/measurements.rs @@ -83,11 +83,81 @@ fn parse_azure_pcr_index(value: &str) -> Result { 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, + ) -> Result { + 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 { + [ + (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), + Dcap(DcapMeasurements), Azure(HashMap), NoAttestation, } @@ -95,9 +165,7 @@ pub enum MultiMeasurements { 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() } @@ -107,17 +175,14 @@ impl fmt::Debug for MultiMeasurements { } /// Used to display DCAP measurements as hex -struct DcapHexDebug<'a>(&'a HashMap); +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(®ister, &hex_value); } map.finish() } @@ -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() @@ -201,7 +266,7 @@ impl MultiMeasurements { )) }) .collect::>()?; - Self::Dcap(measurements_map) + Self::Dcap(DcapMeasurements::from_map(measurements_map)?) } }) } @@ -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) -> Self { @@ -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 @@ -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() { @@ -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, + dcap_measurements: &DcapMeasurements, platform_metadata: Option<&PlatformMetadata>, known_gcp_firmware: Option<&GcpFirmwareCache>, ) -> bool { @@ -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) @@ -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 @@ -667,9 +726,8 @@ 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 @@ -677,13 +735,11 @@ pub(crate) fn compare_portable_dcap_measurement( 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; } @@ -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 @@ -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(); } @@ -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 @@ -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()); } @@ -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()); } @@ -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)));