1use std::collections::HashMap;
4use std::sync::Arc;
5
6use arc_swap::ArcSwap;
7use axum::extract::{FromRequestParts, Request, State};
8use axum::http::{header, request::Parts};
9use axum::middleware::Next;
10use axum::response::Response;
11use jiff::Timestamp;
12use serde::{Deserialize, Serialize};
13use tollgate_auth::CredentialVerifier;
14use tollgate_core::Principal;
15use tollgate_store::Clock;
16
17use crate::error::ApiError;
18use crate::transport::{PeerIdentity, TlsConfig};
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
22#[serde(rename_all = "snake_case")]
23pub enum Role {
24 Instance,
27 Operator,
30}
31
32#[derive(Debug, Clone, PartialEq, Eq)]
34pub struct ControlIdentity {
35 name: String,
36 role: Role,
37}
38
39impl ControlIdentity {
40 pub fn new(name: impl Into<String>, role: Role) -> Result<Self, SecurityError> {
47 let name = name.into();
48 if name.is_empty()
49 || name.len() > 128
50 || !name
51 .bytes()
52 .all(|b| b.is_ascii_alphanumeric() || b"@._:/-".contains(&b))
53 {
54 return Err(SecurityError(
55 "identity must be 1..=128 ASCII identifier characters",
56 ));
57 }
58 Ok(Self { name, role })
59 }
60
61 pub fn name(&self) -> &str {
64 &self.name
65 }
66 pub fn role(&self) -> Role {
68 self.role
69 }
70}
71
72#[derive(Debug, Clone, Copy, PartialEq, Eq)]
78pub struct SecurityError(pub &'static str);
79
80impl std::fmt::Display for SecurityError {
81 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
82 f.write_str(self.0)
83 }
84}
85impl std::error::Error for SecurityError {}
86
87struct BearerScheme {
88 verifier: Arc<dyn CredentialVerifier + Send + Sync>,
89 identities: HashMap<Principal, ControlIdentity>,
90}
91
92#[derive(Default)]
95pub struct SecurityPolicy {
96 bearers: Vec<BearerScheme>,
97 certificates: HashMap<[u8; 32], ControlIdentity>,
98}
99
100impl SecurityPolicy {
101 pub fn new() -> Self {
103 Self::default()
104 }
105
106 pub fn with_bearer(
119 mut self,
120 verifier: Arc<dyn CredentialVerifier + Send + Sync>,
121 identities: impl IntoIterator<Item = (Principal, ControlIdentity)>,
122 ) -> Result<Self, SecurityError> {
123 let mut mapped = HashMap::new();
124 for (principal, identity) in identities {
125 if mapped.insert(principal, identity).is_some() {
126 return Err(SecurityError("duplicate bearer principal"));
127 }
128 }
129 self.bearers.push(BearerScheme {
130 verifier,
131 identities: mapped,
132 });
133 Ok(self)
134 }
135
136 pub fn with_certificate(
138 mut self,
139 fingerprint: [u8; 32],
140 identity: ControlIdentity,
141 ) -> Result<Self, SecurityError> {
142 if self.certificates.insert(fingerprint, identity).is_some() {
143 return Err(SecurityError("duplicate client certificate"));
144 }
145 Ok(self)
146 }
147
148 fn bearer(&self, credential: &[u8], now: Timestamp) -> Result<ControlIdentity, ApiError> {
149 let mut identity = None;
150 for scheme in &self.bearers {
151 if let Some(proof) = scheme.verifier.verify(credential) {
152 if !proof.is_reusable_at(now) {
153 return Err(ApiError::unauthorized());
154 }
155 let found = scheme
156 .identities
157 .get(&proof.principal)
158 .ok_or_else(ApiError::forbidden)?;
159 if identity.as_ref().is_some_and(|previous| previous != found) {
160 return Err(ApiError::unauthorized());
161 }
162 identity = Some(found.clone());
163 }
164 }
165 identity.ok_or_else(ApiError::unauthorized)
166 }
167}
168
169pub(crate) struct SecurityBundle {
170 pub policy: SecurityPolicy,
171 pub tls: Option<TlsConfig>,
172}
173
174pub struct ServerSecurity {
176 pub(crate) current: ArcSwap<SecurityBundle>,
177 encrypted: bool,
178}
179
180impl ServerSecurity {
181 pub fn new(policy: SecurityPolicy, tls: Option<TlsConfig>) -> Result<Arc<Self>, SecurityError> {
190 validate(&policy, tls.as_ref())?;
191 Ok(Arc::new(Self {
192 encrypted: tls.is_some(),
193 current: ArcSwap::from_pointee(SecurityBundle { policy, tls }),
194 }))
195 }
196
197 pub fn replace(
201 &self,
202 policy: SecurityPolicy,
203 tls: Option<TlsConfig>,
204 ) -> Result<(), SecurityError> {
205 validate(&policy, tls.as_ref())?;
206 if self.encrypted != tls.is_some() {
207 return Err(SecurityError(
208 "changing TLS mode requires a listener restart",
209 ));
210 }
211 self.current.store(Arc::new(SecurityBundle { policy, tls }));
212 Ok(())
213 }
214
215 pub fn encrypted(&self) -> bool {
219 self.encrypted
220 }
221}
222
223fn validate(policy: &SecurityPolicy, tls: Option<&TlsConfig>) -> Result<(), SecurityError> {
224 if !policy.certificates.is_empty() && tls.is_none_or(|tls| !tls.verifies_clients()) {
225 return Err(SecurityError(
226 "certificate identities require TLS with a client CA",
227 ));
228 }
229 Ok(())
230}
231
232#[derive(Clone)]
233pub(crate) struct Authorization {
234 pub security: Arc<ServerSecurity>,
235 pub clock: Arc<dyn Clock>,
236 pub role: Role,
237}
238
239fn bearer(request: &Request) -> Result<Option<&[u8]>, ApiError> {
241 let mut values = request.headers().get_all(header::AUTHORIZATION).iter();
242 let Some(header) = values.next() else {
243 return Ok(None);
244 };
245 if values.next().is_some() {
246 return Err(ApiError::unauthorized());
247 }
248 let header = header.as_bytes();
249 if header.len() > 16 * 1024 {
252 return Err(ApiError::unauthorized());
253 }
254 let Some(separator) = header.iter().position(|b| *b == b' ') else {
255 return Err(ApiError::unauthorized());
256 };
257 let (scheme, token) = (&header[..separator], &header[separator + 1..]);
258 if !scheme.eq_ignore_ascii_case(b"Bearer")
259 || token.is_empty()
260 || !token.iter().all(|b| b.is_ascii_graphic())
261 {
262 return Err(ApiError::unauthorized());
263 }
264 Ok(Some(token))
265}
266
267pub(crate) async fn authorize(
268 State(auth): State<Authorization>,
269 mut request: Request,
270 next: Next,
271) -> Result<Response, ApiError> {
272 let bundle = auth.security.current.load_full();
273 let now = auth.clock.now();
274 let bearer = bearer(&request)?;
275 let peer = request
276 .extensions()
277 .get::<axum::extract::ConnectInfo<PeerIdentity>>();
278 let certificate = match peer.and_then(|peer| peer.0.certificates.as_deref()) {
279 Some(chain) => {
280 let tls = bundle.tls.as_ref().ok_or_else(ApiError::unauthorized)?;
281 tls.verify(chain, now)
284 .map_err(|_| ApiError::unauthorized())?;
285 let fingerprint = crate::transport::fingerprint(&chain[0]);
286 Some(
287 bundle
288 .policy
289 .certificates
290 .get(&fingerprint)
291 .ok_or_else(ApiError::forbidden)?
292 .clone(),
293 )
294 }
295 None => None,
296 };
297 let token = bearer
298 .map(|token| bundle.policy.bearer(token, now))
299 .transpose()?;
300 let identity = match (token, certificate) {
301 (Some(a), Some(b)) if a != b => return Err(ApiError::unauthorized()),
302 (Some(identity), _) | (_, Some(identity)) => identity,
303 (None, None) => return Err(ApiError::unauthorized()),
304 };
305 if identity.role != auth.role {
306 return Err(ApiError::forbidden());
307 }
308 tracing::debug!(actor = identity.name(), role = ?identity.role(), "control-plane request authenticated");
309 match identity.role {
310 Role::Instance => {
311 request.extensions_mut().insert(InstanceIdentity);
312 }
313 Role::Operator => {
314 request.extensions_mut().insert(OperatorIdentity(identity));
315 }
316 }
317 Ok(next.run(request).await)
318}
319
320#[derive(Clone)]
323pub(crate) struct InstanceIdentity;
324#[derive(Clone)]
325pub(crate) struct OperatorIdentity(pub ControlIdentity);
326
327macro_rules! identity_extractor {
328 ($name:ident) => {
329 impl<S: Send + Sync> FromRequestParts<S> for $name {
330 type Rejection = ApiError;
331 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, ApiError> {
332 parts
333 .extensions
334 .get::<Self>()
335 .cloned()
336 .ok_or_else(ApiError::unauthorized)
337 }
338 }
339 };
340}
341identity_extractor!(InstanceIdentity);
342identity_extractor!(OperatorIdentity);
343
344#[cfg(test)]
345mod framing_tests {
346 use super::*;
347
348 #[test]
349 fn bearer_framing_checks_each_condition_before_scheme_verification() {
350 let request = |value: &[u8]| {
351 Request::builder()
352 .header(
353 header::AUTHORIZATION,
354 axum::http::HeaderValue::from_bytes(value).unwrap(),
355 )
356 .body(axum::body::Body::empty())
357 .unwrap()
358 };
359 for value in [
360 b"Basic token".as_slice(),
361 b"Bearer ",
362 b"Bearer one\ttwo",
363 b"Bearer token",
364 b"Bearer",
365 b"Bearer \x80",
366 ] {
367 assert!(bearer(&request(value)).is_err());
368 }
369 for size in [1, 16 * 1024 - 7] {
370 let token = vec![b'x'; size];
371 let mut header = b"bEaReR ".to_vec();
372 header.extend_from_slice(&token);
373 let request = request(&header);
374 assert_eq!(bearer(&request).unwrap(), Some(token.as_slice()));
375 }
376 for size in [16 * 1024 + 1, 32 * 1024] {
377 let mut value = b"Bearer ".to_vec();
378 value.resize(size, b'x');
379 assert!(bearer(&request(&value)).is_err());
380 }
381 let mut duplicate = request(b"Bearer token");
382 duplicate
383 .headers_mut()
384 .append(header::AUTHORIZATION, "Bearer token".parse().unwrap());
385 assert!(bearer(&duplicate).is_err());
386 assert_eq!(
387 bearer(&Request::new(axum::body::Body::empty())).unwrap(),
388 None
389 );
390 }
391
392 #[test]
393 fn overlapping_bearer_schemes_must_agree_on_the_verified_identity() {
394 use tollgate_auth::{HmacRegistry, Verified};
395 let verifier = Arc::new(HmacRegistry::new(b"fixture-scheme-agreement-secret"));
396 let principal = verifier.install_credentials([b"shared-credential".as_slice()])[0];
397 let now = Timestamp::from_second(100).unwrap();
398 let instance = ControlIdentity::new("instance", Role::Instance).unwrap();
399 let operator = ControlIdentity::new("operator", Role::Operator).unwrap();
400 for (second, agrees) in [(instance.clone(), true), (operator, false)] {
401 let policy = SecurityPolicy::new()
402 .with_bearer(verifier.clone(), [(principal, instance.clone())])
403 .unwrap()
404 .with_bearer(verifier.clone(), [(principal, second)])
405 .unwrap();
406 assert_eq!(policy.bearer(b"shared-credential", now).is_ok(), agrees);
407 if agrees {
408 assert_eq!(policy.bearer(b"shared-credential", now).unwrap(), instance);
409 }
410 }
411 struct Expired(Principal);
412 impl CredentialVerifier for Expired {
413 fn verify(&self, _: &[u8]) -> Option<Verified> {
414 Some(Verified::until(
415 self.0,
416 Timestamp::from_second(100).unwrap(),
417 ))
418 }
419 }
420 let policy = SecurityPolicy::new()
421 .with_bearer(Arc::new(Expired(principal)), [(principal, instance)])
422 .unwrap();
423 assert!(policy.bearer(b"shared-credential", now).is_err());
424 }
425}
426
427impl OperatorIdentity {
428 pub(crate) async fn run<T, E: Into<ApiError>>(
431 &self,
432 action: &'static str,
433 target: impl std::fmt::Display,
434 clock: &dyn Clock,
435 operation: impl Future<Output = Result<tollgate_store::AdminReceipt<T>, E>>,
436 ) -> Result<T, ApiError> {
437 struct Attempt<'a> {
438 id: tollgate_core::RequestId,
439 identity: &'a ControlIdentity,
440 action: &'static str,
441 target: String,
442 clock: &'a dyn Clock,
443 finished: bool,
444 }
445 impl Drop for Attempt<'_> {
446 fn drop(&mut self) {
447 if !self.finished {
448 tracing::warn!(target: "tollgate::audit", actor = self.identity.name(), action = self.action,
449 operation_id = %self.id,
450 resource = self.target, at = %self.clock.now(), outcome = "cancelled_unknown",
451 "administrative operation abandoned; commit outcome may be unknown");
452 }
453 }
454 }
455 let mut identifier = [0u8; 16];
456 getrandom::fill(&mut identifier).map_err(|_| {
457 ApiError::from(tollgate_store::StoreError(
458 "audit identity entropy unavailable".into(),
459 ))
460 })?;
461 let mut attempt = Attempt {
462 id: tollgate_core::RequestId(u128::from_be_bytes(identifier)),
463 identity: &self.0,
464 action,
465 target: target.to_string(),
466 clock,
467 finished: false,
468 };
469 tracing::info!(target: "tollgate::audit", actor = self.0.name(), action,
470 operation_id = %attempt.id,
471 resource = attempt.target, at = %clock.now(), outcome = "started", "administrative operation started");
472 let result = operation.await;
473 attempt.finished = true;
474 match result {
475 Ok(receipt) => {
476 tracing::info!(target: "tollgate::audit", actor = self.0.name(), action,
477 operation_id = %attempt.id,
478 resource = attempt.target, at = %clock.now(), outcome = "confirmed",
479 before = ?receipt.before, after = ?receipt.after, "administrative operation completed");
480 Ok(receipt.outcome)
481 }
482 Err(error) => {
483 let error = error.into();
484 tracing::warn!(target: "tollgate::audit", actor = self.0.name(), action,
485 operation_id = %attempt.id,
486 resource = attempt.target, at = %clock.now(), outcome = "failed",
487 code = error.code, status = error.status.as_u16(), "administrative operation failed; storage errors may conceal a commit");
488 Err(error)
489 }
490 }
491 }
492}
493
494#[cfg(test)]
495mod tests {
496 use super::*;
497
498 #[tokio::test]
499 async fn cancelled_admin_operations_report_an_unknown_commit_without_a_receipt() {
500 use tracing::instrument::WithSubscriber;
501 #[derive(Clone, Default)]
502 struct Capture(Arc<std::sync::Mutex<Vec<u8>>>);
503 impl std::io::Write for Capture {
504 fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
505 self.0.lock().unwrap().extend_from_slice(bytes);
506 Ok(bytes.len())
507 }
508 fn flush(&mut self) -> std::io::Result<()> {
509 Ok(())
510 }
511 }
512 let capture = Capture::default();
513 let output = capture.clone();
514 let subscriber = tracing_subscriber::fmt()
515 .without_time()
516 .with_ansi(false)
517 .with_writer(move || output.clone())
518 .finish();
519 let identity =
520 OperatorIdentity(ControlIdentity::new("fixture-operator", Role::Operator).unwrap());
521 async {
522 let operation = identity.run(
523 "deposit",
524 "fixture-account",
525 &tollgate_store::SystemClock,
526 std::future::pending::<
527 Result<tollgate_store::AdminReceipt<()>, tollgate_store::StoreError>,
528 >(),
529 );
530 tokio::pin!(operation);
531 tokio::select! {
532 biased;
533 _ = &mut operation => panic!("the backend remains pending"),
534 _ = tokio::task::yield_now() => {},
535 }
536 }
538 .with_subscriber(subscriber)
539 .await;
540 let bytes = capture.0.lock().unwrap();
541 let log = std::str::from_utf8(&bytes).unwrap();
542 assert!(
543 log.contains("started") && log.contains("cancelled_unknown"),
544 "{log}"
545 );
546 assert!(
547 log.contains("fixture-operator")
548 && log.contains("fixture-account")
549 && log.contains("operation_id=")
550 );
551 assert!(!log.contains("confirmed") && !log.contains("before=") && !log.contains("after="));
552 }
553}