#![deny(missing_docs)]
pub(crate) mod instant;
#[cfg(feature = "test-support")]
pub mod test_support;
use instant::StoredInstant;
use std::num::{NonZeroI64, NonZeroUsize};
use std::sync::Arc;
use async_trait::async_trait;
use jiff::{SignedDuration, Timestamp};
use serde::{Deserialize, Serialize};
use sqlx::postgres::PgPoolOptions;
use sqlx::{PgPool, Postgres, Row, Transaction};
use tokio::sync::broadcast;
use tollgate_core::{
AccountId, AccountSnapshot, AccountStatus, BudgetSchedule, BudgetView, CapacityClass,
CostTable, CostUnits, EnforcementMode, FencingToken, Generation, KeyId, LeaseGrant, LeaseId,
Period, PermissionBits, PolicyRevision, Principal, PublishableSnapshot, ResolvedLimits,
Rollover, UsageEvent,
};
use tollgate_store::{
AccountConfig, AccountView, AdminAuthority, AdminReceipt, AdminState, AdminStore,
AllocateError, Allocation, BudgetError, Conservation, CreateAccountError, GrantPolicy,
IngestError, IngestReport, KeyDirectory, KeyError, KeyRecord, KeySnapshotError, KeySummary,
LeaseAllocator, PUSH_CHANNEL_CAPACITY, PublishSnapshotError, ReclaimBatch, ReclaimedLease,
Revocation, RolledAccount, RolloverBatch, SetStatusError, SnapshotPush, SnapshotResolution,
SnapshotSource, StatusChange, StoreError, StoreHealth, UsageSink, pushes_exceed_capacity,
validate_key_page_limit,
};
const STATE_ACTIVE: i16 = 0;
const STATE_RELEASED: i16 = 1;
const STATE_EXPIRED: i16 = 2;
#[derive(Debug, Clone, Copy)]
struct StoredId(u128);
impl Serialize for StoredId {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match u64::try_from(self.0) {
Ok(value) => serializer.serialize_u64(value),
Err(_) => serializer.collect_str(&format_args!("{:032x}", self.0)),
}
}
}
impl<'de> Deserialize<'de> for StoredId {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
struct StoredIdVisitor;
impl serde::de::Visitor<'_> for StoredIdVisitor {
type Value = StoredId;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a legacy u64 number or a canonical 128-bit identifier string")
}
fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E> {
Ok(StoredId(u128::from(value)))
}
fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
value
.parse::<AccountId>()
.map(|id| StoredId(id.0))
.map_err(E::custom)
}
}
deserializer.deserialize_any(StoredIdVisitor)
}
}
fn digest_from(bytes: &[u8], key_id: KeyId) -> Result<[u8; 32], StoreError> {
<[u8; 32]>::try_from(bytes).map_err(|_| {
StoreError(format!(
"credential {key_id} has a {}-byte digest in storage, expected 32",
bytes.len()
))
})
}
#[async_trait]
impl KeyDirectory for PostgresStore {
async fn credential_activity(
&self,
keys: &[KeyId],
) -> Result<Vec<tollgate_store::CredentialActivity>, StoreError> {
use tollgate_store::{CredentialActivity, CredentialActivityState};
let mut result = Vec::with_capacity(keys.len());
for chunk in keys.chunks(tollgate_store::MAX_INGEST_BATCH) {
let ids: Vec<_> = chunk.iter().map(|id| id_bytes(id.0)).collect();
let rows = sqlx::query(
"SELECT evidence.key_id IS NOT NULL, evidence.last_committed_at_us
FROM UNNEST($1::bytea[]) WITH ORDINALITY AS requested(key_id, ordinal)
LEFT JOIN LATERAL (
SELECT k.key_id, a.last_committed_at_us
FROM tollgate_credential_keys k
LEFT JOIN tollgate_credential_activity a ON a.key_id = k.key_id
WHERE k.key_id = requested.key_id LIMIT 1
) AS evidence ON true
ORDER BY requested.ordinal",
)
.bind(&ids)
.fetch_all(&self.pool)
.await
.map_err(storage)?;
if rows.len() != chunk.len() {
return Err(StoreError("incomplete credential activity read".into()));
}
for (&key_id, row) in chunk.iter().zip(rows) {
let state = if !row.get::<bool, _>(0) {
CredentialActivityState::Unknown
} else if let Some(at) = row.get::<Option<i64>, _>(1) {
CredentialActivityState::Committed {
last_committed_at: micros_ts(at, "credential activity")?,
}
} else {
CredentialActivityState::Unobserved
};
result.push(CredentialActivity { key_id, state });
}
}
Ok(result)
}
async fn insert_key(&self, record: KeyRecord) -> Result<(), KeyError> {
self.insert_credential(record, None)
.await
.map(|receipt| receipt.outcome)
}
async fn revoke_key(&self, key_id: KeyId, now: Timestamp) -> Result<Revocation, KeyError> {
self.revoke_key_audited(key_id, now)
.await
.map(|receipt| receipt.outcome)
}
async fn revoke_key_audited(
&self,
key_id: KeyId,
now: Timestamp,
) -> Result<AdminReceipt<Revocation>, KeyError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let row = sqlx::query(
"SELECT account_id, revoked_at_us IS NOT NULL
FROM tollgate_credential_keys WHERE key_id = $1 FOR UPDATE",
)
.bind(id_bytes(key_id.0))
.fetch_optional(&mut *tx)
.await
.map_err(storage)?
.ok_or(KeyError::UnknownKey)?;
let account_bytes: Vec<u8> = row.get(0);
let account_bytes: [u8; 16] = account_bytes
.try_into()
.map_err(|_| StoreError("credential account identifier is not 16 bytes".into()))?;
let account_id = AccountId(u128::from_be_bytes(account_bytes));
let revoked: bool = row.get(1);
let before = AdminState::Credential {
account_id,
key_id,
revoked,
};
let outcome = if revoked {
Revocation::AlreadyRetired
} else {
sqlx::query(
"UPDATE tollgate_credential_keys SET revoked_at_us = $2 WHERE key_id = $1",
)
.bind(id_bytes(key_id.0))
.bind(ts_micros(now))
.execute(&mut *tx)
.await
.map_err(storage)?;
Revocation::Retired
};
Ok(AdminReceipt::new(
outcome,
before,
AdminState::Credential {
account_id,
key_id,
revoked: true,
},
))
}
.await;
finish_transaction(tx, result).await
}
async fn publish_key_snapshot(
&self,
account: AccountId,
key: KeyId,
snapshot: PublishableSnapshot,
) -> Result<AdminReceipt<()>, KeySnapshotError> {
self.publish_key_snapshot_with_generation(account, key, snapshot, false)
.await
}
async fn publish_key_snapshot_next(
&self,
account: AccountId,
key: KeyId,
snapshot: PublishableSnapshot,
) -> Result<AdminReceipt<()>, KeySnapshotError> {
self.publish_key_snapshot_with_generation(account, key, snapshot, true)
.await
}
async fn remove_key_snapshot(
&self,
account: AccountId,
key: KeyId,
) -> Result<AdminReceipt<()>, KeySnapshotError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let (principal, _revoked) = lock_account_key(&mut tx, account, key).await?;
Ok::<_, KeySnapshotError>((principal, remove_in_tx(&mut tx, principal).await?))
}
.await;
let (principal, receipt) = finish_transaction(tx, result).await?;
self.announce_removal(principal, &receipt);
Ok(receipt)
}
async fn active_keys(&self, now: Timestamp) -> Result<Vec<KeyRecord>, StoreError> {
let cutoff = StoredInstant::from(now);
let rows = sqlx::query(
"SELECT key_id, account_id, principal, digest, not_after_floor_us, not_after_submicro_ns, not_after_is_lower_bound
FROM tollgate_credential_keys
WHERE revoked_at_us IS NULL
AND (not_after_floor_us IS NULL OR not_after_submicro_ns IS NULL
OR (not_after_floor_us, not_after_submicro_ns) > ($1, $2))
ORDER BY key_id",
)
.bind(cutoff.micros)
.bind(cutoff.submicro_nanos)
.fetch_all(&self.pool)
.await
.map_err(storage)?;
rows.into_iter().map(credential_from_row).collect()
}
async fn account_keys(
&self,
account: AccountId,
after: Option<KeyId>,
limit: NonZeroUsize,
) -> Result<Vec<KeySummary>, KeyError> {
validate_key_page_limit(limit)?;
let sql = if after.is_some() {
"SELECT page.* FROM tollgate_accounts AS account
LEFT JOIN LATERAL (
SELECT key_id, not_after_floor_us, not_after_submicro_ns,
not_after_is_lower_bound, revoked_at_us
FROM tollgate_credential_keys
WHERE account_id = account.account_id AND key_id > $3
ORDER BY key_id LIMIT $2
) AS page ON TRUE
WHERE account.account_id = $1 ORDER BY page.key_id"
} else {
"SELECT page.* FROM tollgate_accounts AS account
LEFT JOIN LATERAL (
SELECT key_id, not_after_floor_us, not_after_submicro_ns,
not_after_is_lower_bound, revoked_at_us
FROM tollgate_credential_keys
WHERE account_id = account.account_id
ORDER BY key_id LIMIT $2
) AS page ON TRUE
WHERE account.account_id = $1 ORDER BY page.key_id"
};
let mut query = sqlx::query(sql)
.bind(id_bytes(account.0))
.bind(i64::try_from(limit.get()).unwrap_or(i64::MAX));
if let Some(cursor) = after {
query = query.bind(id_bytes(cursor.0));
}
let rows = query.fetch_all(&self.pool).await.map_err(storage)?;
if rows.is_empty() {
return Err(KeyError::UnknownAccount);
}
rows.into_iter()
.filter(|row| row.get::<Option<&[u8]>, _>(0).is_some())
.map(|row| summary_from_row(row).map_err(KeyError::Storage))
.collect()
}
async fn insert_key_within(
&self,
record: KeyRecord,
max_active: NonZeroUsize,
now: Timestamp,
) -> Result<(), KeyError> {
self.insert_key_within_audited(record, max_active, now)
.await
.map(|receipt| receipt.outcome)
}
async fn insert_key_within_audited(
&self,
record: KeyRecord,
max_active: NonZeroUsize,
now: Timestamp,
) -> Result<AdminReceipt<()>, KeyError> {
self.insert_credential(record, Some((max_active, now)))
.await
}
}
impl PostgresStore {
async fn publish_key_snapshot_with_generation(
&self,
account: AccountId,
key: KeyId,
snapshot: PublishableSnapshot,
allocate_generation: bool,
) -> Result<AdminReceipt<()>, KeySnapshotError> {
let generation = if allocate_generation {
None
} else {
Some(i64::try_from(snapshot.generation.0).map_err(|_| {
StoreError("snapshot generation exceeds PostgreSQL BIGINT range".into())
})?)
};
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let (principal, revoked) = lock_account_key(&mut tx, account, key).await?;
if revoked {
return Err(KeySnapshotError::Retired { key_id: key });
}
if snapshot.key_id != Some(key) {
return Err(PublishSnapshotError::CredentialMismatch { key_id: key }.into());
}
let published = publish_in_tx(&mut tx, principal, generation, snapshot).await?;
Ok((principal, published))
}
.await;
let (principal, (written, published, before, after)) =
finish_transaction(tx, result).await?;
if written {
self.push_to_subscribers(SnapshotPush {
principal,
resolution: SnapshotResolution::Present(published),
});
}
Ok(AdminReceipt::new((), before, after))
}
async fn create_account_as(
&self,
config: AccountConfig,
origin: AdminAuthority,
) -> Result<tollgate_store::AdminReceipt<()>, CreateAccountError> {
let result = sqlx::query(
"INSERT INTO tollgate_accounts
(account_id, balance, deposited, status, capacity_class, next_fence,
usage_recorded, settlement_loss, overage_recorded, origin, status_set_by)
VALUES ($1, $2, $2, $3, $4, 1, 0, 0, 0, $5, $5)
ON CONFLICT (account_id) DO NOTHING",
)
.bind(id_bytes(config.account_id.0))
.bind(to_i64(config.initial_balance, "balance").map_err(CreateAccountError::Storage)?)
.bind(config.status.as_str())
.bind(config.capacity_class.as_str())
.bind(origin.as_str())
.execute(&self.pool)
.await
.map_err(|e| CreateAccountError::Storage(storage(e)))?;
if result.rows_affected() == 0 {
return Err(CreateAccountError::AlreadyExists);
}
Ok(AdminReceipt::new(
(),
AdminState::Absent,
AdminState::AccountCreated {
initial_balance: config.initial_balance,
status: config.status,
capacity_class: config.capacity_class,
origin,
},
))
}
async fn set_status_as(
&self,
account: AccountId,
status: AccountStatus,
authority: AdminAuthority,
) -> Result<tollgate_store::AdminReceipt<StatusChange>, SetStatusError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let row = sqlx::query(
"SELECT status, origin, status_set_by
FROM tollgate_accounts WHERE account_id = $1 FOR UPDATE",
)
.bind(id_bytes(account.0))
.fetch_optional(&mut *tx)
.await
.map_err(storage)?
.ok_or(SetStatusError::UnknownAccount)?;
let before = decode_status(row.get::<String, _>(0))?;
let origin = decode_authority(row.get::<String, _>(1))?;
let before_set_by = decode_authority(row.get::<String, _>(2))?;
if authority == AdminAuthority::Provisioner && origin != AdminAuthority::Provisioner {
return Err(SetStatusError::NotProvisioned);
}
if before == AccountStatus::Closed && status != AccountStatus::Closed {
return Err(SetStatusError::AccountClosed);
}
if authority == AdminAuthority::Provisioner
&& before != status
&& before_set_by == AdminAuthority::Operator
{
return Err(SetStatusError::OperatorHold);
}
let set_by = if authority == AdminAuthority::Provisioner && before == status {
before_set_by
} else {
authority
};
sqlx::query(
"UPDATE tollgate_accounts SET status = $2, status_set_by = $3 WHERE account_id = $1",
)
.bind(id_bytes(account.0))
.bind(status.as_str())
.bind(set_by.as_str())
.execute(&mut *tx)
.await
.map_err(storage)?;
let (republished, unreadable) =
republish_patched_snapshots(&mut tx, account, "{status}", status.as_str()).await?;
Ok(((before, before_set_by, set_by), republished, unreadable))
}
.await;
let ((before, before_set_by, set_by), republished, unreadable) =
finish_transaction(tx, result).await?;
if pushes_exceed_capacity(republished.len()) {
tracing::warn!(
%account,
principals = republished.len(),
capacity = PUSH_CHANNEL_CAPACITY,
"status change emitted more pushes than the channel holds; subscribers will resync"
);
}
let count = republished.len();
for (principal, snapshot) in republished {
self.push_to_subscribers(SnapshotPush {
principal,
resolution: SnapshotResolution::Present(snapshot),
});
}
Ok(AdminReceipt::new(
StatusChange {
republished: count + unreadable,
unreadable,
},
AdminState::Status {
status: before,
set_by: before_set_by,
},
AdminState::Status { status, set_by },
))
}
async fn insert_credential(
&self,
record: KeyRecord,
bound: Option<(NonZeroUsize, Timestamp)>,
) -> Result<AdminReceipt<()>, KeyError> {
let mut tx = self
.pool
.begin()
.await
.map_err(|e| KeyError::Storage(storage(e)))?;
let account = sqlx::query(ACCOUNT_LOCK_SQL)
.bind(id_bytes(record.account_id.0))
.fetch_optional(&mut *tx)
.await
.map_err(|e| KeyError::Storage(storage(e)))?;
if account.is_none() {
return Err(KeyError::UnknownAccount);
}
if let Some((max_active, now)) = bound {
let existing: Option<i32> = sqlx::query_scalar(
"SELECT 1 FROM tollgate_credential_keys WHERE key_id = $1 OR principal = $2",
)
.bind(id_bytes(record.key_id.0))
.bind(id_bytes(record.principal.0))
.fetch_optional(&mut *tx)
.await
.map_err(|e| KeyError::Storage(storage(e)))?;
if existing.is_some() {
return Err(KeyError::AlreadyExists);
}
let cutoff = StoredInstant::from(now);
let live: i64 = sqlx::query_scalar(LIVE_KEY_COUNT_SQL)
.bind(id_bytes(record.account_id.0))
.bind(cutoff.micros)
.bind(cutoff.submicro_nanos)
.fetch_one(&mut *tx)
.await
.map_err(|e| KeyError::Storage(storage(e)))?;
if u128::from(live.max(0).unsigned_abs())
>= u128::try_from(max_active.get()).unwrap_or(u128::MAX)
{
return Err(KeyError::ActiveKeyLimit { limit: max_active });
}
}
let expiry = record.not_after.map(StoredInstant::from);
let result = sqlx::query(
"INSERT INTO tollgate_credential_keys
(key_id, account_id, principal, digest, not_after_floor_us,
not_after_submicro_ns, not_after_is_lower_bound, revoked_at_us)
VALUES ($1, $2, $3, $4, $5, $6, FALSE, NULL)
ON CONFLICT (key_id) DO NOTHING",
)
.bind(id_bytes(record.key_id.0))
.bind(id_bytes(record.account_id.0))
.bind(id_bytes(record.principal.0))
.bind(record.digest.to_vec())
.bind(expiry.map(|expiry| expiry.micros))
.bind(expiry.map(|expiry| expiry.submicro_nanos))
.execute(&mut *tx)
.await;
let outcome = match result {
Ok(done) if done.rows_affected() == 0 => Err(KeyError::AlreadyExists),
Ok(_) => Ok(()),
Err(sqlx::Error::Database(e)) if e.is_unique_violation() => {
Err(KeyError::AlreadyExists)
}
Err(e) => Err(KeyError::Storage(storage(e))),
};
outcome?;
tx.commit()
.await
.map_err(|e| KeyError::Storage(storage(e)))?;
Ok(AdminReceipt::new(
(),
AdminState::Absent,
AdminState::Credential {
account_id: record.account_id,
key_id: record.key_id,
revoked: false,
},
))
}
}
fn summary_from_row(row: sqlx::postgres::PgRow) -> Result<KeySummary, StoreError> {
let bytes: Vec<u8> = row.get(0);
let fixed: [u8; 16] = bytes
.as_slice()
.try_into()
.map_err(|_| StoreError("credential identifier is not 16 bytes".into()))?;
let not_after = match (row.get::<Option<i64>, _>(1), row.get::<Option<i16>, _>(2)) {
(None, None) if !row.get::<bool, _>(3) => None,
(Some(micros), Some(submicro_nanos)) => Some(
StoredInstant {
micros,
submicro_nanos,
}
.timestamp()?,
),
_ => return Err(StoreError("incomplete stored credential expiry".into())),
};
Ok(KeySummary {
key_id: KeyId(u128::from_be_bytes(fixed)),
not_after,
revoked_at: row
.get::<Option<i64>, _>(4)
.map(|micros| {
StoredInstant {
micros,
submicro_nanos: 0,
}
.timestamp()
})
.transpose()?,
})
}
fn credential_from_row(row: sqlx::postgres::PgRow) -> Result<KeyRecord, StoreError> {
let id = |index| -> Result<u128, StoreError> {
let bytes: Vec<u8> = row.get(index);
let fixed: [u8; 16] = bytes
.as_slice()
.try_into()
.map_err(|_| StoreError("credential identifier is not 16 bytes".into()))?;
Ok(u128::from_be_bytes(fixed))
};
let key_id = KeyId(id(0)?);
let not_after = match (row.get::<Option<i64>, _>(4), row.get::<Option<i16>, _>(5)) {
(None, None) if !row.get::<bool, _>(6) => None,
(Some(micros), Some(submicro_nanos)) => Some(
StoredInstant {
micros,
submicro_nanos,
}
.timestamp()?,
),
_ => return Err(StoreError("incomplete stored credential expiry".into())),
};
Ok(KeyRecord {
key_id,
account_id: AccountId(id(1)?),
principal: Principal(id(2)?),
digest: digest_from(row.get::<Vec<u8>, _>(3).as_slice(), key_id)?,
not_after,
})
}
#[async_trait]
impl tollgate_store::KeySource for PostgresStore {
async fn active_keys_page(
&self,
now: Timestamp,
after: Option<KeyId>,
limit: NonZeroUsize,
) -> Result<tollgate_store::KeyPage, StoreError> {
validate_key_page_limit(limit)?;
let cutoff = StoredInstant::from(now);
let mut tx = self.pool.begin().await.map_err(storage)?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY")
.execute(&mut *tx)
.await
.map_err(storage)?;
let revision: i64 =
sqlx::query_scalar("SELECT revision FROM tollgate_credential_revision WHERE singleton")
.fetch_one(&mut *tx)
.await
.map_err(storage)?;
let revision = u64::try_from(revision)
.map_err(|_| StoreError("credential revision is negative".into()))?;
let sql = if after.is_some() {
"SELECT key_id, account_id, principal, digest, not_after_floor_us, not_after_submicro_ns, not_after_is_lower_bound
FROM tollgate_credential_keys WHERE revoked_at_us IS NULL
AND (not_after_floor_us IS NULL OR not_after_submicro_ns IS NULL
OR (not_after_floor_us, not_after_submicro_ns) > ($1, $2)) AND key_id > $3
ORDER BY key_id LIMIT $4"
} else {
"SELECT key_id, account_id, principal, digest, not_after_floor_us, not_after_submicro_ns, not_after_is_lower_bound
FROM tollgate_credential_keys WHERE revoked_at_us IS NULL
AND (not_after_floor_us IS NULL OR not_after_submicro_ns IS NULL
OR (not_after_floor_us, not_after_submicro_ns) > ($1, $2)) AND key_id >= $3
ORDER BY key_id LIMIT $4"
};
let rows = sqlx::query(sql)
.bind(cutoff.micros)
.bind(cutoff.submicro_nanos)
.bind(id_bytes(after.unwrap_or(KeyId(0)).0))
.bind((limit.get() + 1) as i64)
.fetch_all(&mut *tx)
.await
.map_err(storage)?;
tx.commit().await.map_err(storage)?;
let mut records = rows
.into_iter()
.map(credential_from_row)
.map(|row| row.map(tollgate_store::CredentialRecord::from))
.collect::<Result<Vec<_>, _>>()?;
let next_after = if records.len() > limit.get() {
records.pop();
records.last().map(|record| record.key_id)
} else {
None
};
tollgate_store::KeyPage::try_new(revision, now, after, limit, records, next_after)
}
}
#[cfg(test)]
mod stored_id_tests {
use super::StoredId;
#[test]
fn malformed_storage_id_explains_both_accepted_representations() {
let error = serde_json::from_value::<StoredId>(serde_json::Value::Bool(true)).unwrap_err();
assert!(
error
.to_string()
.contains("a legacy u64 number or a canonical 128-bit identifier string"),
"unexpected diagnostic: {error}"
);
}
}
#[derive(Serialize)]
struct StoredSnapshotRef<'a> {
account_id: StoredId,
key_id: Option<StoredId>,
status: &'a AccountStatus,
capacity_class: &'a CapacityClass,
enforcement_mode: &'a EnforcementMode,
valid_until: &'a Timestamp,
permissions: &'a PermissionBits,
limits: &'a ResolvedLimits,
cost_table: &'a Arc<CostTable>,
budget: Option<&'a BudgetView>,
policy_revision: &'a PolicyRevision,
}
impl<'a> From<&'a AccountSnapshot> for StoredSnapshotRef<'a> {
fn from(snapshot: &'a AccountSnapshot) -> Self {
StoredSnapshotRef {
account_id: StoredId(snapshot.account_id.0),
key_id: snapshot.key_id.map(|id| StoredId(id.0)),
status: &snapshot.status,
capacity_class: &snapshot.capacity_class,
enforcement_mode: &snapshot.enforcement_mode,
valid_until: &snapshot.valid_until,
permissions: &snapshot.permissions,
limits: &snapshot.limits,
cost_table: &snapshot.cost_table,
budget: snapshot.budget.as_ref(),
policy_revision: &snapshot.policy_revision,
}
}
}
#[derive(Deserialize)]
struct StoredSnapshot {
account_id: StoredId,
key_id: Option<StoredId>,
status: AccountStatus,
#[serde(default)]
capacity_class: CapacityClass,
#[serde(default)]
enforcement_mode: EnforcementMode,
valid_until: Timestamp,
permissions: PermissionBits,
limits: ResolvedLimits,
cost_table: Arc<CostTable>,
#[serde(default)]
budget: Option<BudgetView>,
#[serde(default)]
policy_revision: PolicyRevision,
}
impl StoredSnapshot {
fn into_snapshot(self, generation: Generation) -> AccountSnapshot {
let builder = AccountSnapshot::builder(
AccountId(self.account_id.0),
generation,
self.status,
self.valid_until,
self.permissions,
self.limits,
self.cost_table,
)
.enforcement_mode(self.enforcement_mode)
.capacity_class(self.capacity_class)
.policy_revision(self.policy_revision);
match self.key_id {
Some(key_id) => builder.key_id(tollgate_core::KeyId(key_id.0)).build(),
None => builder.build(),
}
}
}
const ACTIVE_LEASE_SUM_SQL: &str =
"SELECT COALESCE(SUM(granted), 0)::BIGINT, COALESCE(SUM(used), 0)::BIGINT
FROM tollgate_leases WHERE account_id = $1 AND state = 0";
pub(crate) const ACCOUNT_LOCK_SQL: &str =
"SELECT 1 FROM tollgate_accounts WHERE account_id = $1 FOR UPDATE";
pub(crate) const LIVE_KEY_COUNT_SQL: &str = "SELECT count(*) FROM tollgate_credential_keys
WHERE account_id = $1
AND revoked_at_us IS NULL
AND (not_after_floor_us IS NULL OR not_after_submicro_ns IS NULL
OR (not_after_floor_us, not_after_submicro_ns) > ($2, $3))";
const RECLAIM_DUE_LEASES_SQL: &str = "SELECT lease_id, account_id, granted, used
FROM tollgate_leases
WHERE state = 0 AND (expires_at_floor_us, expires_at_submicro_ns) <= ($1, $2)
ORDER BY expires_at_floor_us, expires_at_submicro_ns
LIMIT $3 FOR UPDATE SKIP LOCKED";
const DUE_PERIODS_SQL: &str = "WITH due AS (
SELECT account_id, allowance_balance AS prior, budget_allowance AS allowance,
period_start_us AS crossed_from
FROM tollgate_accounts
WHERE budget_allowance IS NOT NULL
AND budget_period = $1 AND period_start_us < $2
ORDER BY period_start_us LIMIT $3 FOR UPDATE SKIP LOCKED
)
UPDATE tollgate_accounts AS account SET
deposited = account.deposited + due.allowance,
expired = account.expired + due.prior,
balance = account.balance - due.prior + due.allowance,
allowance_balance = due.allowance,
period_start_us = $2
FROM due
WHERE account.account_id = due.account_id
RETURNING due.account_id, due.allowance, due.prior, due.crossed_from";
fn id_bytes(id: u128) -> Vec<u8> {
id.to_be_bytes().to_vec()
}
fn id_from(bytes: &[u8]) -> u128 {
let mut buf = [0u8; 16];
buf.copy_from_slice(bytes);
u128::from_be_bytes(buf)
}
fn to_i64(units: CostUnits, what: &str) -> Result<i64, StoreError> {
i64::try_from(units.get()).map_err(|_| StoreError(format!("{what} exceeds i64 range")))
}
fn stored_fence(value: i64) -> Result<FencingToken, StoreError> {
u64::try_from(value)
.ok()
.filter(|value| *value > 0)
.map(FencingToken)
.ok_or_else(|| StoreError(format!("stored fencing token is not positive: {value}")))
}
fn to_units(value: i64, what: &str) -> Result<CostUnits, StoreError> {
u64::try_from(value)
.map(CostUnits)
.map_err(|_| StoreError(format!("{what} is negative in storage: {value}")))
}
fn ts_micros(ts: Timestamp) -> i64 {
ts.as_microsecond()
}
fn micros_ts(value: i64, what: &str) -> Result<Timestamp, StoreError> {
tollgate_store::clock::timestamp_from_micros(value).map_err(|e| {
StoreError(format!(
"{what} is not a representable instant: {value} ({e})"
))
})
}
fn storage(e: sqlx::Error) -> StoreError {
StoreError(format!("postgres: {e}"))
}
fn alloc_storage(e: sqlx::Error) -> AllocateError {
AllocateError::Storage(storage(e))
}
async fn finish_transaction<T, E>(
tx: Transaction<'_, Postgres>,
result: Result<T, E>,
) -> Result<T, E>
where
E: From<StoreError> + std::fmt::Display,
{
match result {
Ok(value) => {
tx.commit().await.map_err(|e| E::from(storage(e)))?;
Ok(value)
}
Err(error) => {
if let Err(rollback_error) = tx.rollback().await {
return Err(E::from(StoreError(format!(
"operation failed ({error}); transaction rollback failed ({})",
storage(rollback_error)
))));
}
Err(error)
}
}
}
pub struct PostgresStore {
pool: PgPool,
policy: GrantPolicy,
push: broadcast::Sender<SnapshotPush>,
}
#[derive(Debug, Clone, Copy)]
pub struct PoolConfig {
pub max_connections: u32,
pub acquire_timeout: std::time::Duration,
}
impl Default for PoolConfig {
fn default() -> Self {
PoolConfig {
max_connections: 16,
acquire_timeout: std::time::Duration::from_secs(5),
}
}
}
impl PoolConfig {
pub fn validate(&self) -> Result<(), StoreError> {
if self.max_connections == 0 {
return Err(StoreError("max_connections must be positive".into()));
}
if self.acquire_timeout.is_zero() {
return Err(StoreError("acquire_timeout must be positive".into()));
}
Ok(())
}
}
impl PostgresStore {
pub async fn connect(url: &str, policy: GrantPolicy) -> Result<Arc<Self>, StoreError> {
Self::connect_with(url, policy, PoolConfig::default()).await
}
pub async fn connect_with(
url: &str,
policy: GrantPolicy,
pool_config: PoolConfig,
) -> Result<Arc<Self>, StoreError> {
policy
.validate()
.map_err(|e| StoreError(format!("invalid grant policy: {e}")))?;
pool_config.validate()?;
let pool = PgPoolOptions::new()
.max_connections(pool_config.max_connections)
.acquire_timeout(pool_config.acquire_timeout)
.connect(url)
.await
.map_err(storage)?;
sqlx::migrate!("./migrations")
.run(&pool)
.await
.map_err(|e| StoreError(format!("migrate: {e}")))?;
let (push, _) = broadcast::channel(PUSH_CHANNEL_CAPACITY);
Ok(Arc::new(PostgresStore { pool, policy, push }))
}
fn push_to_subscribers(&self, push: SnapshotPush) {
let principal = push.principal;
let subscribers = self.push.send(push).unwrap_or(0);
tracing::debug!(
%principal,
subscribers,
"snapshot pushed to subscribers"
);
}
fn announce_removal(&self, principal: Principal, receipt: &AdminReceipt<()>) {
if receipt.before != receipt.after
&& let AdminState::Snapshot { generation, .. } = receipt.after
{
self.push_to_subscribers(SnapshotPush {
principal,
resolution: SnapshotResolution::Revoked { generation },
});
}
}
pub async fn balance(&self, account: AccountId) -> Result<CostUnits, StoreError> {
let row = sqlx::query("SELECT balance FROM tollgate_accounts WHERE account_id = $1")
.bind(id_bytes(account.0))
.fetch_optional(&self.pool)
.await
.map_err(storage)?;
row.map(|r| to_units(r.get::<i64, _>(0), "account balance"))
.transpose()
.map(|units| units.unwrap_or(CostUnits::ZERO))
}
pub async fn usage_recorded(&self, account: AccountId) -> Result<CostUnits, StoreError> {
let row = sqlx::query("SELECT usage_recorded FROM tollgate_accounts WHERE account_id = $1")
.bind(id_bytes(account.0))
.fetch_optional(&self.pool)
.await
.map_err(storage)?;
row.map(|r| to_units(r.get::<i64, _>(0), "account usage_recorded"))
.transpose()
.map(|units| units.unwrap_or(CostUnits::ZERO))
}
pub async fn conservation(
&self,
account: AccountId,
) -> Result<Option<Conservation>, StoreError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
sqlx::query("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY")
.execute(&mut *tx)
.await
.map_err(storage)?;
let account_row = sqlx::query(
"SELECT deposited, balance, usage_recorded, settlement_loss, overage_recorded,
expired
FROM tollgate_accounts WHERE account_id = $1",
)
.bind(id_bytes(account.0))
.fetch_optional(&mut *tx)
.await
.map_err(storage)?;
let lease_row = sqlx::query(ACTIVE_LEASE_SUM_SQL)
.bind(id_bytes(account.0))
.fetch_one(&mut *tx)
.await
.map_err(storage)?;
Ok::<_, StoreError>((account_row, lease_row))
}
.await;
let (account_row, lease_row) = finish_transaction(tx, result).await?;
let Some(row) = account_row else {
return Ok(None);
};
let active_grants = to_units(lease_row.get::<i64, _>(0), "active lease grants")?;
let active_used = to_units(lease_row.get::<i64, _>(1), "active lease usage")?;
let recorded = to_units(row.get::<i64, _>(2), "usage_recorded")?;
Ok(Some(Conservation {
deposited: to_units(row.get::<i64, _>(0), "deposited")?,
overage_recorded: to_units(row.get::<i64, _>(4), "overage_recorded")?,
balance: to_units(row.get::<i64, _>(1), "balance")?,
active_lease_grants: active_grants,
settled_usage: recorded.checked_sub(active_used).ok_or_else(|| {
StoreError(format!(
"active lease usage {} exceeds recorded usage {} for account {account}",
active_used.get(),
recorded.get()
))
})?,
settlement_loss: to_units(row.get::<i64, _>(3), "settlement_loss")?,
expired: to_units(row.get::<i64, _>(5), "expired")?,
}))
}
}
struct LockedLeaseRow {
account_id: Vec<u8>,
fencing_token: i64,
granted: i64,
used: i64,
credited: i64,
expires_at: StoredInstant,
state: i16,
from_allowance: i64,
period_start_us: i64,
}
struct ReleasedCredit {
account: AccountId,
restored: CostUnits,
preserves_funding: bool,
}
async fn lock_lease(
tx: &mut Transaction<'_, Postgres>,
lease_id: LeaseId,
) -> Result<Option<LockedLeaseRow>, sqlx::Error> {
let row = sqlx::query(
"SELECT account_id, fencing_token, granted, used, credited, expires_at_floor_us, state,
from_allowance, period_start_us, expires_at_submicro_ns
FROM tollgate_leases WHERE lease_id = $1 FOR UPDATE",
)
.bind(id_bytes(lease_id.0))
.fetch_optional(&mut **tx)
.await?;
Ok(row.map(|row| LockedLeaseRow {
account_id: row.get(0),
fencing_token: row.get(1),
granted: row.get(2),
used: row.get(3),
credited: row.get(4),
expires_at: StoredInstant {
micros: row.get(5),
submicro_nanos: row.get(9),
},
state: row.get(6),
from_allowance: row.get(7),
period_start_us: row.get(8),
}))
}
struct Exchange {
floor: CostUnits,
needed: CostUnits,
preserves_funding: bool,
}
impl Exchange {
const ACQUIRE: Exchange = Exchange {
floor: CostUnits::ZERO,
needed: CostUnits::ZERO,
preserves_funding: true,
};
}
impl PostgresStore {
async fn acquire_in_tx(
&self,
tx: &mut Transaction<'_, Postgres>,
account: AccountId,
requested: CostUnits,
expires_at: Timestamp,
exchange: Exchange,
) -> Result<Allocation, AllocateError> {
let Exchange {
floor,
needed,
preserves_funding: settlement_preserves_funding,
} = exchange;
let row = sqlx::query(
"SELECT balance, status, next_fence, allowance_balance, period_start_us
FROM tollgate_accounts WHERE account_id = $1 FOR UPDATE",
)
.bind(id_bytes(account.0))
.fetch_optional(&mut **tx)
.await
.map_err(alloc_storage)?
.ok_or(AllocateError::UnknownAccount)?;
if decode_status(row.get::<String, _>(1)).map_err(AllocateError::Storage)?
!= AccountStatus::Active
{
return Err(AllocateError::AccountInactive);
}
let balance =
to_units(row.get::<i64, _>(0), "account balance").map_err(AllocateError::Storage)?;
let granted = match self
.policy
.consolidation_grant(requested, balance, floor, needed)
{
Some(granted) => granted,
None => {
if settlement_preserves_funding && !requested.is_zero() {
let budget = sqlx::query(BUDGET_VIEW_SQL)
.bind(id_bytes(account.0))
.fetch_one(&mut **tx)
.await
.map_err(alloc_storage)?;
let evidence = budget_view(&budget)
.map_err(AllocateError::Storage)?
.shortfall();
return Err(match evidence.exhaustion() {
Some(exhausted) => AllocateError::BalanceExhausted(exhausted),
None => AllocateError::BalanceInsufficient(evidence),
});
}
return Err(AllocateError::InsufficientBalance);
}
};
let granted_i = to_i64(granted, "grant").map_err(AllocateError::Storage)?;
let allowance_balance = row.get::<i64, _>(3);
let from_allowance = granted_i.min(allowance_balance);
let period_start_us = row.get::<i64, _>(4);
let fence = row.get::<i64, _>(2);
let fence_token = stored_fence(fence).map_err(AllocateError::Storage)?;
let lease_id = LeaseId(uuid::Uuid::new_v4().as_u128());
let budget = sqlx::query(GRANT_DEBIT_SQL)
.bind(id_bytes(account.0))
.bind(granted_i)
.bind(from_allowance)
.fetch_one(&mut **tx)
.await
.map_err(alloc_storage)?;
let funding = budget_view(&budget)
.map_err(AllocateError::Storage)?
.shortfall();
sqlx::query(
"INSERT INTO tollgate_leases
(lease_id, account_id, fencing_token, granted, used, credited, expires_at_floor_us,
state, from_allowance, period_start_us, expires_at_submicro_ns, expiry_is_upper_bound)
VALUES ($1, $2, $3, $4, 0, 0, $5, 0, $6, $7, $8, FALSE)",
)
.bind(id_bytes(lease_id.0))
.bind(id_bytes(account.0))
.bind(fence)
.bind(granted_i)
.bind(StoredInstant::from(expires_at).micros)
.bind(from_allowance)
.bind(period_start_us)
.bind(StoredInstant::from(expires_at).submicro_nanos)
.execute(&mut **tx)
.await
.map_err(alloc_storage)?;
Ok(Allocation {
grant: LeaseGrant {
lease_id,
account_id: account,
fencing_token: fence_token,
units: granted,
expires_at,
},
funding: Some(funding),
})
}
async fn release_in_tx(
&self,
tx: &mut Transaction<'_, Postgres>,
lease_id: LeaseId,
fencing_token: FencingToken,
unspent: CostUnits,
now: Timestamp,
) -> Result<ReleasedCredit, AllocateError> {
let LockedLeaseRow {
account_id,
fencing_token: fence,
granted,
used,
credited: _credited,
expires_at,
state,
from_allowance,
period_start_us,
} = lock_lease(tx, lease_id)
.await
.map_err(alloc_storage)?
.ok_or(AllocateError::UnknownLease)?;
if stored_fence(fence).map_err(AllocateError::Storage)? != fencing_token {
return Err(AllocateError::Fenced);
}
let expires_at = expires_at.timestamp().map_err(AllocateError::Storage)?;
if state != STATE_ACTIVE
|| self
.policy
.reclaim_cutoff(now)
.is_some_and(|cutoff| expires_at <= cutoff)
{
return Err(AllocateError::LeaseNotActive);
}
let account = AccountId(id_from(&account_id));
let unspent_i = to_i64(unspent, "unspent").map_err(AllocateError::Storage)?;
let loss = granted
.checked_sub(
used.checked_add(unspent_i)
.ok_or(AllocateError::InvalidRelease)?,
)
.filter(|l| *l >= 0)
.ok_or(AllocateError::InvalidRelease)?;
sqlx::query("UPDATE tollgate_leases SET state = $2, credited = $3 WHERE lease_id = $1")
.bind(id_bytes(lease_id.0))
.bind(STATE_RELEASED)
.bind(unspent_i)
.execute(&mut **tx)
.await
.map_err(alloc_storage)?;
let from_topup = granted
.checked_sub(from_allowance)
.filter(|t| *t >= 0)
.ok_or_else(|| {
AllocateError::Storage(StoreError(format!(
"lease allowance funding {from_allowance} exceeds its grant {granted}"
)))
})?;
let to_topup = from_topup.min(unspent_i);
let to_allowance = unspent_i - to_topup;
let restored: i64 = sqlx::query_scalar(
"UPDATE tollgate_accounts SET
balance = balance + $2
+ CASE WHEN period_start_us > $5 THEN 0 ELSE $3 END,
allowance_balance = allowance_balance
+ CASE WHEN period_start_us > $5 THEN 0 ELSE $3 END,
expired = expired + CASE WHEN period_start_us > $5 THEN $3 ELSE 0 END,
settlement_loss = settlement_loss + $4
WHERE account_id = $1
RETURNING $2 + CASE WHEN period_start_us > $5 THEN 0 ELSE $3 END",
)
.bind(account_id)
.bind(to_topup)
.bind(to_allowance)
.bind(loss)
.bind(period_start_us)
.fetch_one(&mut **tx)
.await
.map_err(alloc_storage)?;
Ok(ReleasedCredit {
account,
preserves_funding: loss == 0 && restored == unspent_i,
restored: to_units(restored, "restored release credit")
.map_err(AllocateError::Storage)?,
})
}
fn grant_expiry(
&self,
ttl: SignedDuration,
now: Timestamp,
) -> Result<Timestamp, AllocateError> {
if ttl <= SignedDuration::ZERO {
return Err(AllocateError::InvalidTtl);
}
now.checked_add(ttl.min(self.policy.max_ttl))
.map_err(|e| AllocateError::Storage(StoreError(format!("ttl overflow: {e}"))))
}
}
#[async_trait]
impl LeaseAllocator for PostgresStore {
async fn acquire(
&self,
account: AccountId,
requested: CostUnits,
ttl: SignedDuration,
now: Timestamp,
) -> Result<Allocation, AllocateError> {
let expires_at = self.grant_expiry(ttl, now)?;
let mut tx = self.pool.begin().await.map_err(alloc_storage)?;
let result = self
.acquire_in_tx(&mut tx, account, requested, expires_at, Exchange::ACQUIRE)
.await;
finish_transaction(tx, result).await
}
async fn release(
&self,
lease_id: LeaseId,
fencing_token: FencingToken,
unspent: CostUnits,
now: Timestamp,
) -> Result<(), AllocateError> {
let mut tx = self.pool.begin().await.map_err(alloc_storage)?;
let result = self
.release_in_tx(&mut tx, lease_id, fencing_token, unspent, now)
.await
.map(|_| ());
finish_transaction(tx, result).await
}
async fn consolidate(
&self,
lease_id: LeaseId,
fencing_token: FencingToken,
unspent: CostUnits,
requested: CostUnits,
needed: CostUnits,
ttl: SignedDuration,
now: Timestamp,
) -> Result<Allocation, AllocateError> {
let expires_at = self.grant_expiry(ttl, now)?;
let mut tx = self.pool.begin().await.map_err(alloc_storage)?;
let result = async {
let released = self
.release_in_tx(&mut tx, lease_id, fencing_token, unspent, now)
.await?;
self.acquire_in_tx(
&mut tx,
released.account,
requested,
expires_at,
Exchange {
floor: released.restored,
needed,
preserves_funding: released.preserves_funding,
},
)
.await
}
.await;
finish_transaction(tx, result).await
}
async fn reclaim_expired_batch(
&self,
now: Timestamp,
limit: NonZeroUsize,
) -> Result<ReclaimBatch, StoreError> {
let limit_i = i64::try_from(limit.get())
.map_err(|_| StoreError(format!("reclaim batch limit exceeds i64 range: {limit}")))?;
let Some(cutoff) = self.policy.reclaim_cutoff(now) else {
return ReclaimBatch::try_new(Vec::new(), limit);
};
let cutoff = StoredInstant::from(cutoff);
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let rows = sqlx::query(RECLAIM_DUE_LEASES_SQL)
.bind(cutoff.micros)
.bind(cutoff.submicro_nanos)
.bind(limit_i)
.fetch_all(&mut *tx)
.await
.map_err(storage)?;
let mut reclaimed = Vec::with_capacity(rows.len());
let mut lease_ids = Vec::with_capacity(rows.len());
let mut forfeits: std::collections::BTreeMap<Vec<u8>, i64> =
std::collections::BTreeMap::new();
for row in rows {
let lease_bytes: Vec<u8> = row.get(0);
let account_bytes: Vec<u8> = row.get(1);
let granted = row.get::<i64, _>(2);
let used = row.get::<i64, _>(3);
let forfeited = granted.checked_sub(used).ok_or_else(|| {
StoreError(format!(
"reclaim remainder overflow: granted {granted}, used {used}"
))
})?;
let forfeited_units = to_units(forfeited, "reclaim remainder")?;
let total = forfeits.entry(account_bytes.clone()).or_default();
*total = total
.checked_add(forfeited)
.ok_or_else(|| StoreError("reclaim loss sum overflow".into()))?;
lease_ids.push(lease_bytes.clone());
reclaimed.push(ReclaimedLease {
lease_id: LeaseId(id_from(&lease_bytes)),
account_id: AccountId(id_from(&account_bytes)),
forfeited: forfeited_units,
});
}
let batch = ReclaimBatch::try_new(reclaimed, limit)?;
if batch.is_empty() {
return Ok(batch);
}
let (account_ids, account_forfeits): (Vec<_>, Vec<_>) = forfeits.into_iter().unzip();
let expected_lease_rows = u64::try_from(lease_ids.len())
.map_err(|_| StoreError("reclaim lease row count exceeds u64 range".into()))?;
let expected_account_rows = u64::try_from(account_ids.len())
.map_err(|_| StoreError("reclaim account row count exceeds u64 range".into()))?;
let locked_accounts = sqlx::query(
"SELECT account_id FROM tollgate_accounts
WHERE account_id = ANY($1) ORDER BY account_id FOR UPDATE",
)
.bind(&account_ids)
.fetch_all(&mut *tx)
.await
.map_err(storage)?;
if locked_accounts.len() != account_ids.len() {
return Err(StoreError(format!(
"reclaim locked {} of {} referenced account rows",
locked_accounts.len(),
account_ids.len()
)));
}
let updated_leases = sqlx::query(
"UPDATE tollgate_leases
SET state = $2, credited = 0
WHERE lease_id = ANY($1) AND state = $3",
)
.bind(&lease_ids)
.bind(STATE_EXPIRED)
.bind(STATE_ACTIVE)
.execute(&mut *tx)
.await
.map_err(storage)?;
if updated_leases.rows_affected() != expected_lease_rows {
return Err(StoreError(format!(
"reclaim updated {} of {} locked lease rows",
updated_leases.rows_affected(),
lease_ids.len()
)));
}
let updated_accounts = sqlx::query(
"UPDATE tollgate_accounts AS account
SET settlement_loss = account.settlement_loss + delta.forfeited
FROM UNNEST($1::bytea[], $2::bigint[]) AS delta(account_id, forfeited)
WHERE account.account_id = delta.account_id",
)
.bind(&account_ids)
.bind(&account_forfeits)
.execute(&mut *tx)
.await
.map_err(storage)?;
if updated_accounts.rows_affected() != expected_account_rows {
return Err(StoreError(format!(
"reclaim updated {} of {} locked account rows",
updated_accounts.rows_affected(),
account_ids.len()
)));
}
Ok(batch)
}
.await;
finish_transaction(tx, result).await
}
}
#[async_trait]
impl UsageSink for PostgresStore {
async fn ingest(
&self,
events: &[UsageEvent],
_now: Timestamp,
) -> Result<IngestReport, IngestError> {
let mut report = IngestReport {
unattributed: Some(0),
..IngestReport::default()
};
if events.is_empty() {
return Ok(report);
}
struct PreparedEvent<'a> {
event: &'a UsageEvent,
request_id: Vec<u8>,
account_id: Vec<u8>,
key_id: Option<Vec<u8>>,
lease_id: Option<Vec<u8>>,
occurred_at_us: i64,
}
let prepared: Vec<PreparedEvent<'_>> = events
.iter()
.map(|event| {
Ok(PreparedEvent {
event,
request_id: id_bytes(event.request_id.0),
account_id: id_bytes(event.account_id.0),
key_id: event.key_id.map(|id| id_bytes(id.0)),
lease_id: event.source.lease_id().map(|id| id_bytes(id.0)),
occurred_at_us: ts_micros(event.occurred_at),
})
})
.collect::<Result<_, StoreError>>()?;
let mut tx = self.pool.begin().await.map_err(storage)?;
let result: Result<_, IngestError> = async {
let lease_ids: Vec<Vec<u8>> = prepared
.iter()
.filter_map(|event| event.lease_id.clone())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect();
struct LeaseRow {
account_id: Vec<u8>,
fence: i64,
granted: i64,
used: i64,
used_delta: Option<NonZeroI64>,
credited: i64,
settled: bool,
}
let rows = sqlx::query(
"SELECT lease_id, account_id, fencing_token, granted, used, credited, state
FROM tollgate_leases
WHERE lease_id = ANY($1)
ORDER BY account_id, lease_id FOR UPDATE",
)
.bind(&lease_ids)
.fetch_all(&mut *tx)
.await
.map_err(storage)?;
let mut leases: std::collections::BTreeMap<Vec<u8>, LeaseRow> =
std::collections::BTreeMap::new();
for row in rows {
let lease_id: Vec<u8> = row.get(0);
let fence: i64 = row.get(2);
let granted: i64 = row.get(3);
let used: i64 = row.get(4);
let credited: i64 = row.get(5);
leases.insert(
lease_id,
LeaseRow {
account_id: row.get(1),
fence,
granted,
used,
used_delta: None,
credited,
settled: row.get::<i16, _>(6) != STATE_ACTIVE,
},
);
}
let request_ids: Vec<Vec<u8>> = prepared
.iter()
.map(|event| event.request_id.clone())
.collect();
let mut seen: std::collections::HashSet<Vec<u8>> = sqlx::query(
"SELECT request_id FROM tollgate_usage_events WHERE request_id = ANY($1)",
)
.bind(&request_ids)
.fetch_all(&mut *tx)
.await
.map_err(storage)?
.into_iter()
.map(|row| row.get::<Vec<u8>, _>(0))
.collect();
let overage_account_ids: Vec<Vec<u8>> = prepared
.iter()
.filter(|event| event.lease_id.is_none())
.map(|event| event.account_id.clone())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect();
let known_overage_accounts: std::collections::HashSet<Vec<u8>> =
if overage_account_ids.is_empty() {
std::collections::HashSet::new()
} else {
sqlx::query(
"SELECT account_id FROM tollgate_accounts WHERE account_id = ANY($1)",
)
.bind(&overage_account_ids)
.fetch_all(&mut *tx)
.await
.map_err(storage)?
.into_iter()
.map(|row| row.get::<Vec<u8>, _>(0))
.collect()
};
struct Accepted {
event_index: usize,
settled: bool,
fence: Option<i64>,
overage: bool,
units: i64,
}
let mut accepted: Vec<Accepted> = Vec::with_capacity(prepared.len());
for (event_index, event) in prepared.iter().enumerate() {
if seen.contains(event.request_id.as_slice()) {
report.duplicate += 1;
continue;
}
let Some(lease_key) = event.lease_id.as_deref() else {
if !known_overage_accounts.contains(event.account_id.as_slice()) {
report.rejected += 1;
continue;
}
let Ok(units) = i64::try_from(event.event.units.get()) else {
report.rejected += 1;
continue;
};
seen.insert(event.request_id.clone());
accepted.push(Accepted {
event_index,
settled: false,
fence: None,
overage: true,
units,
});
report.accepted += 1;
continue;
};
let Some(lease) = leases.get_mut(lease_key) else {
report.rejected += 1;
continue;
};
if Some(stored_fence(lease.fence)?) != event.event.source.fencing_token()
|| lease.account_id.as_slice() != event.account_id.as_slice()
{
report.rejected += 1;
continue;
}
to_units(lease.granted, "lease granted")?;
to_units(lease.used, "lease used")?;
to_units(lease.credited, "lease credited")?;
let Ok(units) = i64::try_from(event.event.units.get()) else {
report.rejected += 1;
continue;
};
let committed = lease
.used
.checked_add(lease.used_delta.map_or(0, NonZeroI64::get))
.and_then(|used| used.checked_add(lease.credited))
.ok_or_else(|| {
StoreError(format!(
"lease accounting overflow for {:#034x}",
id_from(lease_key)
))
})?;
let remaining = lease.granted.checked_sub(committed).ok_or_else(|| {
StoreError(format!(
"lease accounting exceeds grant for {:#034x}: granted {}, committed {committed}",
id_from(lease_key), lease.granted
))
})?;
if units > remaining {
report.rejected += 1;
continue;
}
let used_delta = lease
.used_delta
.map_or(0, NonZeroI64::get)
.checked_add(units)
.ok_or_else(|| {
StoreError(format!(
"lease usage delta overflow for {:#034x}",
id_from(lease_key)
))
})?;
lease.used_delta = NonZeroI64::new(used_delta);
seen.insert(event.request_id.clone());
accepted.push(Accepted {
event_index,
settled: lease.settled,
fence: Some(lease.fence),
overage: false,
units,
});
report.accepted += 1;
}
if accepted.is_empty() {
return Ok(report);
}
#[derive(Default)]
struct AccountDelta {
usage: i64,
loss: i64,
overage: i64,
}
let (mut rid, mut acct, mut lease, mut fence, mut units, mut at, mut revision, mut keys) = (
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
Vec::with_capacity(accepted.len()),
);
let mut account_deltas: std::collections::BTreeMap<Vec<u8>, AccountDelta> =
std::collections::BTreeMap::new();
for accepted_event in &accepted {
let event = &prepared[accepted_event.event_index];
rid.push(event.request_id.clone());
acct.push(event.account_id.clone());
keys.push(event.key_id.clone());
lease.push(event.lease_id.clone());
fence.push(accepted_event.fence);
units.push(accepted_event.units);
at.push(event.occurred_at_us);
revision.push(event.event.policy_revision.as_bytes().to_vec());
if accepted_event.overage {
debug_assert!(
event.lease_id.is_none() && accepted_event.fence.is_none(),
"an overage row must carry neither half of a capability"
);
}
let entry = account_deltas.entry(event.account_id.clone()).or_default();
entry.usage = entry.usage.checked_add(accepted_event.units).ok_or_else(|| {
IngestError::Refused(StoreError(format!(
"usage delta overflow for account {:#034x}",
event.event.account_id.0
)))
})?;
if accepted_event.settled {
entry.loss = entry.loss.checked_add(accepted_event.units).ok_or_else(|| {
IngestError::Refused(StoreError(format!(
"settlement loss delta overflow for account {:#034x}",
event.event.account_id.0
)))
})?;
}
if accepted_event.overage {
entry.overage =
entry.overage.checked_add(accepted_event.units).ok_or_else(|| {
IngestError::Refused(StoreError(format!(
"overage delta overflow for account {:#034x}",
event.event.account_id.0
)))
})?;
}
}
let inserted = sqlx::query(
"INSERT INTO tollgate_usage_events
(request_id, account_id, lease_id, fencing_token, units, occurred_at_us, policy_revision, key_id)
SELECT * FROM UNNEST($1::bytea[], $2::bytea[], $3::bytea[], $4::bigint[], $5::bigint[], $6::bigint[], $7::bytea[], $8::bytea[])",
)
.bind(&rid)
.bind(&acct)
.bind(&lease)
.bind(&fence)
.bind(&units)
.bind(&at)
.bind(&revision)
.bind(&keys)
.execute(&mut *tx)
.await
.map_err(storage)?;
let expected_event_rows = u64::try_from(accepted.len())
.map_err(|_| StoreError("accepted event count exceeds u64 range".into()))?;
if inserted.rows_affected() != expected_event_rows {
return Err(IngestError::Unavailable(StoreError(format!(
"ingest inserted {} of {} accepted usage rows",
inserted.rows_affected(),
accepted.len()
))));
}
let (lease_update_ids, lease_used_deltas): (Vec<Vec<u8>>, Vec<i64>) = leases
.iter()
.filter_map(|(lease_id, row)| {
row.used_delta
.map(|used_delta| (lease_id.clone(), used_delta.get()))
})
.unzip();
if !lease_update_ids.is_empty() {
let updated_leases = sqlx::query(
"UPDATE tollgate_leases AS lease
SET used = lease.used + delta.used
FROM UNNEST($1::bytea[], $2::bigint[]) AS delta(lease_id, used)
WHERE lease.lease_id = delta.lease_id",
)
.bind(&lease_update_ids)
.bind(&lease_used_deltas)
.execute(&mut *tx)
.await
.map_err(storage)?;
let expected_lease_rows = u64::try_from(lease_update_ids.len())
.map_err(|_| StoreError("ingest lease row count exceeds u64 range".into()))?;
if updated_leases.rows_affected() != expected_lease_rows {
return Err(IngestError::Unavailable(StoreError(format!(
"ingest updated {} of {} locked lease rows",
updated_leases.rows_affected(),
lease_update_ids.len()
))));
}
}
let mut account_ids = Vec::with_capacity(account_deltas.len());
let mut account_usage_deltas = Vec::with_capacity(account_deltas.len());
let mut account_loss_deltas = Vec::with_capacity(account_deltas.len());
let mut account_overage_deltas = Vec::with_capacity(account_deltas.len());
for (account_id, delta) in &account_deltas {
account_ids.push(account_id.clone());
account_usage_deltas.push(delta.usage);
account_loss_deltas.push(delta.loss);
account_overage_deltas.push(delta.overage);
}
let locked_accounts = sqlx::query(
"SELECT account_id, usage_recorded, settlement_loss, overage_recorded
FROM tollgate_accounts
WHERE account_id = ANY($1)
ORDER BY account_id FOR UPDATE",
)
.bind(&account_ids)
.fetch_all(&mut *tx)
.await
.map_err(storage)?;
if locked_accounts.len() != account_ids.len() {
return Err(IngestError::Unavailable(StoreError(format!(
"ingest locked {} of {} referenced account rows",
locked_accounts.len(),
account_ids.len()
))));
}
for row in locked_accounts {
let account_id: Vec<u8> = row.get(0);
let usage_recorded: i64 = row.get(1);
let settlement_loss: i64 = row.get(2);
let overage_recorded: i64 = row.get(3);
to_units(usage_recorded, "account usage_recorded")?;
to_units(settlement_loss, "account settlement_loss")?;
to_units(overage_recorded, "account overage_recorded")?;
let delta = account_deltas.get(&account_id).ok_or_else(|| {
StoreError(format!(
"ingest locked unexpected account {:#034x}",
id_from(&account_id)
))
})?;
usage_recorded.checked_add(delta.usage).ok_or_else(|| {
IngestError::Refused(StoreError(format!(
"usage_recorded overflow for account {:#034x}",
id_from(&account_id)
)))
})?;
overage_recorded.checked_add(delta.overage).ok_or_else(|| {
IngestError::Refused(StoreError(format!(
"overage_recorded overflow for account {:#034x}",
id_from(&account_id)
)))
})?;
if settlement_loss < delta.loss {
return Err(IngestError::Unavailable(StoreError(format!(
"settlement_loss underflow for account {:#034x}: settled straggler \
usage {} exceeds recorded loss",
id_from(&account_id),
delta.loss
))));
}
}
let updated_accounts = sqlx::query(
"UPDATE tollgate_accounts AS account
SET usage_recorded = account.usage_recorded + delta.usage,
settlement_loss = account.settlement_loss - delta.loss,
overage_recorded = account.overage_recorded + delta.overage
FROM UNNEST($1::bytea[], $2::bigint[], $3::bigint[], $4::bigint[])
AS delta(account_id, usage, loss, overage)
WHERE account.account_id = delta.account_id
AND account.settlement_loss >= delta.loss",
)
.bind(&account_ids)
.bind(&account_usage_deltas)
.bind(&account_loss_deltas)
.bind(&account_overage_deltas)
.execute(&mut *tx)
.await
.map_err(storage)?;
let expected_account_rows = u64::try_from(account_ids.len())
.map_err(|_| StoreError("ingest account row count exceeds u64 range".into()))?;
if updated_accounts.rows_affected() != expected_account_rows {
return Err(IngestError::Unavailable(StoreError(format!(
"ingest updated {} of {} locked account rows",
updated_accounts.rows_affected(),
account_ids.len()
))));
}
if keys.iter().all(Option::is_none) {
report.unattributed = Some(expected_event_rows);
return Ok(report);
}
let attributed: i64 = sqlx::query_scalar(
"WITH matched AS MATERIALIZED (
SELECT k.key_id, b.occurred_at_us
FROM UNNEST($1::bytea[], $2::bytea[], $3::bigint[])
AS b(key_id, account_id, occurred_at_us)
JOIN LATERAL (
SELECT key_id FROM tollgate_credential_keys
WHERE key_id = b.key_id AND account_id = b.account_id LIMIT 1
) k ON true
), updated AS (
INSERT INTO tollgate_credential_activity AS activity (key_id, last_committed_at_us)
SELECT key_id, MAX(occurred_at_us) FROM matched GROUP BY key_id ORDER BY key_id
ON CONFLICT (key_id) DO UPDATE
SET last_committed_at_us = EXCLUDED.last_committed_at_us
WHERE activity.last_committed_at_us < EXCLUDED.last_committed_at_us
RETURNING key_id
) SELECT COUNT(*) FROM matched"
).bind(&keys).bind(&acct).bind(&at).fetch_one(&mut *tx).await.map_err(storage)?;
report.unattributed = Some(expected_event_rows.checked_sub(
u64::try_from(attributed).map_err(|_| StoreError("negative attribution count".into()))?
).ok_or_else(|| StoreError("attribution count exceeds accepted events".into()))?);
Ok(report)
}
.await;
finish_transaction(tx, result).await
}
}
fn decode_status(stored: String) -> Result<AccountStatus, StoreError> {
match stored.as_str() {
s if s == AccountStatus::Active.as_str() => Ok(AccountStatus::Active),
s if s == AccountStatus::Suspended.as_str() => Ok(AccountStatus::Suspended),
s if s == AccountStatus::Closed.as_str() => Ok(AccountStatus::Closed),
other => Err(StoreError(format!("unrecognized account status {other:?}"))),
}
}
fn decode_capacity_class(stored: String) -> Result<CapacityClass, StoreError> {
match stored.as_str() {
s if s == CapacityClass::Assured.as_str() => Ok(CapacityClass::Assured),
s if s == CapacityClass::BestEffort.as_str() => Ok(CapacityClass::BestEffort),
other => Err(StoreError(format!("unrecognized capacity class {other:?}"))),
}
}
fn decode_authority(stored: String) -> Result<AdminAuthority, StoreError> {
match stored.as_str() {
s if s == AdminAuthority::Operator.as_str() => Ok(AdminAuthority::Operator),
s if s == AdminAuthority::Provisioner.as_str() => Ok(AdminAuthority::Provisioner),
other => Err(StoreError(format!(
"unrecognized admin authority {other:?}"
))),
}
}
fn decode_schedule(
allowance: Option<i64>,
period: Option<String>,
rollover: Option<String>,
) -> Result<Option<BudgetSchedule>, StoreError> {
let populated = [allowance.is_some(), period.is_some(), rollover.is_some()];
let (Some(allowance), Some(period), Some(rollover)) = (allowance, period, rollover) else {
if populated.iter().any(|present| *present) {
return Err(StoreError(
"stored budget schedule is partially populated".into(),
));
}
return Ok(None);
};
let period = match period.as_str() {
s if s == Period::UtcCalendarMonth.as_str() => Period::UtcCalendarMonth,
other => return Err(StoreError(format!("unrecognized budget period {other:?}"))),
};
let rollover = match rollover.as_str() {
s if s == Rollover::None.as_str() => Rollover::None,
other => {
return Err(StoreError(format!(
"unrecognized budget rollover {other:?}"
)));
}
};
Ok(Some(BudgetSchedule {
allowance: to_units(allowance, "budget allowance")?,
period,
rollover,
}))
}
const BUDGET_VIEW_SQL: &str = "SELECT account_id, deposited, overage_recorded, usage_recorded,
settlement_loss, expired, budget_allowance, budget_period,
budget_rollover, period_start_us
FROM tollgate_accounts WHERE account_id = $1";
const GRANT_DEBIT_SQL: &str = "UPDATE tollgate_accounts
SET balance = balance - $2,
allowance_balance = allowance_balance - $3,
next_fence = next_fence + 1
WHERE account_id = $1
RETURNING account_id, deposited, overage_recorded, usage_recorded,
settlement_loss, expired, budget_allowance, budget_period,
budget_rollover, period_start_us";
fn budget_view(row: &sqlx::postgres::PgRow) -> Result<BudgetView, StoreError> {
let deposited = to_units(row.get::<i64, _>(1), "deposited")?;
let overage = to_units(row.get::<i64, _>(2), "overage_recorded")?;
let usage = to_units(row.get::<i64, _>(3), "usage_recorded")?;
let loss = to_units(row.get::<i64, _>(4), "settlement_loss")?;
let expired = to_units(row.get::<i64, _>(5), "expired")?;
let schedule = decode_schedule(row.get(6), row.get(7), row.get(8))?;
let period_start = micros_ts(row.get::<i64, _>(9), "period_start_us")?;
let funded = deposited
.checked_add(overage)
.ok_or_else(|| StoreError("account funding total overflows".into()))?;
let consumed = usage
.checked_add(loss)
.and_then(|spent| spent.checked_add(expired))
.ok_or_else(|| StoreError("account consumption total overflows".into()))?;
Ok(BudgetView {
balance_at_publish: funded.checked_sub(consumed).ok_or_else(|| {
StoreError(format!(
"consumption {} exceeds funding {}",
consumed.get(),
funded.get()
))
})?,
period_end: schedule.map(|schedule| schedule.period.end_after(period_start)),
})
}
fn decode_publishable(
principal: Principal,
generation: i64,
value: serde_json::Value,
) -> Result<PublishableSnapshot, StoreError> {
let generation = generation_from(generation)?;
let snapshot: StoredSnapshot =
serde_json::from_value(value).map_err(|e| StoreError(format!("snapshot decode: {e}")))?;
let budget = snapshot.budget;
let publishable = PublishableSnapshot::try_new(Arc::new(snapshot.into_snapshot(generation)))
.map_err(|error| {
StoreError(format!(
"invalid stored snapshot for principal {:#034x}: {error}",
principal.0
))
})?;
Ok(match budget {
Some(budget) => publishable.with_budget(Some(budget)),
None => publishable,
})
}
fn generation_from(column: i64) -> Result<Generation, StoreError> {
u64::try_from(column)
.map(Generation)
.map_err(|_| StoreError("stored snapshot generation is negative".into()))
}
#[async_trait]
impl SnapshotSource for PostgresStore {
async fn snapshot(&self, principal: Principal) -> Result<SnapshotResolution, StoreError> {
let row = sqlx::query(
"SELECT generation, snapshot, deleted FROM tollgate_snapshots WHERE principal = $1",
)
.bind(id_bytes(principal.0))
.fetch_optional(&self.pool)
.await
.map_err(storage)?;
match row {
Some(row) if row.get::<bool, _>(2) => Ok(SnapshotResolution::Revoked {
generation: generation_from(row.get::<i64, _>(0))?,
}),
Some(row) => Ok(SnapshotResolution::Present(decode_publishable(
principal,
row.get::<i64, _>(0),
row.get(1),
)?)),
None => Ok(SnapshotResolution::Unknown),
}
}
fn subscribe(&self) -> broadcast::Receiver<SnapshotPush> {
self.push.subscribe()
}
async fn principals(&self) -> Result<Option<Vec<Principal>>, StoreError> {
let rows = sqlx::query("SELECT principal FROM tollgate_snapshots ORDER BY principal")
.fetch_all(&self.pool)
.await
.map_err(storage)?;
Ok(Some(
rows.iter()
.map(|row| Principal(id_from(row.get::<Vec<u8>, _>(0).as_slice())))
.collect(),
))
}
}
#[async_trait]
impl StoreHealth for PostgresStore {
async fn ping(&self) -> Result<(), StoreError> {
sqlx::query("SELECT 1")
.execute(&self.pool)
.await
.map_err(storage)?;
Ok(())
}
}
async fn republish_patched_snapshots(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
account: AccountId,
json_path: &'static str,
value: &str,
) -> Result<(Vec<(Principal, PublishableSnapshot)>, usize), SetStatusError> {
let rows = sqlx::query(
"UPDATE tollgate_snapshots
SET generation = generation + 1,
snapshot = jsonb_set(snapshot, $3::text[], to_jsonb($2::text))
WHERE account_id = $1
AND deleted = FALSE
AND snapshot #>> $3::text[] IS DISTINCT FROM $2::text
RETURNING principal, generation, snapshot",
)
.bind(id_bytes(account.0))
.bind(value)
.bind(json_path)
.fetch_all(&mut **tx)
.await
.map_err(storage)?;
let mut republished = Vec::with_capacity(rows.len());
let mut unreadable = 0usize;
for row in rows {
let principal = Principal(id_from(row.get::<Vec<u8>, _>(0).as_slice()));
match decode_publishable(
principal,
row.get::<i64, _>(1),
row.get::<serde_json::Value, _>(2),
) {
Ok(snapshot) => republished.push((principal, snapshot)),
Err(error) => {
unreadable += 1;
tracing::warn!(
%principal,
%error,
"restamped snapshot could not be decoded for push"
);
}
}
}
republished.sort_unstable_by_key(|(principal, _)| *principal);
Ok((republished, unreadable))
}
#[async_trait]
impl AdminStore for PostgresStore {
async fn create_account(
&self,
config: AccountConfig,
) -> Result<tollgate_store::AdminReceipt<()>, CreateAccountError> {
self.create_account_as(config, AdminAuthority::Operator)
.await
}
async fn create_provisioned_account(
&self,
account: AccountId,
) -> Result<tollgate_store::AdminReceipt<()>, CreateAccountError> {
self.create_account_as(
AccountConfig {
account_id: account,
initial_balance: CostUnits::ZERO,
status: AccountStatus::Suspended,
capacity_class: CapacityClass::BestEffort,
},
AdminAuthority::Provisioner,
)
.await
}
async fn deposit(
&self,
account: AccountId,
units: CostUnits,
) -> Result<AdminReceipt<()>, AllocateError> {
let row = sqlx::query(
"UPDATE tollgate_accounts
SET balance = balance + $2, deposited = deposited + $2
WHERE account_id = $1
RETURNING balance - $2 AS old_topup, deposited - $2 AS old_deposited,
balance AS new_topup, deposited AS new_deposited",
)
.bind(id_bytes(account.0))
.bind(i64::try_from(units.get()).map_err(|_| AllocateError::BalanceOverflow)?)
.fetch_optional(&self.pool)
.await
.map_err(|error| {
if error
.as_database_error()
.and_then(|db| db.code())
.as_deref()
== Some("22003")
{
AllocateError::BalanceOverflow
} else {
alloc_storage(error)
}
})?
.ok_or(AllocateError::UnknownAccount)?;
let state = |topup: &str, deposited: &str| -> Result<AdminState, StoreError> {
Ok(AdminState::Funding {
topup: to_units(row.get(topup), "audit topup")?,
deposited: to_units(row.get(deposited), "audit deposited")?,
})
};
Ok(AdminReceipt::new(
(),
state("old_topup", "old_deposited")?,
state("new_topup", "new_deposited")?,
))
}
async fn set_budget_schedule(
&self,
account: AccountId,
schedule: Option<BudgetSchedule>,
) -> Result<AdminReceipt<()>, BudgetError> {
let allowance = schedule
.map(|s| to_i64(s.allowance, "budget allowance"))
.transpose()
.map_err(BudgetError::Storage)?;
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let row = sqlx::query(
"SELECT budget_allowance, budget_period, budget_rollover
FROM tollgate_accounts WHERE account_id = $1 FOR UPDATE",
)
.bind(id_bytes(account.0))
.fetch_optional(&mut *tx)
.await
.map_err(storage)?
.ok_or(BudgetError::UnknownAccount)?;
let before = decode_schedule(row.get(0), row.get(1), row.get(2))?;
sqlx::query(
"UPDATE tollgate_accounts
SET budget_allowance = $2, budget_period = $3, budget_rollover = $4
WHERE account_id = $1",
)
.bind(id_bytes(account.0))
.bind(allowance)
.bind(schedule.map(|s| s.period.as_str()))
.bind(schedule.map(|s| s.rollover.as_str()))
.execute(&mut *tx)
.await
.map_err(storage)?;
Ok(AdminReceipt::new(
(),
AdminState::Budget { schedule: before },
AdminState::Budget { schedule },
))
}
.await;
finish_transaction(tx, result).await
}
async fn account_view(&self, account: AccountId) -> Result<Option<AccountView>, StoreError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
sqlx::query("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY")
.execute(&mut *tx)
.await
.map_err(storage)?;
let account_row = sqlx::query(
"SELECT deposited, balance, usage_recorded, settlement_loss, overage_recorded,
expired, status, capacity_class, budget_allowance, budget_period,
budget_rollover, period_start_us, origin, status_set_by
FROM tollgate_accounts WHERE account_id = $1",
)
.bind(id_bytes(account.0))
.fetch_optional(&mut *tx)
.await
.map_err(storage)?;
let lease_row = sqlx::query(ACTIVE_LEASE_SUM_SQL)
.bind(id_bytes(account.0))
.fetch_one(&mut *tx)
.await
.map_err(storage)?;
Ok::<_, StoreError>((account_row, lease_row))
}
.await;
let (account_row, lease_row) = finish_transaction(tx, result).await?;
let Some(row) = account_row else {
return Ok(None);
};
let active_grants = to_units(lease_row.get::<i64, _>(0), "active lease grants")?;
let active_used = to_units(lease_row.get::<i64, _>(1), "active lease usage")?;
let recorded = to_units(row.get::<i64, _>(2), "usage_recorded")?;
Ok(Some(AccountView {
account_id: account,
status: decode_status(row.get::<String, _>(6))?,
capacity_class: decode_capacity_class(row.get::<String, _>(7))?,
origin: decode_authority(row.get::<String, _>(12))?,
status_set_by: decode_authority(row.get::<String, _>(13))?,
schedule: decode_schedule(
row.get::<Option<i64>, _>(8),
row.get::<Option<String>, _>(9),
row.get::<Option<String>, _>(10),
)?,
period_start: StoredInstant {
micros: row.get::<i64, _>(11),
submicro_nanos: 0,
}
.timestamp()?,
conservation: Conservation {
deposited: to_units(row.get::<i64, _>(0), "deposited")?,
overage_recorded: to_units(row.get::<i64, _>(4), "overage_recorded")?,
balance: to_units(row.get::<i64, _>(1), "balance")?,
active_lease_grants: active_grants,
settled_usage: recorded.checked_sub(active_used).ok_or_else(|| {
StoreError(format!(
"active lease usage {} exceeds recorded usage {} for account {account}",
active_used.get(),
recorded.get()
))
})?,
settlement_loss: to_units(row.get::<i64, _>(3), "settlement_loss")?,
expired: to_units(row.get::<i64, _>(5), "expired")?,
},
}))
}
async fn roll_due_periods(
&self,
now: Timestamp,
limit: NonZeroUsize,
) -> Result<RolloverBatch, StoreError> {
let limit_i = i64::try_from(limit.get())
.map_err(|_| StoreError(format!("rollover batch limit exceeds i64 range: {limit}")))?;
let mut rolled = Vec::new();
for period in Period::ALL {
let boundary_us = ts_micros(period.start_of(now));
let remaining = limit_i
- i64::try_from(rolled.len())
.map_err(|_| StoreError("rollover batch row count exceeds i64 range".into()))?;
if remaining <= 0 {
break;
}
let rows = sqlx::query(DUE_PERIODS_SQL)
.bind(period.as_str())
.bind(boundary_us)
.bind(remaining)
.fetch_all(&self.pool)
.await
.map_err(storage)?;
for row in rows {
rolled.push((
row.get::<i64, _>(3),
RolledAccount {
account_id: AccountId(id_from(&row.get::<Vec<u8>, _>(0))),
deposited: to_units(row.get::<i64, _>(1), "budget allowance")?,
expired: to_units(row.get::<i64, _>(2), "expiring allowance")?,
},
));
}
}
rolled
.sort_unstable_by_key(|(crossed_from, account)| (*crossed_from, account.account_id.0));
RolloverBatch::try_new(
rolled.into_iter().map(|(_, account)| account).collect(),
limit,
)
}
async fn set_account_status(
&self,
account: AccountId,
status: AccountStatus,
) -> Result<tollgate_store::AdminReceipt<StatusChange>, SetStatusError> {
self.set_status_as(account, status, AdminAuthority::Operator)
.await
}
async fn activate_provisioned(
&self,
account: AccountId,
) -> Result<tollgate_store::AdminReceipt<StatusChange>, SetStatusError> {
self.set_status_as(account, AccountStatus::Active, AdminAuthority::Provisioner)
.await
}
async fn set_capacity_class(
&self,
account: AccountId,
class: CapacityClass,
) -> Result<tollgate_store::AdminReceipt<StatusChange>, SetStatusError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = async {
let row = sqlx::query(
"SELECT status, capacity_class FROM tollgate_accounts WHERE account_id = $1 FOR UPDATE",
)
.bind(id_bytes(account.0))
.fetch_optional(&mut *tx)
.await
.map_err(storage)?
.ok_or(SetStatusError::UnknownAccount)?;
let before = decode_capacity_class(row.get::<String, _>(1))?;
if decode_status(row.get::<String, _>(0))? == AccountStatus::Closed {
return Err(SetStatusError::AccountClosed);
}
sqlx::query("UPDATE tollgate_accounts SET capacity_class = $2 WHERE account_id = $1")
.bind(id_bytes(account.0))
.bind(class.as_str())
.execute(&mut *tx)
.await
.map_err(storage)?;
let (republished, unreadable) =
republish_patched_snapshots(&mut tx, account, "{capacity_class}", class.as_str()).await?;
Ok((before, republished, unreadable))
}
.await;
let (before, republished, unreadable) = finish_transaction(tx, result).await?;
if pushes_exceed_capacity(republished.len()) {
tracing::warn!(
%account,
principals = republished.len(),
capacity = PUSH_CHANNEL_CAPACITY,
"capacity class change emitted more pushes than the channel holds; \
subscribers will resync"
);
}
let count = republished.len();
for (principal, snapshot) in republished {
self.push_to_subscribers(SnapshotPush {
principal,
resolution: SnapshotResolution::Present(snapshot),
});
}
Ok(AdminReceipt::new(
StatusChange {
republished: count + unreadable,
unreadable,
},
AdminState::CapacityClass {
capacity_class: before,
},
AdminState::CapacityClass {
capacity_class: class,
},
))
}
async fn publish_snapshot(
&self,
principal: Principal,
snapshot: PublishableSnapshot,
) -> Result<tollgate_store::AdminReceipt<()>, PublishSnapshotError> {
let generation = i64::try_from(snapshot.generation.0).map_err(|_| {
StoreError("snapshot generation exceeds PostgreSQL BIGINT range".into())
})?;
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = publish_in_tx(&mut tx, principal, Some(generation), snapshot).await;
let (written, published, before, after) = finish_transaction(tx, result).await?;
if written {
self.push_to_subscribers(SnapshotPush {
principal,
resolution: SnapshotResolution::Present(published),
});
}
Ok(AdminReceipt::new((), before, after))
}
async fn remove_snapshot(&self, principal: Principal) -> Result<AdminReceipt<()>, StoreError> {
let mut tx = self.pool.begin().await.map_err(storage)?;
let result = remove_in_tx(&mut tx, principal).await;
let receipt = finish_transaction(tx, result).await?;
self.announce_removal(principal, &receipt);
Ok(receipt)
}
}
async fn lock_account_key(
tx: &mut Transaction<'_, Postgres>,
account: AccountId,
key: KeyId,
) -> Result<(Principal, bool), KeySnapshotError> {
let row = sqlx::query(
"SELECT principal, revoked_at_us IS NOT NULL
FROM tollgate_credential_keys WHERE key_id = $1 AND account_id = $2 FOR SHARE",
)
.bind(id_bytes(key.0))
.bind(id_bytes(account.0))
.fetch_optional(&mut **tx)
.await
.map_err(storage)?
.ok_or(KeySnapshotError::UnknownCredential)?;
let principal: Vec<u8> = row.get(0);
let principal: [u8; 16] = principal
.try_into()
.map_err(|_| StoreError("credential principal is not 16 bytes".into()))?;
Ok((Principal(u128::from_be_bytes(principal)), row.get(1)))
}
async fn publish_in_tx(
tx: &mut Transaction<'_, Postgres>,
principal: Principal,
generation: Option<i64>,
snapshot: PublishableSnapshot,
) -> Result<(bool, PublishableSnapshot, AdminState, AdminState), PublishSnapshotError> {
if let Some(key_id) = snapshot.key_id {
let matches: bool = sqlx::query_scalar(
"SELECT EXISTS(SELECT 1 FROM tollgate_credential_keys
WHERE key_id = $1 AND principal = $2 AND account_id = $3)",
)
.bind(id_bytes(key_id.0))
.bind(id_bytes(principal.0))
.bind(id_bytes(snapshot.account_id.0))
.fetch_one(&mut **tx)
.await
.map_err(storage)?;
if !matches {
return Err(PublishSnapshotError::CredentialMismatch { key_id });
}
}
let ledger = sqlx::query(
"SELECT status, deposited, overage_recorded, usage_recorded, settlement_loss,
expired, budget_allowance, budget_period, budget_rollover, period_start_us,
capacity_class
FROM tollgate_accounts WHERE account_id = $1 FOR SHARE",
)
.bind(id_bytes(snapshot.account_id.0))
.fetch_optional(&mut **tx)
.await
.map_err(storage)?;
let view = match &ledger {
Some(row) => {
let ledger = decode_status(row.get::<String, _>(0))?;
if ledger != snapshot.status {
return Err(PublishSnapshotError::StatusMismatch {
ledger,
submitted: snapshot.status,
});
}
let ledger_class = decode_capacity_class(row.get::<String, _>(10))?;
if ledger_class != snapshot.capacity_class {
return Err(PublishSnapshotError::CapacityClassMismatch {
ledger: ledger_class,
submitted: snapshot.capacity_class,
});
}
Some(budget_view(row)?)
}
None => None,
};
let published = snapshot.with_budget(view);
let value = serde_json::to_value(StoredSnapshotRef::from(published.as_snapshot()))
.map_err(|e| StoreError(format!("snapshot encode: {e}")))?;
let (written, before, after) = write_snapshot_audited(tx, principal, generation, value).await?;
let published = if generation.is_none() {
let AdminState::Snapshot { generation, .. } = after else {
return Err(StoreError("published snapshot has no generation".into()).into());
};
published.restamped(published.status, generation)
} else {
published
};
Ok((written, published, before, after))
}
async fn remove_in_tx(
tx: &mut Transaction<'_, Postgres>,
principal: Principal,
) -> Result<AdminReceipt<()>, StoreError> {
let before = snapshot_audit_row(tx, principal).await?;
let after = match before {
AdminState::Snapshot {
generation,
revoked: false,
} => {
sqlx::query("UPDATE tollgate_snapshots SET deleted = TRUE WHERE principal = $1")
.bind(id_bytes(principal.0))
.execute(&mut **tx)
.await
.map_err(storage)?;
AdminState::Snapshot {
generation,
revoked: true,
}
}
state => state,
};
Ok(AdminReceipt::new((), before, after))
}
async fn snapshot_audit_row(
tx: &mut Transaction<'_, Postgres>,
principal: Principal,
) -> Result<AdminState, StoreError> {
let row = sqlx::query(
"SELECT generation, deleted FROM tollgate_snapshots WHERE principal = $1 FOR UPDATE",
)
.bind(id_bytes(principal.0))
.fetch_optional(&mut **tx)
.await
.map_err(storage)?;
row.map(|row| {
Ok(AdminState::Snapshot {
generation: generation_from(row.get(0))?,
revoked: row.get(1),
})
})
.unwrap_or(Ok(AdminState::Absent))
}
async fn write_snapshot_audited(
tx: &mut Transaction<'_, Postgres>,
principal: Principal,
generation: Option<i64>,
value: serde_json::Value,
) -> Result<(bool, AdminState, AdminState), StoreError> {
let mut before = snapshot_audit_row(tx, principal).await?;
if before == AdminState::Absent {
let initial = generation.unwrap_or(1);
let after = AdminState::Snapshot {
generation: generation_from(initial)?,
revoked: false,
};
let inserted = sqlx::query(
"INSERT INTO tollgate_snapshots (principal, generation, snapshot, deleted)
VALUES ($1, $2, $3, FALSE) ON CONFLICT (principal) DO NOTHING",
)
.bind(id_bytes(principal.0))
.bind(initial)
.bind(&value)
.execute(&mut **tx)
.await
.map_err(storage)?;
if inserted.rows_affected() == 1 {
return Ok((true, before, after));
}
before = snapshot_audit_row(tx, principal).await?;
}
let AdminState::Snapshot {
generation: previous,
..
} = before
else {
return Err(StoreError("snapshot disappeared during publication".into()));
};
let generation = match generation {
Some(stated) => stated,
None => i64::try_from(previous.0)
.ok()
.and_then(|previous| previous.checked_add(1))
.ok_or_else(|| StoreError("snapshot generation overflow".into()))?,
};
let after = AdminState::Snapshot {
generation: generation_from(generation)?,
revoked: false,
};
if previous >= generation_from(generation)? {
return Ok((false, before, before));
}
sqlx::query("UPDATE tollgate_snapshots SET generation = $2, snapshot = $3, deleted = FALSE WHERE principal = $1")
.bind(id_bytes(principal.0)).bind(generation).bind(value)
.execute(&mut **tx).await.map_err(storage)?;
Ok((true, before, after))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stored_authorities_reject_unknown_vocabulary() {
assert_eq!(
decode_authority("Operator".into()).unwrap(),
AdminAuthority::Operator
);
assert_eq!(
decode_authority("Provisioner".into()).unwrap(),
AdminAuthority::Provisioner
);
for invalid in [
"",
"operator",
"provisioner",
"Administrator",
"Provisioner ",
] {
assert!(decode_authority(invalid.into()).is_err(), "{invalid:?}");
}
}
#[test]
fn stored_fences_use_the_exact_positive_bigint_domain() {
for invalid in [i64::MIN, -1, 0] {
assert!(stored_fence(invalid).is_err());
}
for valid in [1, 2, i64::MAX] {
assert_eq!(stored_fence(valid).unwrap(), FencingToken(valid as u64));
}
}
#[tokio::test]
async fn ping_surfaces_a_closed_pool() {
let pool = PgPoolOptions::new()
.connect_lazy("postgres://localhost/tollgate")
.unwrap();
pool.close().await;
let (push, _) = broadcast::channel(1);
let store = PostgresStore {
pool,
policy: GrantPolicy::default(),
push,
};
assert!(store.ping().await.is_err());
assert!(matches!(
AdminStore::deposit(&store, AccountId(1), CostUnits(1)).await,
Err(AllocateError::Storage(_))
));
}
}