affine_core 0.0.2

AFFiNE primitive core.
Documentation
use std::{
  collections::VecDeque,
  sync::{Mutex, OnceLock},
};

use base64::{Engine, engine::general_purpose::STANDARD as BASE64};
use chrono::{DateTime, Duration, Utc};
use p256::{
  ecdsa::{
    Signature, SigningKey, VerifyingKey,
    signature::{Signer, Verifier},
  },
  pkcs8::{DecodePrivateKey, DecodePublicKey},
};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use thiserror::Error;

use super::MAX_SEAT_QUANTITY;

const LICENSE_FORMAT_VERSION: u32 = 1;
const LICENSE_AUDIENCE: &str = "affine-selfhost";
const LICENSE_PLAN: &str = "selfhost_team";
const LICENSE_ENVELOPE_LIFETIME_HOURS: i64 = 24;
const LICENSE_NOT_BEFORE_SKEW_MINUTES: i64 = 5;
const CACHE_CAPACITY: usize = 256;
const CACHE_TTL_MINUTES: i64 = 10;

static CACHE: OnceLock<Mutex<VecDeque<CacheEntry>>> = OnceLock::new();

#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(deny_unknown_fields, rename_all = "camelCase")]
struct LicenseClaimsV1 {
  format_version: u32,
  license_id: String,
  workspace_id: String,
  audience: String,
  plan: String,
  seat_quantity: i32,
  issued_at: String,
  not_before: String,
  expires_at: String,
}

#[derive(Debug, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
struct LicenseEnvelopeV1 {
  claims: LicenseClaimsV1,
  signature: String,
}

#[derive(Clone, Debug, Eq, PartialEq)]
pub struct VerifiedLicenseClaims {
  license_id: String,
  workspace_id: String,
  seat_quantity: i32,
  issued_at: DateTime<Utc>,
  not_before: DateTime<Utc>,
  expires_at: DateTime<Utc>,
}

impl VerifiedLicenseClaims {
  pub fn license_id(&self) -> &str {
    &self.license_id
  }

  pub fn workspace_id(&self) -> &str {
    &self.workspace_id
  }

  pub const fn audience(&self) -> &'static str {
    LICENSE_AUDIENCE
  }

  pub const fn seat_quantity(&self) -> i32 {
    self.seat_quantity
  }

  pub const fn issued_at(&self) -> DateTime<Utc> {
    self.issued_at
  }

  pub const fn expires_at(&self) -> DateTime<Utc> {
    self.expires_at
  }
}

pub struct LicenseVerifier;

pub struct LicenseIssuer;

#[derive(Clone, Copy, Debug)]
pub struct LicenseIssuance<'a> {
  pub license_id: &'a str,
  pub workspace_id: &'a str,
  pub seat_quantity: i32,
  pub subscription_end: Option<DateTime<Utc>>,
  pub now: DateTime<Utc>,
}

#[derive(Clone, Copy, Debug, Eq, Error, PartialEq)]
pub enum LicenseIssuanceError {
  #[error("invalid private key")]
  InvalidPrivateKey,
  #[error("invalid license claims")]
  InvalidClaims,
  #[error("invalid license time window")]
  InvalidTimeWindow,
  #[error("failed to encode license envelope")]
  Encoding,
}

#[derive(Clone)]
struct CacheEntry {
  payload_digest: [u8; 32],
  public_key_digest: [u8; 32],
  claims: VerifiedLicenseClaims,
  cache_until: DateTime<Utc>,
}

#[derive(Clone, Copy, Debug, Eq, Error, PartialEq)]
pub enum LicenseError {
  #[error("invalid license envelope")]
  InvalidEnvelope,
  #[error("invalid license signature")]
  InvalidSignature,
  #[error("invalid public key")]
  InvalidPublicKey,
  #[error("unknown license format")]
  UnknownFormat,
  #[error("license audience mismatched")]
  AudienceMismatch,
  #[error("unknown self-hosted license plan")]
  UnknownPlan,
  #[error("invalid license claims")]
  InvalidClaims,
  #[error("workspace mismatched with license")]
  WorkspaceMismatch,
  #[error("invalid license timestamp")]
  InvalidTime,
  #[error("invalid license time window")]
  InvalidTimeWindow,
  #[error("license is not valid yet")]
  NotBefore,
  #[error("license expired")]
  Expired,
}

