Skip to content
Merged
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
210 changes: 131 additions & 79 deletions dstack/guest-agent/src/rpc_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -61,8 +61,8 @@ const IDENTITY_RETRY_INTERVAL: Duration = Duration::from_secs(30);
///
/// Costs a quote and an event-log replay, which is why its result is cached for
/// the life of the process.
fn decode_identity(inner: &AppStateInner) -> Result<AppIdentity> {
let attestation = inner.info_attestation()?.into_v1();
fn decode_identity(platform: &dyn PlatformBackend) -> Result<AppIdentity> {
let attestation = platform.attestation_for_info()?.into_v1();
let app_info = attestation
.decode_app_info(false)
.context("failed to decode app info")?;
Expand Down Expand Up @@ -107,12 +107,12 @@ struct AppStateInner {
/// when each parsed the key itself.
app_root_signing_key: Option<SigningKey>,
/// Identity as decoded from the boot attestation. See [`AppIdentity`].
identity: RwLock<Option<Arc<AppIdentity>>>,
identity: tokio::sync::OnceCell<AppIdentity>,
/// When the last identity decode failed, so the retry can be throttled.
/// See [`IDENTITY_RETRY_INTERVAL`]. Separate from the cache above because
/// it is written only on the degraded path, and read only when the cache
/// is empty.
/// See [`IDENTITY_RETRY_INTERVAL`].
identity_last_failure: Mutex<Option<Instant>>,
/// Queues platform calls; see [`AppStateInner::attest`].
attest_lock: Arc<tokio::sync::Mutex<()>>,
/// `sys_vendor` and `product_name`, read once. Neither changes while the
/// VM is running.
cloud_vendor: String,
Expand All @@ -139,8 +139,26 @@ pub(crate) struct AppIdentity {
}

impl AppStateInner {
fn info_attestation(&self) -> Result<VersionedAttestation> {
self.platform.attestation_for_info()
/// Run a platform call on the blocking pool, one at a time.
///
/// Platform calls block on the hardware quote (about a second on TDX), so
/// running one on the executor stalls every connection the agent serves,
/// including the systemd watchdog heartbeat. Callers queue here rather than
/// on `dstack-attest`'s quote lock, so a flood of requests holds no threads.
/// The guard moves into the task, so a cancelled caller cannot let the
/// next call in while its own is still running.
async fn attest<T: Send + 'static>(
&self,
work: impl FnOnce(&dyn PlatformBackend) -> Result<T> + Send + 'static,
) -> Result<T> {
let guard = self.attest_lock.clone().lock_owned().await;
let platform = self.platform.clone();
tokio::task::spawn_blocking(move || {
let _guard = guard;
work(platform.as_ref())
})
.await
.context("the attestation task panicked")?
}

async fn issue_cert(
Expand All @@ -150,9 +168,10 @@ impl AppStateInner {
wire: AttestationWire,
) -> Result<Vec<String>> {
let pubkey = key.public_key_der();
let attested = pubkey.clone();
let attestation = wire.apply(
self.platform
.certificate_attestation(&pubkey)
self.attest(move |platform| platform.certificate_attestation(&attested))
.await
.context("Failed to get certificate attestation")?,
);
let csr = CertSigningRequestV2 {
Expand Down Expand Up @@ -269,8 +288,9 @@ impl AppState {
health,
gpu_attestor,
app_root_signing_key,
identity: RwLock::new(None),
identity_last_failure: Mutex::new(None),
identity: Default::default(),
identity_last_failure: Default::default(),
attest_lock: Arc::default(),
cloud_vendor: read_dmi_file("sys_vendor"),
cloud_product: read_dmi_file("product_name"),
}),
Expand All @@ -280,7 +300,7 @@ impl AppState {
// down, and `identity()` retries on demand. A failure here counts as
// the first attempt and arms the retry throttle, so the boot attempt
// and a request-driven one are on the same budget.
if let Err(err) = me.identity() {
if let Err(err) = me.identity().await {
error!("failed to decode app identity at startup: {err:?}");
}
me.maybe_request_demo_cert();
Expand All @@ -299,21 +319,23 @@ impl AppState {
self.inner.health.as_deref()
}

pub(crate) fn quote_response(&self, report_data: [u8; 64]) -> Result<GetQuoteResponse> {
pub(crate) async fn quote_response(&self, report_data: [u8; 64]) -> Result<GetQuoteResponse> {
let vm_config = self.inner.vm_config.clone();
self.inner
.platform
.quote_response(report_data, &self.inner.vm_config)
.attest(move |platform| platform.quote_response(report_data, &vm_config))
.await
}

/// Encoded attestation bytes for `Attest`, in the wire form the calling
/// surface commits to.
pub(crate) fn attest_cvm(
pub(crate) async fn attest_cvm(
&self,
report_data: [u8; 64],
wire: AttestationWire,
) -> Result<Vec<u8>> {
wire.apply(self.inner.platform.attest_cvm(report_data)?)
.to_bytes()
self.inner
.attest(move |platform| wire.apply(platform.attest_cvm(report_data)?).to_bytes())
.await
}

/// The application's root secp256k1 key, the root of every derived key and
Expand Down Expand Up @@ -354,54 +376,29 @@ impl AppState {
/// include anonymous ones. Within the window the caller is told the
/// identity is unavailable and when the next attempt is, and the platform
/// is not touched at all.
pub(crate) fn identity(&self) -> Result<Arc<AppIdentity>> {
if let Some(identity) = self
.inner
.identity
.read()
.or_panic("lock should never fail")
.as_ref()
{
return Ok(identity.clone());
}
// Two callers arriving together on a cold cache may both attempt once.
// That is the same race the cache has always had, and one extra quote
// is not worth holding a lock across the decode for.
if let Some(failed_at) = *self
.inner
.identity_last_failure
.lock()
.or_panic("lock should never fail")
{
let elapsed = failed_at.elapsed();
if elapsed < IDENTITY_RETRY_INTERVAL {
let retry_in = (IDENTITY_RETRY_INTERVAL - elapsed).as_secs() + 1;
anyhow::bail!(
"the app identity is unavailable: decoding it failed and the next attempt is at most {retry_in}s away"
);
}
}
// Blocking, and deliberately left on the executor: the throttle above
// caps this at one quote per interval for the whole process, which is
// far short of what would justify a `spawn_blocking` hop and making
// every caller of `identity()` async to reach it.
let identity = match decode_identity(&self.inner) {
Ok(identity) => Arc::new(identity),
Err(err) => {
*self
.inner
.identity_last_failure
.lock()
.or_panic("lock should never fail") = Some(Instant::now());
return Err(err);
}
};
*self
.inner
pub(crate) async fn identity(&self) -> Result<&AppIdentity> {
self.inner
.identity
.write()
.or_panic("lock should never fail") = Some(identity.clone());
Ok(identity)
.get_or_try_init(|| async {
let last_failure = &self.inner.identity_last_failure;
if let Some(failed_at) = *last_failure.lock().or_panic("lock should never fail") {
let elapsed = failed_at.elapsed();
if elapsed < IDENTITY_RETRY_INTERVAL {
let retry_in = (IDENTITY_RETRY_INTERVAL - elapsed).as_secs() + 1;
anyhow::bail!(
"the app identity is unavailable: decoding it failed and the next attempt is at most {retry_in}s away"
);
}
}
self.inner
.attest(decode_identity)
.await
.inspect_err(|_| {
*last_failure.lock().or_panic("lock should never fail") =
Some(Instant::now());
})
})
.await
}

/// The VM's hardware configuration, as the VMM produced it.
Expand Down Expand Up @@ -535,7 +532,10 @@ impl InternalRpcHandler {

pub async fn get_info(state: &AppState, external: bool) -> Result<AppInfo> {
let hide_tcb_info = external && !state.config().app_compose.public_tcbinfo;
let versioned_attestation = state.inner.info_attestation()?;
let versioned_attestation = state
.inner
.attest(|platform| platform.attestation_for_info())
.await?;
let attestation = versioned_attestation.into_v1();
let app_info = attestation
.decode_app_info(false)
Expand Down Expand Up @@ -678,7 +678,7 @@ impl DstackGuestRpc for InternalRpcHandler {

async fn get_quote(self, request: RawQuoteArgs) -> Result<GetQuoteResponse> {
let report_data = pad64(&request.report_data).context("Report data is too long")?;
self.state.quote_response(report_data)
self.state.quote_response(report_data).await
}

/// Always fails. See the RPC's doc comment in agent_rpc.proto: the method
Expand Down Expand Up @@ -805,7 +805,8 @@ impl DstackGuestRpc for InternalRpcHandler {
Ok(AttestResponse {
attestation: self
.state
.attest_cvm(report_data, AttestationWire::Legacy)?,
.attest_cvm(report_data, AttestationWire::Legacy)
.await?,
})
}

Expand Down Expand Up @@ -912,7 +913,7 @@ impl TappdRpc for InternalRpcHandlerV0 {
};
let report_data =
content_type.to_report_data_with_hash(&request.report_data, &request.hash_algorithm)?;
let response = self.state.quote_response(report_data)?;
let response = self.state.quote_response(report_data).await?;
Ok(TdxQuoteResponse {
quote: response.quote,
event_log: response.event_log,
Expand Down Expand Up @@ -992,7 +993,7 @@ impl WorkerRpc for ExternalRpcHandler {
request: GetAttestationForAppKeyRequest,
) -> Result<GetQuoteResponse> {
let report_data = self.app_key_report_data(&request.algorithm).await?;
self.state.quote_response(report_data)
self.state.quote_response(report_data).await
}
}

Expand Down Expand Up @@ -1130,18 +1131,25 @@ pub(crate) mod tests {
/// How many times the fixture platform was asked to attest for `Info`, and
/// whether it should refuse. Counting the calls is the only way to see the
/// identity throttle work: what it changes is how often the platform is
/// touched, not what any single call returns.
/// touched, not what any single call returns. The gate holds the next
/// quote inside the platform, as a slow quote does.
#[derive(Default)]
struct InfoAttestationProbe {
calls: AtomicUsize,
failing: AtomicBool,
gate: Mutex<
Option<(
tokio::sync::oneshot::Sender<()>,
std::sync::mpsc::Receiver<()>,
)>,
>,
}

impl InfoAttestationProbe {
fn failing() -> Self {
Self {
calls: AtomicUsize::new(0),
failing: AtomicBool::new(true),
..Self::default()
}
}

Expand All @@ -1152,6 +1160,27 @@ pub(crate) mod tests {
fn set_failing(&self, failing: bool) {
self.failing.store(failing, Ordering::Relaxed);
}

/// Arm the gate: the receiver fires once a platform call is blocked
/// in it, and sending on the sender lets that call finish.
fn arm_gate(
&self,
) -> (
tokio::sync::oneshot::Receiver<()>,
std::sync::mpsc::Sender<()>,
) {
let (entered_tx, entered_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = std::sync::mpsc::channel();
*self.gate.lock().unwrap() = Some((entered_tx, release_rx));
(entered_rx, release_tx)
}

fn pass_gate(&self) {
if let Some((entered, release)) = self.gate.lock().unwrap().take() {
entered.send(()).unwrap();
release.recv_timeout(Duration::from_secs(5)).unwrap();
}
}
}

/// The fixture state, built against a platform the probe watches.
Expand Down Expand Up @@ -1350,6 +1379,7 @@ pNs85uhOZE8z2jr8Pg==
report_data: [u8; 64],
vm_config: &str,
) -> Result<GetQuoteResponse> {
self.probe.pass_gate();
let attestation = patch_report_data(&self.attestation, report_data);
let Some(quote) = attestation.platform.tdx_quote().map(ToOwned::to_owned) else {
return Err(anyhow::anyhow!(
Expand Down Expand Up @@ -1410,8 +1440,9 @@ pNs85uhOZE8z2jr8Pg==
"/nonexistent/nvattest",
),
app_root_signing_key: SigningKey::from_slice(&DUMMY_K256_KEY).ok(),
identity: RwLock::new(None),
identity_last_failure: Mutex::new(None),
identity: Default::default(),
identity_last_failure: Default::default(),
attest_lock: Arc::default(),
// Read the same way production does, so a test comparing v1 `Info`
// against v0 `get_info` compares like with like.
cloud_vendor: read_dmi_file("sys_vendor"),
Expand Down Expand Up @@ -2012,6 +2043,22 @@ pNs85uhOZE8z2jr8Pg==
assert!(err.contains("removed in dstack 0.6.0"), "{err}");
}

/// A quote must not park the runtime. On a single-threaded runtime the
/// test only sees the quote enter the platform if it left the executor.
#[tokio::test(flavor = "current_thread")]
async fn a_blocking_quote_does_not_park_the_runtime() {
let probe = Arc::new(InfoAttestationProbe::default());
let (state, _guard) = setup_test_state_with_probe(probe.clone()).await;
let (entered_rx, release_tx) = probe.arm_gate();
let quote = tokio::spawn(async move { state.quote_response([0u8; 64]).await });
tokio::time::timeout(Duration::from_secs(2), entered_rx)
.await
.unwrap()
.unwrap();
release_tx.send(()).unwrap();
quote.await.unwrap().unwrap();
}

/// A failed decode must be as cheap to repeat as a cached success is.
/// `identity()` is reached from the anonymous `/prpc/v1/Info`, and each
/// attempt is a hardware quote plus an RTMR replay under the global quote
Expand All @@ -2024,12 +2071,13 @@ pNs85uhOZE8z2jr8Pg==

let err = state
.identity()
.await
.expect_err("the platform refuses to attest");
assert!(err.to_string().contains("cannot attest"), "{err}");
assert_eq!(probe.calls(), 1);

for _ in 0..8 {
let err = state.identity().expect_err("still throttled");
let err = state.identity().await.expect_err("still throttled");
assert!(
err.to_string().contains("the app identity is unavailable"),
"{err}"
Expand All @@ -2051,6 +2099,7 @@ pNs85uhOZE8z2jr8Pg==
let (state, _guard) = setup_test_state_with_probe(probe.clone()).await;
state
.identity()
.await
.expect_err("the platform refuses to attest");

// Age the recorded failure rather than sleeping out the interval.
Expand All @@ -2065,11 +2114,14 @@ pNs85uhOZE8z2jr8Pg==
);
probe.set_failing(false);

let identity = state.identity().expect("the platform recovered");
let identity = state.identity().await.expect("the platform recovered");
assert_eq!(probe.calls(), 2);

let again = state.identity().expect("a decoded identity is cached");
assert!(Arc::ptr_eq(&identity, &again));
let again = state
.identity()
.await
.expect("a decoded identity is cached");
assert!(std::ptr::eq(identity, again));
assert_eq!(
probe.calls(),
2,
Expand Down
Loading
Loading