diff --git a/.agents/knowledge/domain-facts.md b/.agents/knowledge/domain-facts.md index ba1df1551..5ead47e2b 100644 --- a/.agents/knowledge/domain-facts.md +++ b/.agents/knowledge/domain-facts.md @@ -56,6 +56,15 @@ bypassed under which build feature. - Hold `spdm::AppContextGuard` across migration and rebind exchanges so dropping either role's future wipes the buffer. Keep cancellation coverage for all four entry points in `src/migtd/src/spdm/tests.rs`. +- SPDM `FINISH` clears `runtime_info.last_session_id` while the established + session still holds keys. Retire every exchange-owned `SpdmContext::session` + slot on return or cancellation, not just the last handshake ID + (`deps/spdm-rs/spdmlib/src/requester/finish_req.rs`, `responder/context.rs`; + upstream cleanup: `7b40eccb`). +- The guard also tears down those slots before `finalize_spdm_session` attempts + transport shutdown. Shutdown must run on protocol failure or timeout without + replacing the primary error. Keep the teardown assertions for both roles, + handshaking/established states, cancellation, and repeated cleanup. ## TDINFO / MROwner / MROwnerConfig semantics diff --git a/doc/MigTD_Functionality_Summary.md b/doc/MigTD_Functionality_Summary.md index 0a9e2b118..6920f00c4 100644 --- a/doc/MigTD_Functionality_Summary.md +++ b/doc/MigTD_Functionality_Summary.md @@ -84,7 +84,11 @@ Implemented in `src/migtd/src/migration/{data.rs, session.rs, event.rs}`: - `StartMigration` — run the MSK key-exchange flow. - `StartRebinding` — approve rebinding the user TD to a new MigTD (policy v2). - `GetTdReport` — return a TD report using MigTD's fixed report data; the - request payload contains only the migration request ID. + canonical request payload is the 8-byte migration request ID. + **REVERT_ME (OS-transition testing only):** temporarily also accept the + legacy 72-byte payload with a trailing 64-byte `REPORT_DATA`. That tail + is discarded silently, never stored or used to generate the report. + Remove this compatibility path once older test OS versions are retired. - `EnableLogArea` — enable/raise the VMM log level for a request. - `GetMigtdData` — return MigTD attestation data (policy v2). - `ReportStatus` returns the per-request result (`ReportStatusResponse` carries diff --git a/src/migtd/src/migration/mod.rs b/src/migtd/src/migration/mod.rs index 3b1858ba4..2dba8e618 100644 --- a/src/migtd/src/migration/mod.rs +++ b/src/migtd/src/migration/mod.rs @@ -345,10 +345,16 @@ impl ReportInfo { data_length: u32, payload: &[u8], ) -> core::result::Result { - if data_length != core::mem::size_of::() as u32 { + let request_id_size = core::mem::size_of::() as u32; + let legacy_size = + request_id_size + tdx_tdcall::tdreport::TD_REPORT_ADDITIONAL_DATA_SIZE as u32; + // REVERT_ME: tolerate the retired REPORT_DATA tail during OS-transition testing. + if data_length != request_id_size && data_length != legacy_size { return Err(MigrationResult::InvalidParameter); } payload + .get(..data_length as usize) + .ok_or(MigrationResult::InvalidParameter)? .pread(0) .map_err(|_| MigrationResult::InvalidParameter) } diff --git a/src/migtd/src/migration/session.rs b/src/migtd/src/migration/session.rs index 4972a1af4..3f222c370 100644 --- a/src/migtd/src/migration/session.rs +++ b/src/migtd/src/migration/session.rs @@ -1463,7 +1463,7 @@ mod test { use crate::migration::MIGTD_MIGRATION_INFO_HEADER_SIZE; use crate::migration::{ data::{RequestDataBufferHeader, WaitForRequestResponse}, - EnableLogAreaInfo, MigrationResult, + EnableLogAreaInfo, MigrationResult, ReportInfo, }; use core::mem::size_of; use core::task::Poll; @@ -1653,33 +1653,84 @@ mod test { } _ => panic!("Expected GetTdReport, got unexpected variant"), } + assert!(pending.is_none()); cleanup_request(request_id); } #[test] - fn test_parse_get_td_report_legacy_report_data_rejected() { + fn test_parse_get_td_report_legacy_report_data_ignored() { + assert_eq!(size_of::(), size_of::()); let request_id: u64 = 0xCC00_0000_0000_0003; - let mut payload = vec![0xCC; 72]; - payload[0..8].copy_from_slice(&request_id.to_le_bytes()); - let buf = build_request_buffer(3, &payload); - let mut pending = None; - let result = parse_request(&buf, HDR_LEN, &mut pending); - assert!(matches!( - result, - Poll::Ready(Err(MigrationResult::InvalidParameter)) - )); + for report_data_byte in [0x00, 0x5a, 0xff] { + let mut payload = vec![report_data_byte; 72]; + payload[0..8].copy_from_slice(&request_id.to_le_bytes()); + let buf = build_request_buffer(3, &payload); + let mut pending = None; + let result = parse_request(&buf, HDR_LEN, &mut pending); + match result { + Poll::Ready(Ok(WaitForRequestResponse::GetTdReport(info))) => { + assert_eq!(info.mig_request_id, request_id); + } + _ => panic!("Expected GetTdReport for legacy payload"), + } + assert!(pending.is_none()); + cleanup_request(request_id); + } } #[test] fn test_parse_get_td_report_wrong_size_rejected() { - // Only the exact 8-byte request ID is accepted. - let buf = build_request_buffer(3, &[0u8; 16]); + for payload_size in 0..=80 { + if payload_size == 8 || payload_size == 72 { + continue; + } + let buf = build_request_buffer(3, &vec![0u8; payload_size]); + let mut pending = None; + let result = parse_request(&buf, HDR_LEN, &mut pending); + assert!(matches!( + result, + Poll::Ready(Err(MigrationResult::InvalidParameter)) + )); + } + } + + #[test] + fn test_parse_get_td_report_truncated_declared_payload_rejected() { + for data_length in [8, 72] { + for payload_size in 0..data_length as usize { + let payload = vec![0u8; payload_size]; + assert!(matches!( + ReportInfo::read_from_bytes(data_length, &payload), + Err(MigrationResult::InvalidParameter) + )); + let buf = build_raw_buffer(0x0301, data_length, &payload); + let mut pending = None; + assert!(matches!( + parse_request(&buf, HDR_LEN, &mut pending), + Poll::Ready(Err(MigrationResult::InvalidParameter)) + )); + } + } + } + + #[cfg(feature = "policy_v2")] + #[test] + fn test_parse_get_migtd_data_retains_report_data() { + let request_id: u64 = 0xCD00_0000_0000_0004; + let report_data = [0x5a; 64]; + let mut payload = request_id.to_le_bytes().to_vec(); + payload.extend_from_slice(&report_data); + let buf = build_request_buffer(5, &payload); let mut pending = None; - let result = parse_request(&buf, HDR_LEN, &mut pending); - assert!(matches!( - result, - Poll::Ready(Err(MigrationResult::InvalidParameter)) - )); + match parse_request(&buf, HDR_LEN, &mut pending) { + Poll::Ready(Ok(WaitForRequestResponse::GetMigtdData(info))) => { + assert_eq!(info.mig_request_id, request_id); + assert_eq!(info.reportdata, report_data); + } + _ => panic!("Expected GetMigtdData, got unexpected variant"), + } + assert!(pending.is_none()); + cleanup_request(request_id); } #[test] diff --git a/src/migtd/src/migration/spdm_session.rs b/src/migtd/src/migration/spdm_session.rs index 4d88f89c1..56d0cb9a6 100644 --- a/src/migtd/src/migration/spdm_session.rs +++ b/src/migtd/src/migration/spdm_session.rs @@ -31,14 +31,15 @@ pub(super) fn map_spdm_setup_err(mig_request_id: u64) -> MigrationResult { MigrationResult::SecureSessionError } -/// Run an SPDM session `body` under [`SPDM_TIMEOUT`], then shut down the -/// transport associated with `io_ref`. Returns any value produced by `body`. +/// Run an SPDM session `body` under [`SPDM_TIMEOUT`] and always attempt transport +/// shutdown. A protocol error or timeout takes precedence over a shutdown error. /// /// `body` is the already-constructed future returned by an SPDM exchange /// function (e.g. `spdm_requester_transfer_msk`, `spdm_responder_rebind_new`). /// The caller owns the SPDM context that `body` borrows; this helper only /// drives `body` to completion and then takes the device-IO lock to invoke -/// `shutdown_transport`. +/// `shutdown_transport`. The exchange's `AppContextGuard` retires its session +/// keys before shutdown, including when the timeout drops `body`. pub(super) async fn finalize_spdm_session( body: Fut, io_ref: SpdmDeviceIoArc, @@ -47,7 +48,7 @@ pub(super) async fn finalize_spdm_session( where Fut: Future>, { - let value = with_timeout(SPDM_TIMEOUT, body) + let session_result = with_timeout(SPDM_TIMEOUT, body) .await .map_err(|e| { log::error!( @@ -55,18 +56,22 @@ where "finalize_spdm_session: body timeout: {e:?}\n" ); MigrationResult::from(e) - })? - .map_err(|e| { - log::error!( - migration_request_id = mig_request_id; - "finalize_spdm_session: body error: {e:?}\n" - ); - crate::spdm::decode_spdm_session_err(e) - })?; + }) + .and_then(|result| { + result.map_err(|e| { + log::error!( + migration_request_id = mig_request_id; + "finalize_spdm_session: body error: {e:?}\n" + ); + crate::spdm::decode_spdm_session_err(e) + }) + }); let mut transport_lock = io_ref.lock(); let transport = transport_lock.deref_mut(); - shutdown_transport(&mut transport.transport, mig_request_id).await?; + let shutdown_result = shutdown_transport(&mut transport.transport, mig_request_id).await; + let value = session_result?; + shutdown_result?; Ok(value) } diff --git a/src/migtd/src/spdm/handshake.rs b/src/migtd/src/spdm/handshake.rs index d8b526d16..3347fbf08 100644 --- a/src/migtd/src/spdm/handshake.rs +++ b/src/migtd/src/spdm/handshake.rs @@ -57,8 +57,9 @@ pub(super) async fn requester_handshake_prelude( /// /// This helper does NOT zeroize `app_context_data_buffer` itself: the buffer /// may carry the caller's ephemeral private key, so callers hold an -/// `AppContextGuard` across the await to wipe it on success, error, or -/// cancellation (see `spdm_responder_transfer_msk` / `spdm_responder_rebind_new`). +/// `AppContextGuard` across the await to wipe it and retire session keys on +/// success, error, or cancellation (see `spdm_responder_transfer_msk` / +/// `spdm_responder_rebind_new`). /// /// Callers that need to expose extra context to VDM handlers (e.g. setting /// `spdm_responder_ex.info = RebindInformation(..)`) must do so before diff --git a/src/migtd/src/spdm/mod.rs b/src/migtd/src/spdm/mod.rs index 0c63f60af..b7a0ebeb9 100644 --- a/src/migtd/src/spdm/mod.rs +++ b/src/migtd/src/spdm/mod.rs @@ -23,7 +23,7 @@ use codec::Codec; use codec::Reader; use codec::Writer; use log::error; -use spdmlib::common::SpdmDeviceIo; +use spdmlib::common::{SpdmContext, SpdmDeviceIo}; use spdmlib::error::*; use spdmlib::protocol::{SpdmDigestStruct, SPDM_MAX_HASH_SIZE}; use spin::Mutex; @@ -50,17 +50,26 @@ use crate::spdm::vmcall_msg::VMCALL_SPDM_MESSAGE_HEADER_SIZE; pub(crate) type SpdmDeviceIoArc = Arc>>; -// The raw application buffer holds the ephemeral signing key. Borrow the whole -// context so the exchange can keep using it while the buffer is wiped on both -// normal return and future cancellation. +// Borrow the whole context so the exchange can keep using it while its session +// keys and raw application-buffer signing key are cleared on return or cancellation. struct AppContextGuard<'a, T> { context: &'a mut T, - buffer: fn(&mut T) -> &mut [u8], + common: fn(&mut T) -> &mut SpdmContext, } impl Drop for AppContextGuard<'_, T> { fn drop(&mut self) { - (self.buffer)(self.context).zeroize(); + let common = (self.common)(self.context); + teardown_sessions(common); + common.app_context_data_buffer.zeroize(); + } +} + +pub(crate) fn teardown_sessions(context: &mut SpdmContext) { + // FINISH clears last_session_id while the established session still holds keys. + // Each context belongs to one exchange, so retire every slot before shutdown. + for session in &mut context.session { + session.teardown(); } } diff --git a/src/migtd/src/spdm/spdm_rebind.rs b/src/migtd/src/spdm/spdm_rebind.rs index c61fa9045..060ccdeb8 100644 --- a/src/migtd/src/spdm/spdm_rebind.rs +++ b/src/migtd/src/spdm/spdm_rebind.rs @@ -21,7 +21,7 @@ pub async fn spdm_requester_rebind_old( ) -> Result<(), SpdmStatus> { let guard = super::AppContextGuard { context: spdm_requester, - buffer: |context| &mut context.common.app_context_data_buffer, + common: |context| &mut context.common, }; spdm_requester_rebind_old_inner(guard.context, rebind_info, peer_data).await } @@ -61,7 +61,7 @@ pub async fn spdm_responder_rebind_new<'a>( ) -> Result<(), SpdmStatus> { let guard = super::AppContextGuard { context: spdm_responder_ex, - buffer: |context| &mut context.responder_context.common.app_context_data_buffer, + common: |context| &mut context.responder_context.common, }; let spdm_responder_ex = &mut *guard.context; diff --git a/src/migtd/src/spdm/spdm_req.rs b/src/migtd/src/spdm/spdm_req.rs index 2e0018e9a..b9e116d0e 100644 --- a/src/migtd/src/spdm/spdm_req.rs +++ b/src/migtd/src/spdm/spdm_req.rs @@ -91,7 +91,7 @@ pub async fn spdm_requester_transfer_msk( ) -> Result { let guard = super::AppContextGuard { context: spdm_requester, - buffer: |context| &mut context.common.app_context_data_buffer, + common: |context| &mut context.common, }; spdm_requester_transfer_msk_inner( guard.context, diff --git a/src/migtd/src/spdm/spdm_rsp.rs b/src/migtd/src/spdm/spdm_rsp.rs index 1199bd5e5..98ab8f21e 100644 --- a/src/migtd/src/spdm/spdm_rsp.rs +++ b/src/migtd/src/spdm/spdm_rsp.rs @@ -169,7 +169,7 @@ pub async fn spdm_responder_transfer_msk<'a>( ) -> Result<(), SpdmStatus> { let guard = super::AppContextGuard { context: spdm_responder_ex, - buffer: |context| &mut context.responder_context.common.app_context_data_buffer, + common: |context| &mut context.responder_context.common, }; let spdm_responder_ex = &mut *guard.context; @@ -180,6 +180,7 @@ pub async fn spdm_responder_transfer_msk<'a>( mig_info, exchange_information, }; + spdm_responder_ex.mig_info_exchanged = false; spdm_responder_ex.remote_information = None; // The VDM handler reads `mig_info` / `exchange_information` from @@ -209,17 +210,7 @@ pub async fn rsp_handle_message(spdm_responder: &mut ResponderContext) -> Result match res { Ok(Ok(_)) => {} - Ok(Err(spdm_status)) => { - if spdm_status.severity == StatusSeverity::ERROR - && matches!(spdm_status.status_code, StatusCode::VDM(_)) - { - return Err(spdm_status); - } - if spdm_status == SPDM_STATUS_INVALID_STATE_LOCAL { - // Terminate the responder upon invalid state. - return Err(spdm_status); - } - } + Ok(Err(spdm_status)) => return Err(spdm_status), Err(_) => return Err(SPDM_STATUS_RECEIVE_FAIL), } diff --git a/src/migtd/src/spdm/tests.rs b/src/migtd/src/spdm/tests.rs index e3ed00a02..65a7c34a5 100644 --- a/src/migtd/src/spdm/tests.rs +++ b/src/migtd/src/spdm/tests.rs @@ -6,18 +6,31 @@ use super::*; use crate::migration::session::ExchangeInformation; use core::future::{pending, Future}; use core::task::{Context, Poll, Waker}; +use spdmlib::{ + common::{session::SpdmSessionState, SpdmContext, INVALID_SESSION_ID}, + message::{SpdmMessageHeader, SpdmRequestResponseCode}, + protocol::SpdmVersion, +}; +#[derive(Default)] struct TestTransport { fail_read: bool, + incoming: Vec, + offset: usize, } impl AsyncRead for TestTransport { - async fn read(&mut self, _buffer: &mut [u8]) -> async_io::Result { + async fn read(&mut self, buffer: &mut [u8]) -> async_io::Result { if self.fail_read { - Err(async_io::ErrorKind::ConnectionAborted.into()) - } else { - pending().await + return Err(async_io::ErrorKind::ConnectionAborted.into()); } + let size = buffer.len().min(self.incoming.len() - self.offset); + if size == 0 { + return pending().await; + } + buffer[..size].copy_from_slice(&self.incoming[self.offset..self.offset + size]); + self.offset += size; + Ok(size) } } @@ -42,7 +55,11 @@ fn finish_or_cancel(future: impl Future>, fail #[test] fn requester_migration_app_context_is_wiped_on_cancellation_and_error() { for fail_read in [false, true] { - let (mut requester, _) = spdm_requester(TestTransport { fail_read }).unwrap(); + let (mut requester, _) = spdm_requester(TestTransport { + fail_read, + ..Default::default() + }) + .unwrap(); requester.common.app_context_data_buffer.fill(0xa5); let mig_info = MigtdMigrationInformation::default(); let exchange_information = ExchangeInformation::default(); @@ -69,7 +86,12 @@ fn requester_migration_app_context_is_wiped_on_cancellation_and_error() { #[test] fn responder_migration_app_context_is_wiped_on_cancellation_and_error() { for fail_read in [false, true] { - let (mut responder, _) = spdm_responder(TestTransport { fail_read }).unwrap(); + let (mut responder, _) = spdm_responder(TestTransport { + fail_read, + ..Default::default() + }) + .unwrap(); + responder.mig_info_exchanged = true; responder .responder_context .common @@ -89,6 +111,7 @@ fn responder_migration_app_context_is_wiped_on_cancellation_and_error() { fail_read, ); + assert!(!responder.mig_info_exchanged); assert!(responder .responder_context .common @@ -102,7 +125,11 @@ fn responder_migration_app_context_is_wiped_on_cancellation_and_error() { #[test] fn requester_rebind_app_context_is_wiped_on_cancellation_and_error() { for fail_read in [false, true] { - let (mut requester, _) = spdm_requester(TestTransport { fail_read }).unwrap(); + let (mut requester, _) = spdm_requester(TestTransport { + fail_read, + ..Default::default() + }) + .unwrap(); requester.common.app_context_data_buffer.fill(0xa5); let mig_info = MigtdMigrationInformation::default(); @@ -123,7 +150,11 @@ fn requester_rebind_app_context_is_wiped_on_cancellation_and_error() { #[test] fn responder_rebind_app_context_is_wiped_on_cancellation_and_error() { for fail_read in [false, true] { - let (mut responder, _) = spdm_responder(TestTransport { fail_read }).unwrap(); + let (mut responder, _) = spdm_responder(TestTransport { + fail_read, + ..Default::default() + }) + .unwrap(); responder .responder_context .common @@ -144,3 +175,105 @@ fn responder_rebind_app_context_is_wiped_on_cancellation_and_error() { .all(|byte| *byte == 0)); } } + +fn assert_teardown_clears_sessions(context: &mut SpdmContext) { + for last_session_id in [Some(1), None] { + for cancel in [false, true] { + context.app_context_data_buffer.fill(0xa5); + for (index, session) in context.session.iter_mut().enumerate() { + session.setup(u32::try_from(index + 1).unwrap()).unwrap(); + session.set_session_state(if last_session_id.is_some() { + SpdmSessionState::SpdmSessionHandshaking + } else { + SpdmSessionState::SpdmSessionEstablished + }); + let mut secret = session.get_application_secret(); + secret.request_direction.encryption_key.data.fill(0xa5); + secret.request_direction.encryption_key.data_size = 32; + session.set_application_secret(secret); + } + // FINISH clears this field without removing the established session. + context.runtime_info.set_last_session_id(last_session_id); + + for _ in 0..2 { + let mut future = Box::pin(async { + let _guard = AppContextGuard { + context: &mut *context, + common: |context| context, + }; + if cancel { + pending::<()>().await; + } + }); + let mut task_context = Context::from_waker(Waker::noop()); + assert_eq!(future.as_mut().poll(&mut task_context).is_pending(), cancel); + drop(future); + + assert!(context + .app_context_data_buffer + .iter() + .all(|byte| *byte == 0)); + for session in &context.session { + assert_eq!(session.get_session_id(), INVALID_SESSION_ID); + assert_eq!( + session.get_session_state(), + SpdmSessionState::SpdmSessionNotStarted + ); + assert_eq!(session.get_application_secret(), Default::default()); + } + } + } + } +} + +#[test] +fn requester_teardown_clears_handshaking_and_established_sessions() { + let (mut requester, _) = spdm_requester(TestTransport::default()).unwrap(); + assert_teardown_clears_sessions(&mut requester.common); +} + +#[test] +fn responder_teardown_clears_handshaking_and_established_sessions() { + let (mut responder, _) = spdm_responder(TestTransport::default()).unwrap(); + assert_teardown_clears_sessions(&mut responder.responder_context.common); +} + +#[test] +fn responder_propagates_malformed_message_error() { + // GET_VERSION needs a two-byte payload after its header. + let mut incoming = vec![0u8; VMCALL_SPDM_MESSAGE_HEADER_SIZE + 2]; + let mut writer = Writer::init(&mut incoming); + vmcall_msg::VmCallMessageHeader { + version: vmcall_msg::VMCALL_SPDM_VERSION, + msg_type: vmcall_msg::VmCallMessageType::SpdmMessage, + length: 2, + } + .encode(&mut writer) + .unwrap(); + SpdmMessageHeader { + version: SpdmVersion::SpdmVersion10, + request_response_code: SpdmRequestResponseCode::SpdmRequestGetVersion, + } + .encode(&mut writer) + .unwrap(); + + let (mut responder, _) = spdm_responder(TestTransport { + incoming, + ..Default::default() + }) + .unwrap(); + let mig_info = MigtdMigrationInformation::default(); + let exchange_information = ExchangeInformation::default(); + let mut future = Box::pin(spdm_responder_transfer_msk( + &mut responder, + &mig_info, + &exchange_information, + #[cfg(feature = "policy_v2")] + Vec::new(), + )); + let mut context = Context::from_waker(Waker::noop()); + assert_eq!( + future.as_mut().poll(&mut context), + Poll::Ready(Err(SPDM_STATUS_INVALID_MSG_FIELD)) + ); +}