use std::time::Duration;
use mkit_attest::grant::{Namespace, is_loopback_origin};
use mkit_core::repo_identity::RepositoryIdentity;
use mkit_transport_connect::{Completion, audience_from_url, repository_identity_from_url};
use crate::config::{self, LayeredConfig};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Target {
pub name: String,
pub endpoint: String,
pub repo_chosen: bool,
}
impl Target {
pub fn audience(&self) -> Result<String, String> {
audience_from_url(&self.endpoint).ok_or_else(|| {
format!(
"`{}` is not an mkit+https:// (or loopback mkit+http://) remote",
self.endpoint
)
})
}
#[must_use]
pub fn repository(&self) -> Option<RepositoryIdentity> {
repository_identity_from_url(&self.endpoint)
.ok()
.filter(|id| id.namespace().is_some())
}
#[must_use]
pub fn namespace(&self) -> Option<Namespace> {
self.repository().and_then(|id| id.namespace().copied())
}
#[must_use]
pub fn is_loopback_dev(&self) -> bool {
self.endpoint.starts_with("mkit+http://")
&& self.audience().is_ok_and(|a| is_loopback_origin(&a))
}
}
pub fn resolve_target(cfg: &LayeredConfig, arg: Option<&str>) -> Result<Target, String> {
match arg {
Some(url) if url.starts_with("mkit+") => Ok(Target {
name: url.to_owned(),
endpoint: url.to_owned(),
repo_chosen: false,
}),
Some(name) => {
let resolved = config::resolve_remote(cfg, name).ok_or_else(|| {
format!("no remote named `{name}` (add one with `mkit remote add`, or pass an mkit+https:// URL)")
})?;
Ok(Target {
name: resolved.name,
endpoint: resolved.endpoint,
repo_chosen: resolved.repo_chosen,
})
}
None => {
let trusted = cfg.user.trusted_remote_endpoint.trim();
if trusted.is_empty() {
return Err(
"no remote given and no trusted remote configured (set one with `mkit config trusted_remote_endpoint <url>`)"
.to_owned(),
);
}
Ok(Target {
name: "trusted".to_owned(),
endpoint: trusted.to_owned(),
repo_chosen: false,
})
}
}
}
pub fn check_audiences(audiences: &[String], remote: Option<&Target>) -> Result<(), String> {
for audience in audiences {
if is_loopback_origin(audience) && !remote.is_some_and(Target::is_loopback_dev) {
return Err(format!(
"audience {audience} is a loopback address, which every local server shares; \
refusing to sign for it unless the remote is itself a loopback mkit+http:// development remote"
));
}
}
Ok(())
}
#[derive(Debug, PartialEq, Eq)]
pub enum Driven<T> {
Done(T),
TimedOut {
waited: Duration,
},
Cancelled,
}
pub fn drive<T, E>(
call: impl FnMut() -> Result<Completion<T>, E>,
timeout: Duration,
sleep: impl FnMut(Duration) -> bool,
) -> Result<Driven<T>, E> {
let start = std::time::Instant::now();
drive_with_clock(call, timeout, sleep, || start.elapsed())
}
fn drive_with_clock<T, E>(
mut call: impl FnMut() -> Result<Completion<T>, E>,
timeout: Duration,
mut sleep: impl FnMut(Duration) -> bool,
elapsed: impl Fn() -> Duration,
) -> Result<Driven<T>, E> {
loop {
match call()? {
Completion::Done(value) => return Ok(Driven::Done(value)),
Completion::Pending { retry_after } => {
let waited = elapsed();
if waited.saturating_add(retry_after) > timeout {
return Ok(Driven::TimedOut { waited });
}
if !sleep(retry_after) {
return Ok(Driven::Cancelled);
}
}
}
}
}
#[must_use]
pub fn interruptible_sleep(duration: Duration) -> bool {
let slice = Duration::from_millis(100);
let mut left = duration;
while !left.is_zero() {
if crate::signal::is_shutdown() {
return false;
}
let step = left.min(slice);
std::thread::sleep(step);
left -= step;
}
!crate::signal::is_shutdown()
}
#[cfg(test)]
mod tests {
use super::*;
fn target(url: &str) -> Target {
Target {
name: "t".into(),
endpoint: url.into(),
repo_chosen: false,
}
}
#[test]
fn loopback_audiences_need_a_loopback_dev_remote() {
let audiences = |a: &str| vec![a.to_owned()];
let dev = target("mkit+http://127.0.0.1:8080");
let prod = target("mkit+https://git.example.com");
for loopback in [
"http://localhost:8080",
"http://127.0.0.1:8080",
"http://127.9.9.9",
"http://[::1]:8080",
] {
assert!(
check_audiences(&audiences(loopback), Some(&prod)).is_err(),
"{loopback}"
);
assert!(
check_audiences(&audiences(loopback), None).is_err(),
"{loopback}"
);
}
assert!(check_audiences(&audiences("http://127.0.0.1:8080"), Some(&dev)).is_ok());
assert!(check_audiences(&audiences("https://git.example.com"), Some(&prod)).is_ok());
assert!(target("mkit+http://127.0.0.1:8080").is_loopback_dev());
assert!(!target("mkit+https://git.example.com").is_loopback_dev());
}
#[test]
fn audience_and_namespace_come_from_the_url() {
let t =
target("mkit+https://git.example.com/0x8ba1f109551bd432803012645ac136ddd64dba72/site");
assert_eq!(t.audience().unwrap(), "https://git.example.com");
assert_eq!(
t.namespace().unwrap().to_string(),
"0x8ba1f109551bd432803012645ac136ddd64dba72"
);
assert!(target("mkit+https://git.example.com").namespace().is_none());
assert!(target("https://git.example.com").audience().is_err());
}
#[test]
fn drive_resends_until_done_and_sums_the_waits() {
let mut calls: u64 = 0;
let mut slept = Vec::new();
let clock = std::cell::Cell::new(Duration::ZERO);
let outcome: Result<Driven<u64>, ()> = drive_with_clock(
|| {
calls += 1;
Ok(if calls < 4 {
Completion::Pending {
retry_after: Duration::from_secs(calls),
}
} else {
Completion::Done(9)
})
},
Duration::from_mins(1),
|d| {
slept.push(d);
clock.set(clock.get() + d);
true
},
|| clock.get(),
);
assert_eq!(outcome, Ok(Driven::Done(9)));
assert_eq!(calls, 4);
assert_eq!(
slept,
[
Duration::from_secs(1),
Duration::from_secs(2),
Duration::from_secs(3)
]
);
}
#[test]
fn drive_stops_at_the_total_timeout_and_on_cancel() {
let pending = || -> Result<Completion<()>, ()> {
Ok(Completion::Pending {
retry_after: Duration::from_secs(4),
})
};
let mut sleeps = 0;
let clock = std::cell::Cell::new(Duration::ZERO);
let outcome = drive_with_clock(
pending,
Duration::from_secs(10),
|d| {
sleeps += 1;
clock.set(clock.get() + d);
true
},
|| clock.get(),
);
assert_eq!(
outcome,
Ok(Driven::TimedOut {
waited: Duration::from_secs(8)
})
);
assert_eq!(sleeps, 2);
let outcome = drive(pending, Duration::from_secs(10), |_| false);
assert_eq!(outcome, Ok(Driven::Cancelled));
let clock = std::cell::Cell::new(Duration::ZERO);
let outcome = drive_with_clock(
|| {
clock.set(clock.get() + Duration::from_secs(7));
pending()
},
Duration::from_secs(10),
|_| true,
|| clock.get(),
);
assert_eq!(
outcome,
Ok(Driven::TimedOut {
waited: Duration::from_secs(7)
})
);
let outcome: Result<Driven<()>, &str> =
drive(|| Err("boom"), Duration::from_secs(1), |_| true);
assert_eq!(outcome, Err("boom"));
}
}