impl LicenseError {
  pub const fn code(self) -> &'static str {
    match self {
      Self::InvalidEnvelope => "invalid_envelope",
      Self::InvalidSignature => "invalid_signature",
      Self::InvalidPublicKey => "invalid_public_key",
      Self::UnknownFormat => "unknown_format",
      Self::AudienceMismatch => "audience_mismatch",
      Self::UnknownPlan => "unknown_plan",
      Self::InvalidClaims => "invalid_claims",
      Self::WorkspaceMismatch => "workspace_mismatch",
      Self::InvalidTime => "invalid_time",
      Self::InvalidTimeWindow => "invalid_time_window",
      Self::NotBefore => "not_before",
      Self::Expired => "expired",
    }
  }
}

impl LicenseVerifier {
  pub fn verify(
    payload: &[u8],
    public_key: &str,
    expected_workspace_id: Option<&str>,
    now: DateTime<Utc>,
  ) -> Result<VerifiedLicenseClaims, LicenseError> {
    let payload_digest = Sha256::digest(payload).into();
    let public_key_digest = Sha256::digest(public_key.as_bytes()).into();
    if let Some(claims) = cached_claims(payload_digest, public_key_digest, expected_workspace_id, now) {
      return claims;
    }

    let envelope: LicenseEnvelopeV1 = serde_json::from_slice(payload).map_err(|_| LicenseError::InvalidEnvelope)?;
    let canonical = serde_json::to_vec(&envelope.claims).map_err(|_| LicenseError::InvalidEnvelope)?;
    let signature = BASE64
      .decode(&envelope.signature)
      .map_err(|_| LicenseError::InvalidSignature)?;
    let signature = Signature::from_der(&signature).map_err(|_| LicenseError::InvalidSignature)?;
    let verifying_key = VerifyingKey::from_public_key_pem(public_key).map_err(|_| LicenseError::InvalidPublicKey)?;
    verifying_key
      .verify(&canonical, &signature)
      .map_err(|_| LicenseError::InvalidSignature)?;
    let claims = validate_claims(envelope.claims, expected_workspace_id, now)?;
    cache_claims(payload_digest, public_key_digest, claims.clone(), now);
    Ok(claims)
  }
}

impl LicenseIssuer {
  pub fn validate_seat_quantity(seat_quantity: i32) -> Result<(), LicenseIssuanceError> {
    if !(1..=MAX_SEAT_QUANTITY).contains(&seat_quantity) {
      return Err(LicenseIssuanceError::InvalidClaims);
    }
    Ok(())
  }

  pub fn issue(input: LicenseIssuance<'_>, private_key: &str) -> Result<Vec<u8>, LicenseIssuanceError> {
    if input.license_id.is_empty() || input.workspace_id.is_empty() {
      return Err(LicenseIssuanceError::InvalidClaims);
    }
    Self::validate_seat_quantity(input.seat_quantity)?;
    let expires_at = input
      .subscription_end
      .unwrap_or(DateTime::<Utc>::MAX_UTC)
      .min(input.now + Duration::hours(LICENSE_ENVELOPE_LIFETIME_HOURS));
    if expires_at <= input.now {
      return Err(LicenseIssuanceError::InvalidTimeWindow);
    }
    let claims = LicenseClaimsV1 {
      format_version: LICENSE_FORMAT_VERSION,
      license_id: input.license_id.to_string(),
      workspace_id: input.workspace_id.to_string(),
      audience: LICENSE_AUDIENCE.to_string(),
      plan: LICENSE_PLAN.to_string(),
      seat_quantity: input.seat_quantity,
      issued_at: input.now.to_rfc3339(),
      not_before: (input.now - Duration::minutes(LICENSE_NOT_BEFORE_SKEW_MINUTES)).to_rfc3339(),
      expires_at: expires_at.to_rfc3339(),
    };
    let canonical = serde_json::to_vec(&claims).map_err(|_| LicenseIssuanceError::Encoding)?;
    let signing_key = SigningKey::from_pkcs8_pem(private_key).map_err(|_| LicenseIssuanceError::InvalidPrivateKey)?;
    let signature: Signature = signing_key.sign(&canonical);
    serde_json::to_vec(&LicenseEnvelopeV1 {
      claims,
      signature: BASE64.encode(signature.to_der()),
    })
    .map_err(|_| LicenseIssuanceError::Encoding)
  }
}

