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;