mod transport;
mod types;
use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
use pb_mapper_auth::{
discard_staged_admin_key, generate_admin_key, stage_admin_key_candidate, write_admin_key_file,
};
use pb_mapper_core::checksum::parse_credential;
use pb_mapper_core::config::control_io_timeout;
use pb_mapper_core::paging::MAX_PAGE_SIZE;
use pb_mapper_protocol::command::{AdminRequest, AdminResponse};
use self::transport::send_admin_request;
use self::types::Paged;
use super::Error;
use super::client::ClientInner;
use super::error::Result;
use super::types::LegacyProtocol;
pub use self::types::{
AuthStatusInfo, ConnectionInfo, ConnectionPage, IssuedKey, KeyListPage, KeyMetadata,
ServiceInfo, ServicePage,
};
const COLLECT_PAGE_SIZE: u16 = MAX_PAGE_SIZE;
const MAX_PAGES: u32 = 10_000;
#[derive(Clone)]
pub struct Admin {
pub(crate) inner: Arc<ClientInner>,
}
macro_rules! admin_rpc {
(
$(#[$doc:meta])*
$name:ident($($arg:ident: $arg_ty:ty),* $(,)?)
-> $output:ty,
request: $request:expr,
response: $variant:ident($binding:pat) => $value:expr
) => {
$(#[$doc])*
pub async fn $name(&self, $($arg: $arg_ty),*) -> Result<$output> {
match self.request($request).await? {
AdminResponse::$variant($binding) => Ok($value),
other => unexpected(stringify!($variant), &other),
}
}
};
}
macro_rules! admin_rpc_struct {
(
$(#[$doc:meta])*
$name:ident($($arg:ident: $arg_ty:ty),* $(,)?)
-> $output:ty,
request: $request:expr,
response: $variant:ident { $($binding:tt)* } => $value:expr
) => {
$(#[$doc])*
pub async fn $name(&self, $($arg: $arg_ty),*) -> Result<$output> {
match self.request($request).await? {
AdminResponse::$variant { $($binding)* } => Ok($value),
other => unexpected(stringify!($variant), &other),
}
}
};
}
impl Admin {
pub async fn request(&self, request: AdminRequest) -> Result<AdminResponse> {
self.request_with_timeout(request, control_io_timeout())
.await
}
pub async fn request_with_timeout(
&self,
request: AdminRequest,
io_timeout: Duration,
) -> Result<AdminResponse> {
send_admin_request(&self.inner.server, self.credential(), request, io_timeout).await
}
admin_rpc!(
issue_key(ttl: Duration, label: Option<String>) -> IssuedKey,
request: AdminRequest::KeyIssue { ttl_seconds: ttl.as_secs(), label },
response: KeyIssued(issued) => IssuedKey::from(issued)
);
admin_rpc!(
show_key(key_id: u64) -> IssuedKey,
request: AdminRequest::KeyShow { key_id },
response: KeyShown(issued) => IssuedKey::from(issued)
);
admin_rpc!(
reveal_key(key_id: u64) -> IssuedKey,
request: AdminRequest::KeyReveal { key_id },
response: KeyShown(issued) => IssuedKey::from(issued)
);
admin_rpc!(
renew_key(key_id: u64, ttl: Duration) -> IssuedKey,
request: AdminRequest::KeyRenew { key_id, ttl_seconds: ttl.as_secs() },
response: KeyRenewed(issued) => IssuedKey::from(issued)
);
admin_rpc!(
revoke_key(key_id: u64) -> KeyMetadata,
request: AdminRequest::KeyRevoke { key_id },
response: KeyRevoked(meta) => KeyMetadata::from(meta)
);
admin_rpc!(
auth_status() -> AuthStatusInfo,
request: AdminRequest::AuthStatus,
response: AuthStatus(status) => AuthStatusInfo::from(status)
);
admin_rpc_struct!(
gc_keys() -> u64,
request: AdminRequest::KeyGc,
response: KeyGc { removed } => removed
);
admin_rpc_struct!(
reset_auth_state() -> (),
request: AdminRequest::AuthStateReset { confirm: true },
response: Ok { .. } => ()
);
admin_rpc_struct!(
set_legacy_protocol(policy: LegacyProtocol) -> (),
request: AdminRequest::LegacyProtocolSet { policy: policy.into() },
response: Ok { .. } => ()
);
admin_rpc!(
list_keys(page: u32, page_size: u16) -> KeyListPage,
request: AdminRequest::KeyList { page, page_size: validate_page_size(page_size)? },
response: KeyList(page) => KeyListPage::from(page)
);
admin_rpc!(
list_services(key_id: Option<u64>, page: u32, page_size: u16) -> ServicePage,
request: AdminRequest::ServiceList {
key_id,
page,
page_size: validate_page_size(page_size)?,
},
response: Services(page) => ServicePage::from(page)
);
admin_rpc!(
list_connections(key_id: Option<u64>, page: u32, page_size: u16) -> ConnectionPage,
request: AdminRequest::ConnectionList {
key_id,
page,
page_size: validate_page_size(page_size)?,
},
response: Connections(page) => ConnectionPage::from(page)
);
admin_rpc_struct!(
retire_connections(key_id: Option<u64>, service_name: String, conn_id: Option<u32>) -> u32,
request: AdminRequest::ConnectionRetire { key_id, service_name, conn_id },
response: ConnectionsRetired { retired } => retired
);
pub async fn list_keys_all(&self) -> Result<Vec<KeyMetadata>> {
collect_pages(|page| self.list_keys(page, COLLECT_PAGE_SIZE)).await
}
pub async fn list_services_all(&self, key_id: Option<u64>) -> Result<Vec<ServiceInfo>> {
collect_pages(|page| self.list_services(key_id, page, COLLECT_PAGE_SIZE)).await
}
pub async fn list_connections_all(&self, key_id: Option<u64>) -> Result<Vec<ConnectionInfo>> {
collect_pages(|page| self.list_connections(key_id, page, COLLECT_PAGE_SIZE)).await
}
pub async fn rotate_root_key(&self, new_key: Option<String>) -> Result<String> {
let caller_supplied = new_key.is_some();
let new_key = new_key.unwrap_or_else(generate_admin_key);
let parsed = parse_credential(new_key.trim()).map_err(Error::invalid_config)?;
if !parsed.is_admin() {
return Err(Error::invalid_config(
"root rotation requires a 32-byte administrator key",
));
}
let preserve = |error: Error| {
if caller_supplied {
return error;
}
Error::RootRotationUncertain {
candidate: new_key.clone(),
message: error.to_string(),
}
};
let response = self
.request(AdminRequest::RootKeyRotate {
new_admin_key: new_key.clone(),
})
.await
.map_err(preserve)?;
match response {
AdminResponse::Ok { .. } => {
*self
.inner
.credential
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = parsed;
Ok(new_key)
}
other => unexpected::<String>("Ok", &other).map_err(preserve),
}
}
pub async fn rotate_root_key_to_file(
&self,
path: &Path,
new_key: Option<String>,
) -> Result<String> {
let new_key = new_key.unwrap_or_else(generate_admin_key);
let staged_path = stage_admin_key_candidate(path, &new_key).map_err(auth_file_error)?;
let staged_note = || format!("the candidate key remains at `{}`", staged_path.display());
let rotated = self.rotate_root_key(Some(new_key)).await.map_err(|error| {
auth_file_message(format!(
"root rotation request failed; {}: {error}",
staged_note()
))
})?;
self.auth_status().await.map_err(|error| {
auth_file_message(format!(
"new administrator key did not pass the post-rotation status check; {}: {error}",
staged_note()
))
})?;
write_admin_key_file(path, &rotated, true).map_err(|error| {
auth_file_message(format!(
"administrator key rotated and verified, but `{}` could not be updated; recover \
the key from `{}`: {error}",
path.display(),
staged_path.display()
))
})?;
discard_staged_admin_key(&staged_path);
Ok(rotated)
}
fn credential(&self) -> pb_mapper_core::checksum::Credential {
*self
.inner
.credential
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
}
fn auth_file_error(error: pb_mapper_auth::AuthFailure) -> Error {
auth_file_message(error.to_string())
}
fn auth_file_message(message: impl Into<String>) -> Error {
Error::AuthFile {
message: message.into(),
}
}
fn validate_page_size(page_size: u16) -> Result<u16> {
if !(1..=MAX_PAGE_SIZE).contains(&page_size) {
return Err(Error::invalid_config(format!(
"page_size must be between 1 and {MAX_PAGE_SIZE}"
)));
}
Ok(page_size)
}
fn unexpected<T>(expected: &str, actual: &AdminResponse) -> Result<T> {
Err(Error::protocol(format!(
"expected {expected}, got {actual:?}"
)))
}
async fn collect_pages<P, F, Fut>(mut fetch: F) -> Result<Vec<P::Item>>
where
P: Paged,
F: FnMut(u32) -> Fut,
Fut: std::future::Future<Output = Result<P>>,
{
let mut page = 0_u32;
let mut items = Vec::new();
for _ in 0..MAX_PAGES {
let (chunk, next) = fetch(page).await?.into_parts();
items.extend(chunk);
match next {
Some(next_page) => page = next_page,
None => return Ok(items),
}
}
Err(Error::protocol(format!(
"pagination exceeded {MAX_PAGES} pages"
)))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sdk::admin::types::KeyListPage;
fn page(next_page: Option<u32>, ids: &[u64]) -> KeyListPage {
KeyListPage {
schema_version: 1,
items: ids
.iter()
.map(|&key_id| KeyMetadata {
key_id,
state: "active".into(),
issued_at: 0,
expires_at: 0,
label: None,
})
.collect(),
next_page,
}
}
#[test]
fn the_collect_page_size_can_drain_a_relay_at_capacity() {
assert_eq!(
COLLECT_PAGE_SIZE, MAX_PAGE_SIZE,
"paging below the relay's maximum multiplies its per-page re-sort \
and shrinks what `*_all` can drain"
);
let pages_at_capacity =
pb_mapper_auth::MAX_TEMP_KEY_CAPACITY.div_ceil(usize::from(COLLECT_PAGE_SIZE));
assert!(
pages_at_capacity <= MAX_PAGES as usize,
"a full inventory needs {pages_at_capacity} pages, over the {MAX_PAGES} cap"
);
}
#[test]
fn a_page_size_outside_the_relays_range_is_rejected_locally() {
assert!(validate_page_size(0).is_err());
assert!(validate_page_size(MAX_PAGE_SIZE + 1).is_err());
assert_eq!(validate_page_size(1).unwrap(), 1);
assert_eq!(validate_page_size(MAX_PAGE_SIZE).unwrap(), MAX_PAGE_SIZE);
}
#[tokio::test]
async fn collect_pages_follows_the_cursor_and_concatenates() {
let requested = std::sync::Mutex::new(Vec::new());
let items = collect_pages::<KeyListPage, _, _>(|page_number| {
requested.lock().expect("not poisoned").push(page_number);
async move {
Ok(match page_number {
0 => page(Some(7), &[1, 2]),
7 => page(Some(9), &[3]),
_ => page(None, &[4]),
})
}
})
.await
.expect("three pages is well inside the cap");
assert_eq!(
items.iter().map(|item| item.key_id).collect::<Vec<_>>(),
vec![1, 2, 3, 4]
);
assert_eq!(
*requested.lock().expect("not poisoned"),
vec![0, 7, 9],
"the cursor the relay hands back is what gets asked for next"
);
}
#[tokio::test]
async fn collect_pages_gives_up_on_a_cursor_that_never_ends() {
let calls = std::sync::atomic::AtomicU32::new(0);
let error = collect_pages::<KeyListPage, _, _>(|page_number| {
calls.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
async move { Ok(page(Some(page_number + 1), &[u64::from(page_number)])) }
})
.await
.expect_err("an endless cursor must fail rather than hang");
assert!(error.to_string().contains("pagination exceeded"));
assert_eq!(
calls.load(std::sync::atomic::Ordering::Relaxed),
MAX_PAGES,
"the cap is what stops it, and it stops exactly there"
);
}
}