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