Skip to content

Commit e4a3156

Browse files
committed
refactor(guest): keep gateway failover loop inline
1 parent ed97346 commit e4a3156

1 file changed

Lines changed: 23 additions & 112 deletions

File tree

dstack/dstack-util/src/system_setup.rs

Lines changed: 23 additions & 112 deletions
Original file line numberDiff line numberDiff line change
@@ -387,34 +387,6 @@ fn gateway_rpc_url(base: &str) -> String {
387387
}
388388
}
389389

390-
async fn register_first_available_gateway<T, F, Fut>(
391-
gateway_urls: &[String],
392-
mut register: F,
393-
) -> Result<T>
394-
where
395-
F: FnMut(String) -> Fut,
396-
Fut: std::future::Future<Output = Result<T>>,
397-
{
398-
if gateway_urls.is_empty() {
399-
bail!("Missing gateway urls");
400-
}
401-
let mut first_error = None;
402-
for gateway_url in gateway_urls {
403-
let gateway_url = gateway_url.trim_end_matches('/').to_string();
404-
match register(gateway_url).await {
405-
Ok(response) => return Ok(response),
406-
Err(err) => {
407-
warn!("Failed to register CVM: {err:?}, retrying with next dstack-gateway");
408-
if first_error.is_none() {
409-
first_error = Some(err);
410-
}
411-
}
412-
}
413-
}
414-
Err(first_error.unwrap_or_else(|| anyhow!("unknown error")))
415-
.context("Failed to register CVM, all dstack-gateway urls are down")
416-
}
417-
418390
struct GatewayContext<'a> {
419391
shared: &'a HostShared,
420392
keys: &'a AppKeys,
@@ -613,13 +585,28 @@ impl<'a> GatewayContext<'a> {
613585
warn!("failed to save gateway cache: {e:?}");
614586
}
615587

616-
// Read config and make the API call against the first healthy gateway.
617-
let response =
618-
register_first_available_gateway(&self.shared.sys_config.gateway_urls, |gateway_url| {
619-
let key_store = key_store.clone();
620-
async move { self.register_cvm(&gateway_url, &key_store).await }
621-
})
622-
.await?;
588+
if self.shared.sys_config.gateway_urls.is_empty() {
589+
bail!("Missing gateway urls");
590+
}
591+
// Read config and make API call
592+
let response = 'out: {
593+
let mut error = anyhow!("unknown error");
594+
for (i, url) in self.shared.sys_config.gateway_urls.iter().enumerate() {
595+
let response = self.register_cvm(url, &key_store).await;
596+
match response {
597+
Ok(response) => {
598+
break 'out response;
599+
}
600+
Err(err) => {
601+
warn!("Failed to register CVM: {err:?}, retrying with next dstack-gateway");
602+
if i == 0 {
603+
error = err;
604+
}
605+
}
606+
}
607+
}
608+
return Err(error).context("Failed to register CVM, all dstack-gateway urls are down");
609+
};
623610
let mut wg_info = response.wg.context("Missing wg info")?;
624611

625612
let client_ip = &wg_info.client_ip;
@@ -3632,10 +3619,8 @@ mod kms_provider_inventory_tests {
36323619

36333620
#[cfg(test)]
36343621
mod gateway_registration_refresh_tests {
3635-
use super::{gateway_rpc_url, register_first_available_gateway, GatewayKeyStore};
3636-
use anyhow::anyhow;
3622+
use super::{gateway_rpc_url, GatewayKeyStore};
36373623
use std::os::unix::fs::PermissionsExt as _;
3638-
use std::sync::{Arc, Mutex};
36393624

36403625
fn key_store(cert_not_after: u64) -> GatewayKeyStore {
36413626
GatewayKeyStore {
@@ -3701,78 +3686,4 @@ mod gateway_registration_refresh_tests {
37013686
assert!(!key_store(1_600).is_cert_valid_at(1_000));
37023687
assert!(!key_store(u64::MAX).is_cert_valid_at(u64::MAX));
37033688
}
3704-
3705-
#[tokio::test]
3706-
async fn ordered_outage_wrong_identity_and_malformed_fail_over() {
3707-
let urls = ["outage", "wrong-identity", "malformed", "healthy"]
3708-
.map(|name| format!("https://{name}.test/"));
3709-
let attempts = Arc::new(Mutex::new(Vec::new()));
3710-
let observed = attempts.clone();
3711-
let response = register_first_available_gateway(&urls, move |url| {
3712-
observed.lock().unwrap().push(url.clone());
3713-
async move {
3714-
if url.contains("healthy") {
3715-
Ok("stable-instance")
3716-
} else {
3717-
Err(anyhow!("injected registration failure"))
3718-
}
3719-
}
3720-
})
3721-
.await
3722-
.unwrap();
3723-
assert_eq!(response, "stable-instance");
3724-
assert_eq!(attempts.lock().unwrap().len(), 4);
3725-
}
3726-
3727-
#[tokio::test]
3728-
async fn first_success_short_circuits_and_all_failed_preserves_first_error() {
3729-
let healthy = ["first".to_string(), "must-not-run".to_string()];
3730-
let attempts = Arc::new(Mutex::new(0));
3731-
let observed = attempts.clone();
3732-
register_first_available_gateway(&healthy, move |_| {
3733-
*observed.lock().unwrap() += 1;
3734-
async { Ok::<_, anyhow::Error>(()) }
3735-
})
3736-
.await
3737-
.unwrap();
3738-
assert_eq!(*attempts.lock().unwrap(), 1);
3739-
3740-
let failed = ["first".to_string(), "second".to_string()];
3741-
let error = register_first_available_gateway::<(), _, _>(&failed, |url| async move {
3742-
Err(anyhow!("failure-at-{url}"))
3743-
})
3744-
.await
3745-
.unwrap_err();
3746-
assert!(format!("{error:#}").contains("failure-at-first"));
3747-
}
3748-
3749-
#[tokio::test]
3750-
async fn concurrent_refreshes_have_isolated_selection_state() {
3751-
let urls = ["down".to_string(), "healthy".to_string()];
3752-
let refresh = || async {
3753-
register_first_available_gateway(&urls, |url| async move {
3754-
if url == "healthy" {
3755-
Ok(url)
3756-
} else {
3757-
Err(anyhow!("down"))
3758-
}
3759-
})
3760-
.await
3761-
};
3762-
let (left, right) = tokio::join!(refresh(), refresh());
3763-
assert_eq!(left.unwrap(), "healthy");
3764-
assert_eq!(right.unwrap(), "healthy");
3765-
}
3766-
3767-
#[tokio::test]
3768-
async fn empty_gateway_inventory_fails_closed() {
3769-
let error = register_first_available_gateway::<(), _, _>(&[], |_| async {
3770-
panic!("registration must not run");
3771-
#[allow(unreachable_code)]
3772-
Ok(())
3773-
})
3774-
.await
3775-
.unwrap_err();
3776-
assert!(error.to_string().contains("Missing gateway urls"));
3777-
}
37783689
}

0 commit comments

Comments
 (0)