use serde::{Deserialize, Serialize};
use crate::principal::{PrincipalId, PrincipalKind};
use crate::scope::Scope;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum CeilingError {
#[error("ceiling principal must be a service principal; got {0}")]
PrincipalNotService(PrincipalKind),
#[error("expires_at must be strictly greater than issued_at")]
ExpiresBeforeIssued,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[non_exhaustive]
pub struct ServiceCeiling {
pub principal: PrincipalId,
pub audiences: Vec<String>,
pub scopes: Vec<Scope>,
pub issued_at: i64,
pub expires_at: i64,
}
impl ServiceCeiling {
pub fn new(
principal: PrincipalId,
audiences: Vec<String>,
scopes: Vec<Scope>,
issued_at: i64,
expires_at: i64,
) -> Result<Self, CeilingError> {
if principal.kind != PrincipalKind::Service {
return Err(CeilingError::PrincipalNotService(principal.kind));
}
if expires_at <= issued_at {
return Err(CeilingError::ExpiresBeforeIssued);
}
Ok(Self {
principal,
audiences,
scopes,
issued_at,
expires_at,
})
}
pub fn is_expired_at(&self, now: i64) -> bool {
self.expires_at <= now
}
pub fn allows_audience(&self, aud: &str) -> bool {
self.audiences.iter().any(|a| a == aud)
}
pub fn covers<'a>(&self, requested: impl IntoIterator<Item = &'a str>) -> bool {
requested
.into_iter()
.all(|s| self.scopes.iter().any(|c| c.as_wire() == s))
}
}
#[derive(Deserialize)]
struct RawServiceCeiling {
principal: PrincipalId,
audiences: Vec<String>,
scopes: Vec<Scope>,
issued_at: i64,
expires_at: i64,
}
impl<'de> Deserialize<'de> for ServiceCeiling {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
let r = RawServiceCeiling::deserialize(d)?;
Self::new(r.principal, r.audiences, r.scopes, r.issued_at, r.expires_at)
.map_err(serde::de::Error::custom)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::yah_scopes;
fn ceiling() -> ServiceCeiling {
ServiceCeiling::new(
PrincipalId::service("issues"),
vec!["https://inference.example".into()],
vec![yah_scopes::CLOUD_READ],
100,
200,
)
.unwrap()
}
#[test]
fn rejects_non_service_principal() {
let e = ServiceCeiling::new(PrincipalId::user("u"), vec![], vec![], 1, 2).unwrap_err();
assert_eq!(e, CeilingError::PrincipalNotService(PrincipalKind::User));
}
#[test]
fn roundtrips_and_revalidates_on_the_wire() {
let c = ceiling();
let json = serde_json::to_string(&c).unwrap();
assert_eq!(serde_json::from_str::<ServiceCeiling>(&json).unwrap(), c);
let bad = json.replace("\"expires_at\":200", "\"expires_at\":50");
assert!(serde_json::from_str::<ServiceCeiling>(&bad).is_err());
}
#[test]
fn covers_and_audience_checks() {
let c = ceiling();
assert!(c.covers([yah_scopes::CLOUD_READ.as_wire()]));
assert!(!c.covers([yah_scopes::CLOUD_DEPLOY.as_wire()]));
assert!(c.allows_audience("https://inference.example"));
assert!(!c.allows_audience("https://other.example"));
assert!(c.is_expired_at(200) && !c.is_expired_at(199));
}
}