fn cached_claims(
  payload_digest: [u8; 32],
  public_key_digest: [u8; 32],
  expected_workspace_id: Option<&str>,
  now: DateTime<Utc>,
) -> Option<Result<VerifiedLicenseClaims, LicenseError>> {
  let mut cache = CACHE.get_or_init(Default::default).lock().ok()?;
  cache.retain(|entry| now < entry.cache_until && now < entry.claims.expires_at);
  let index = cache
    .iter()
    .position(|entry| entry.payload_digest == payload_digest && entry.public_key_digest == public_key_digest)?;
  let entry = cache.remove(index)?;
  let claims = entry.claims.clone();
  cache.push_front(entry);
  Some(recheck_claims(claims, expected_workspace_id, now))
}

fn cache_claims(
  payload_digest: [u8; 32],
  public_key_digest: [u8; 32],
  claims: VerifiedLicenseClaims,
  now: DateTime<Utc>,
) {
  let cache_until = claims.expires_at.min(now + Duration::minutes(CACHE_TTL_MINUTES));
  if let Ok(mut cache) = CACHE.get_or_init(Default::default).lock() {
    cache.retain(|entry| {
      (entry.payload_digest != payload_digest || entry.public_key_digest != public_key_digest)
        && now < entry.cache_until
        && now < entry.claims.expires_at
    });
    cache.push_front(CacheEntry {
      payload_digest,
      public_key_digest,
      claims,
      cache_until,
    });
    cache.truncate(CACHE_CAPACITY);
  }
}

fn recheck_claims(
  claims: VerifiedLicenseClaims,
  expected_workspace_id: Option<&str>,
  now: DateTime<Utc>,
) -> Result<VerifiedLicenseClaims, LicenseError> {
  if expected_workspace_id.is_some_and(|workspace_id| workspace_id != claims.workspace_id) {
    return Err(LicenseError::WorkspaceMismatch);
  }
  if now < claims.not_before {
    return Err(LicenseError::NotBefore);
  }
  if now >= claims.expires_at {
    return Err(LicenseError::Expired);
  }
  Ok(claims)
}

fn validate_claims(
  claims: LicenseClaimsV1,
  expected_workspace_id: Option<&str>,
  now: DateTime<Utc>,
) -> Result<VerifiedLicenseClaims, LicenseError> {
  if claims.format_version != LICENSE_FORMAT_VERSION {
    return Err(LicenseError::UnknownFormat);
  }
  if claims.audience != LICENSE_AUDIENCE {
    return Err(LicenseError::AudienceMismatch);
  }
  if claims.plan != LICENSE_PLAN {
    return Err(LicenseError::UnknownPlan);
  }
  if claims.license_id.is_empty()
    || claims.workspace_id.is_empty()
    || claims.seat_quantity <= 0
    || claims.seat_quantity > MAX_SEAT_QUANTITY
  {
    return Err(LicenseError::InvalidClaims);
  }
  let issued_at = parse_time(&claims.issued_at)?;
  let not_before = parse_time(&claims.not_before)?;
  let expires_at = parse_time(&claims.expires_at)?;
  if not_before > issued_at || issued_at >= expires_at {
    return Err(LicenseError::InvalidTimeWindow);
  }
  recheck_claims(
    VerifiedLicenseClaims {
      license_id: claims.license_id,
      workspace_id: claims.workspace_id,
      seat_quantity: claims.seat_quantity,
      issued_at,
      not_before,
      expires_at,
    },
    expected_workspace_id,
    now,
  )
}

fn parse_time(value: &str) -> Result<DateTime<Utc>, LicenseError> {
  DateTime::parse_from_rfc3339(value)
    .map(|value| value.with_timezone(&Utc))
    .map_err(|_| LicenseError::InvalidTime)
}

#[cfg(test)]
#[path = "../tests/access_control/license/tests.rs"]
mod tests;