use std::collections::HashMap;
use std::sync::Arc;
use arc_swap::ArcSwap;
use axum::extract::{FromRequestParts, MatchedPath, OriginalUri, Request, State};
use axum::http::{header, request::Parts};
use axum::middleware::Next;
use axum::response::Response;
use jiff::Timestamp;
use serde::{Deserialize, Serialize};
use tollgate_auth::CredentialVerifier;
use tollgate_core::{
AccountId, AccountSnapshot, AccountStatus, CostTable, CostUnits, EnforcementMode, Generation,
PermissionBits, PolicyRevision, Principal, PublishableSnapshot, ResolvedLimits,
};
use tollgate_store::Clock;
use crate::error::ApiError;
use crate::transport::{PeerIdentity, TlsConfig};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Role {
Instance,
Operator,
Provisioner,
}
impl Role {
pub fn as_str(self) -> &'static str {
match self {
Role::Instance => "instance",
Role::Operator => "operator",
Role::Provisioner => "provisioner",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ProvisionerPolicyTemplate {
pub cost_table: Arc<CostTable>,
pub limits: ResolvedLimits,
pub permissions: PermissionBits,
#[serde(default)]
pub policy_revision: PolicyRevision,
}
impl ProvisionerPolicyTemplate {
pub(crate) fn validate(&self) -> Result<(), SecurityError> {
let snapshot = AccountSnapshot::builder(
AccountId(0),
Generation(0),
AccountStatus::Active,
Timestamp::UNIX_EPOCH,
self.permissions,
self.limits,
Arc::clone(&self.cost_table),
)
.build();
PublishableSnapshot::try_new(Arc::new(snapshot))
.map(|_| ())
.map_err(|_| SecurityError("invalid provisioner policy template"))
}
fn matches(&self, snapshot: &AccountSnapshot) -> bool {
snapshot.enforcement_mode == EnforcementMode::Strict
&& snapshot.cost_table == self.cost_table
&& snapshot.limits == self.limits
&& snapshot.permissions == self.permissions
&& snapshot.policy_revision == self.policy_revision
}
}
#[derive(Debug, Clone, Eq)]
pub struct ProvisionerLimits {
max_budget_allowance: CostUnits,
policy_templates: Arc<[ProvisionerPolicyTemplate]>,
}
impl PartialEq for ProvisionerLimits {
fn eq(&self, other: &Self) -> bool {
self.max_budget_allowance == other.max_budget_allowance
&& self
.policy_templates
.iter()
.all(|policy| other.policy_templates.contains(policy))
&& other
.policy_templates
.iter()
.all(|policy| self.policy_templates.contains(policy))
}
}
impl ProvisionerLimits {
pub fn new(
max_budget_allowance: CostUnits,
policy_templates: Vec<ProvisionerPolicyTemplate>,
) -> Result<Self, SecurityError> {
if policy_templates.is_empty() {
return Err(SecurityError(
"a provisioner identity requires approved policy templates",
));
}
for template in &policy_templates {
template.validate()?;
}
Ok(Self {
max_budget_allowance,
policy_templates: policy_templates.into(),
})
}
pub fn max_budget_allowance(&self) -> CostUnits {
self.max_budget_allowance
}
pub fn allows_snapshot(&self, snapshot: &AccountSnapshot) -> bool {
self.policy_templates
.iter()
.any(|template| template.matches(snapshot))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum Grant {
Instance,
Operator,
Provisioner(ProvisionerLimits),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ControlIdentity {
name: String,
grant: Grant,
}
impl ControlIdentity {
pub fn new(name: impl Into<String>, role: Role) -> Result<Self, SecurityError> {
let grant = match role {
Role::Instance => Grant::Instance,
Role::Operator => Grant::Operator,
Role::Provisioner => {
return Err(SecurityError(
"a provisioner identity requires budget and approved policy limits",
));
}
};
Self::with_grant(name.into(), grant)
}
pub fn provisioner(
name: impl Into<String>,
limits: ProvisionerLimits,
) -> Result<Self, SecurityError> {
Self::with_grant(name.into(), Grant::Provisioner(limits))
}
fn with_grant(name: String, grant: Grant) -> Result<Self, SecurityError> {
if name.is_empty()
|| name.len() > 128
|| !name
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b"@._:/-".contains(&b))
{
return Err(SecurityError(
"identity must be 1..=128 ASCII identifier characters",
));
}
Ok(Self { name, grant })
}
pub fn name(&self) -> &str {
&self.name
}
pub fn role(&self) -> Role {
match &self.grant {
Grant::Instance => Role::Instance,
Grant::Operator => Role::Operator,
Grant::Provisioner(_) => Role::Provisioner,
}
}
pub fn provisioner_limits(&self) -> Option<ProvisionerLimits> {
match &self.grant {
Grant::Provisioner(limits) => Some(limits.clone()),
Grant::Instance | Grant::Operator => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SecurityError(pub &'static str);
impl std::fmt::Display for SecurityError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.0)
}
}
impl std::error::Error for SecurityError {}
struct BearerScheme {
verifier: Arc<dyn CredentialVerifier + Send + Sync>,
identities: HashMap<Principal, ControlIdentity>,
}
#[derive(Default)]
pub struct SecurityPolicy {
bearers: Vec<BearerScheme>,
certificates: HashMap<[u8; 32], ControlIdentity>,
}
impl SecurityPolicy {
pub fn new() -> Self {
Self::default()
}
pub fn with_bearer(
mut self,
verifier: Arc<dyn CredentialVerifier + Send + Sync>,
identities: impl IntoIterator<Item = (Principal, ControlIdentity)>,
) -> Result<Self, SecurityError> {
let mut mapped = HashMap::new();
for (principal, identity) in identities {
if mapped.insert(principal, identity).is_some() {
return Err(SecurityError("duplicate bearer principal"));
}
}
self.bearers.push(BearerScheme {
verifier,
identities: mapped,
});
Ok(self)
}
pub fn with_certificate(
mut self,
fingerprint: [u8; 32],
identity: ControlIdentity,
) -> Result<Self, SecurityError> {
if self.certificates.insert(fingerprint, identity).is_some() {
return Err(SecurityError("duplicate client certificate"));
}
Ok(self)
}
fn bearer(&self, credential: &[u8], now: Timestamp) -> Result<ControlIdentity, ApiError> {
let mut identity = None;
for scheme in &self.bearers {
if let Some(proof) = scheme.verifier.verify(credential) {
if !proof.is_reusable_at(now) {
return Err(ApiError::unauthorized());
}
let found = scheme
.identities
.get(&proof.principal)
.ok_or_else(ApiError::forbidden)?;
if identity.as_ref().is_some_and(|previous| previous != found) {
return Err(ApiError::unauthorized());
}
identity = Some(found.clone());
}
}
identity.ok_or_else(ApiError::unauthorized)
}
}
pub(crate) struct SecurityBundle {
pub policy: SecurityPolicy,
pub tls: Option<TlsConfig>,
}
pub struct ServerSecurity {
pub(crate) current: ArcSwap<SecurityBundle>,
encrypted: bool,
}
impl ServerSecurity {
pub fn new(policy: SecurityPolicy, tls: Option<TlsConfig>) -> Result<Arc<Self>, SecurityError> {
validate(&policy, tls.as_ref())?;
Ok(Arc::new(Self {
encrypted: tls.is_some(),
current: ArcSwap::from_pointee(SecurityBundle { policy, tls }),
}))
}
pub fn replace(
&self,
policy: SecurityPolicy,
tls: Option<TlsConfig>,
) -> Result<(), SecurityError> {
validate(&policy, tls.as_ref())?;
if self.encrypted != tls.is_some() {
return Err(SecurityError(
"changing TLS mode requires a listener restart",
));
}
self.current.store(Arc::new(SecurityBundle { policy, tls }));
Ok(())
}
pub fn encrypted(&self) -> bool {
self.encrypted
}
}
fn validate(policy: &SecurityPolicy, tls: Option<&TlsConfig>) -> Result<(), SecurityError> {
if !policy.certificates.is_empty() && tls.is_none_or(|tls| !tls.verifies_clients()) {
return Err(SecurityError(
"certificate identities require TLS with a client CA",
));
}
Ok(())
}
#[derive(Clone)]
pub(crate) struct Authorization {
pub security: Arc<ServerSecurity>,
pub clock: Arc<dyn Clock>,
pub roles: &'static [Role],
}
fn bearer(request: &Request) -> Result<Option<&[u8]>, ApiError> {
let mut values = request.headers().get_all(header::AUTHORIZATION).iter();
let Some(header) = values.next() else {
return Ok(None);
};
if values.next().is_some() {
return Err(ApiError::unauthorized());
}
let header = header.as_bytes();
if header.len() > 16 * 1024 {
return Err(ApiError::unauthorized());
}
let Some(separator) = header.iter().position(|b| *b == b' ') else {
return Err(ApiError::unauthorized());
};
let (scheme, token) = (&header[..separator], &header[separator + 1..]);
if !scheme.eq_ignore_ascii_case(b"Bearer")
|| token.is_empty()
|| !token.iter().all(|b| b.is_ascii_graphic())
{
return Err(ApiError::unauthorized());
}
Ok(Some(token))
}
pub(crate) async fn authorize(
State(auth): State<Authorization>,
mut request: Request,
next: Next,
) -> Result<Response, ApiError> {
let bundle = auth.security.current.load_full();
let now = auth.clock.now();
let bearer = bearer(&request)?;
let peer = request
.extensions()
.get::<axum::extract::ConnectInfo<PeerIdentity>>();
let certificate = match peer.and_then(|peer| peer.0.certificates.as_deref()) {
Some(chain) => {
let tls = bundle.tls.as_ref().ok_or_else(ApiError::unauthorized)?;
tls.verify(chain, now)
.map_err(|_| ApiError::unauthorized())?;
let fingerprint = crate::transport::fingerprint(&chain[0]);
Some(
bundle
.policy
.certificates
.get(&fingerprint)
.ok_or_else(ApiError::forbidden)?
.clone(),
)
}
None => None,
};
let token = bearer
.map(|token| bundle.policy.bearer(token, now))
.transpose()?;
let identity = match (token, certificate) {
(Some(a), Some(b)) if a != b => return Err(ApiError::unauthorized()),
(Some(identity), _) | (_, Some(identity)) => identity,
(None, None) => return Err(ApiError::unauthorized()),
};
if !auth.roles.contains(&identity.role()) {
let action = request
.extensions()
.get::<MatchedPath>()
.map_or("unmatched", MatchedPath::as_str);
tracing::warn!(target: "tollgate::audit", actor = identity.name(),
role = identity.role().as_str(), action = %format_args!("{} {action}", request.method()),
resource = request.extensions().get::<OriginalUri>().map_or_else(|| request.uri().path(), |uri| uri.path()),
at = %now, outcome = "refused",
code = "scope-forbidden", "control-plane request refused for its role");
return Err(ApiError::forbidden());
}
tracing::debug!(
actor = identity.name(),
role = identity.role().as_str(),
"control-plane request authenticated"
);
match identity.grant.clone() {
Grant::Instance => {
request.extensions_mut().insert(InstanceIdentity);
}
Grant::Operator => {
request
.extensions_mut()
.insert(AdminIdentity::Operator(identity.clone()));
request.extensions_mut().insert(OperatorIdentity(identity));
}
Grant::Provisioner(limits) => {
request
.extensions_mut()
.insert(AdminIdentity::Provisioner(identity, limits));
}
}
Ok(next.run(request).await)
}
#[derive(Clone)]
pub(crate) struct InstanceIdentity;
#[derive(Clone)]
pub(crate) struct OperatorIdentity(pub ControlIdentity);
#[derive(Clone)]
pub(crate) enum AdminIdentity {
Operator(ControlIdentity),
Provisioner(ControlIdentity, ProvisionerLimits),
}
macro_rules! identity_extractor {
($name:ident) => {
impl<S: Send + Sync> FromRequestParts<S> for $name {
type Rejection = ApiError;
async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, ApiError> {
parts
.extensions
.get::<Self>()
.cloned()
.ok_or_else(ApiError::unauthorized)
}
}
};
}
identity_extractor!(InstanceIdentity);
identity_extractor!(OperatorIdentity);
identity_extractor!(AdminIdentity);
#[cfg(test)]
mod framing_tests {
use super::*;
#[test]
fn bearer_framing_checks_each_condition_before_scheme_verification() {
let request = |value: &[u8]| {
Request::builder()
.header(
header::AUTHORIZATION,
axum::http::HeaderValue::from_bytes(value).unwrap(),
)
.body(axum::body::Body::empty())
.unwrap()
};
for value in [
b"Basic token".as_slice(),
b"Bearer ",
b"Bearer one\ttwo",
b"Bearer token",
b"Bearer",
b"Bearer \x80",
] {
assert!(bearer(&request(value)).is_err());
}
for size in [1, 16 * 1024 - 7] {
let token = vec![b'x'; size];
let mut header = b"bEaReR ".to_vec();
header.extend_from_slice(&token);
let request = request(&header);
assert_eq!(bearer(&request).unwrap(), Some(token.as_slice()));
}
for size in [16 * 1024 + 1, 32 * 1024] {
let mut value = b"Bearer ".to_vec();
value.resize(size, b'x');
assert!(bearer(&request(&value)).is_err());
}
let mut duplicate = request(b"Bearer token");
duplicate
.headers_mut()
.append(header::AUTHORIZATION, "Bearer token".parse().unwrap());
assert!(bearer(&duplicate).is_err());
assert_eq!(
bearer(&Request::new(axum::body::Body::empty())).unwrap(),
None
);
}
#[test]
fn overlapping_bearer_schemes_must_agree_on_the_verified_identity() {
use tollgate_auth::{HmacRegistry, Verified};
let verifier = Arc::new(HmacRegistry::new(b"fixture-scheme-agreement-secret"));
let principal = verifier.install_credentials([b"shared-credential".as_slice()])[0];
let now = Timestamp::from_second(100).unwrap();
let instance = ControlIdentity::new("instance", Role::Instance).unwrap();
let operator = ControlIdentity::new("operator", Role::Operator).unwrap();
for (second, agrees) in [(instance.clone(), true), (operator, false)] {
let policy = SecurityPolicy::new()
.with_bearer(verifier.clone(), [(principal, instance.clone())])
.unwrap()
.with_bearer(verifier.clone(), [(principal, second)])
.unwrap();
assert_eq!(policy.bearer(b"shared-credential", now).is_ok(), agrees);
if agrees {
assert_eq!(policy.bearer(b"shared-credential", now).unwrap(), instance);
}
}
struct Expired(Principal);
impl CredentialVerifier for Expired {
fn verify(&self, _: &[u8]) -> Option<Verified> {
Some(Verified::until(
self.0,
Timestamp::from_second(100).unwrap(),
))
}
}
let policy = SecurityPolicy::new()
.with_bearer(Arc::new(Expired(principal)), [(principal, instance)])
.unwrap();
assert!(policy.bearer(b"shared-credential", now).is_err());
}
}
impl OperatorIdentity {
pub(crate) async fn run<T, E: Into<ApiError>>(
&self,
action: &'static str,
target: impl std::fmt::Display,
clock: &dyn Clock,
operation: impl Future<Output = Result<tollgate_store::AdminReceipt<T>, E>>,
) -> Result<T, ApiError> {
audited(&self.0, action, target, clock, operation).await
}
}
impl AdminIdentity {
pub(crate) fn identity(&self) -> &ControlIdentity {
match self {
AdminIdentity::Operator(identity) | AdminIdentity::Provisioner(identity, _) => identity,
}
}
pub(crate) async fn run<T, E: Into<ApiError>>(
&self,
action: &'static str,
target: impl std::fmt::Display,
clock: &dyn Clock,
operation: impl Future<Output = Result<tollgate_store::AdminReceipt<T>, E>>,
) -> Result<T, ApiError> {
audited(self.identity(), action, target, clock, operation).await
}
pub(crate) fn refuse(
&self,
action: &'static str,
target: impl std::fmt::Display,
clock: &dyn Clock,
code: &'static str,
title: &'static str,
) -> ApiError {
let identity = self.identity();
tracing::warn!(target: "tollgate::audit", actor = identity.name(),
role = identity.role().as_str(), action, resource = %target, at = %clock.now(),
outcome = "refused", code, reason = title, "administrative operation refused for its scope");
ApiError::refused_scope(code, title)
}
pub(crate) async fn check_account<S: tollgate_store::AdminStore + ?Sized>(
&self,
store: &S,
account: tollgate_core::AccountId,
action: &'static str,
target: impl std::fmt::Display,
clock: &dyn Clock,
) -> Result<(), ApiError> {
if let AdminIdentity::Operator(_) = self {
return Ok(());
}
match store.account_view(account).await? {
None => Err(tollgate_store::SetStatusError::UnknownAccount.into()),
Some(view) if view.origin == tollgate_store::AdminAuthority::Provisioner => Ok(()),
Some(_) => Err(self.refuse(
action,
target,
clock,
"account-not-provisioned",
"account was not created by a provisioner",
)),
}
}
}
async fn audited<T, E: Into<ApiError>>(
identity: &ControlIdentity,
action: &'static str,
target: impl std::fmt::Display,
clock: &dyn Clock,
operation: impl Future<Output = Result<tollgate_store::AdminReceipt<T>, E>>,
) -> Result<T, ApiError> {
struct Attempt<'a> {
id: tollgate_core::RequestId,
identity: &'a ControlIdentity,
action: &'static str,
target: String,
clock: &'a dyn Clock,
finished: bool,
}
impl Drop for Attempt<'_> {
fn drop(&mut self) {
if !self.finished {
tracing::warn!(target: "tollgate::audit", actor = self.identity.name(),
role = self.identity.role().as_str(), action = self.action,
operation_id = %self.id,
resource = self.target, at = %self.clock.now(), outcome = "cancelled_unknown",
"administrative operation abandoned; commit outcome may be unknown");
}
}
}
let mut identifier = [0u8; 16];
getrandom::fill(&mut identifier).map_err(|_| {
ApiError::from(tollgate_store::StoreError(
"audit identity entropy unavailable".into(),
))
})?;
let mut attempt = Attempt {
id: tollgate_core::RequestId(u128::from_be_bytes(identifier)),
identity,
action,
target: target.to_string(),
clock,
finished: false,
};
tracing::info!(target: "tollgate::audit", actor = identity.name(),
role = identity.role().as_str(), action,
operation_id = %attempt.id,
resource = attempt.target, at = %clock.now(), outcome = "started", "administrative operation started");
let result = operation.await;
attempt.finished = true;
match result {
Ok(receipt) => {
tracing::info!(target: "tollgate::audit", actor = identity.name(),
role = identity.role().as_str(), action,
operation_id = %attempt.id,
resource = attempt.target, at = %clock.now(), outcome = "confirmed",
before = ?receipt.before, after = ?receipt.after, "administrative operation completed");
Ok(receipt.outcome)
}
Err(error) => {
let error = error.into();
tracing::warn!(target: "tollgate::audit", actor = identity.name(),
role = identity.role().as_str(), action,
operation_id = %attempt.id,
resource = attempt.target, at = %clock.now(), outcome = "failed",
code = error.code, status = error.status.as_u16(), "administrative operation failed; storage errors may conceal a commit");
Err(error)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn cancelled_admin_operations_report_an_unknown_commit_without_a_receipt() {
use tracing::instrument::WithSubscriber;
#[derive(Clone, Default)]
struct Capture(Arc<std::sync::Mutex<Vec<u8>>>);
impl std::io::Write for Capture {
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
let capture = Capture::default();
let output = capture.clone();
let subscriber = tracing_subscriber::fmt()
.without_time()
.with_ansi(false)
.with_writer(move || output.clone())
.finish();
let identity =
OperatorIdentity(ControlIdentity::new("fixture-operator", Role::Operator).unwrap());
async {
let operation = identity.run(
"deposit",
"fixture-account",
&tollgate_store::SystemClock,
std::future::pending::<
Result<tollgate_store::AdminReceipt<()>, tollgate_store::StoreError>,
>(),
);
tokio::pin!(operation);
tokio::select! {
biased;
_ = &mut operation => panic!("the backend remains pending"),
_ = tokio::task::yield_now() => {},
}
}
.with_subscriber(subscriber)
.await;
let bytes = capture.0.lock().unwrap();
let log = std::str::from_utf8(&bytes).unwrap();
assert!(
log.contains("started") && log.contains("cancelled_unknown"),
"{log}"
);
assert!(
log.contains("fixture-operator")
&& log.contains("fixture-account")
&& log.contains("operation_id=")
);
assert!(!log.contains("confirmed") && !log.contains("before=") && !log.contains("after="));
}
}
#[cfg(test)]
mod provisioner_policy_tests {
use super::*;
fn template(cost: u64, permissions: u64) -> ProvisionerPolicyTemplate {
ProvisionerPolicyTemplate {
cost_table: Arc::new(CostTable::builder(CostUnits(cost), CostUnits(cost)).build()),
limits: ResolvedLimits::new(1),
permissions: PermissionBits(permissions),
policy_revision: PolicyRevision::UNSTATED,
}
}
fn snapshot(template: &ProvisionerPolicyTemplate) -> AccountSnapshot {
AccountSnapshot::builder(
AccountId(1),
Generation(1),
AccountStatus::Active,
Timestamp::UNIX_EPOCH,
template.permissions,
template.limits,
Arc::clone(&template.cost_table),
)
.build()
}
#[test]
fn provisioner_templates_match_whole_policies_and_are_identity_scoped() {
let first = template(1, 1);
let second = template(2, 2);
let limits =
ProvisionerLimits::new(CostUnits(100), vec![first.clone(), second.clone()]).unwrap();
assert_eq!(
limits,
ProvisionerLimits::new(
CostUnits(100),
vec![second.clone(), first.clone(), first.clone()]
)
.unwrap()
);
assert_ne!(
limits,
ProvisionerLimits::new(CostUnits(101), vec![first.clone(), second.clone()]).unwrap()
);
let subset = ProvisionerLimits::new(CostUnits(100), vec![first.clone()]).unwrap();
assert_ne!(limits, subset);
assert_ne!(subset, limits);
let mut candidate = snapshot(&first);
assert!(limits.allows_snapshot(&candidate));
candidate.generation = Generation(u64::MAX);
candidate.valid_until = Timestamp::MAX;
candidate.account_id = AccountId(2);
candidate.key_id = Some(tollgate_core::KeyId(3));
assert!(
limits.allows_snapshot(&candidate),
"binding and publication metadata are separate store contracts"
);
candidate.cost_table = Arc::clone(&second.cost_table);
assert!(
!limits.allows_snapshot(&candidate),
"a mixed policy is not an approved policy"
);
candidate.permissions = second.permissions;
assert!(
limits.allows_snapshot(&candidate),
"the second complete policy is approved too"
);
let other_identity = ProvisionerLimits::new(CostUnits(100), vec![first]).unwrap();
assert!(!other_identity.allows_snapshot(&candidate));
candidate.enforcement_mode = EnforcementMode::Elastic {
overage_cap: CostUnits(10),
};
assert!(
!limits.allows_snapshot(&candidate),
"templates cannot approve unfunded credit"
);
}
#[test]
fn provisioner_limits_refuse_missing_or_unpublishable_templates() {
assert!(ProvisionerLimits::new(CostUnits(100), vec![]).is_err());
let mut invalid = template(1, 1);
invalid.limits = ResolvedLimits::new(1).with_weighted_rate(0, 1);
assert!(ProvisionerLimits::new(CostUnits(100), vec![invalid]).is_err());
}
}