feat(server): impl storage runtime (#15181)

#### PR Dependency Tree


* **PR #15181** 👈

This tree was auto-generated by
[Charcoal](https://github.com/danerwilliams/charcoal)

<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added an additional storage backend option: asset-pack based storage
(provider for avatar, blob, and copilot).
* Introduced a dedicated storage runtime with provider capability
reporting and expanded object operations (put/head/get/list/delete),
including presigned and multipart flows where supported.
* Cloudflare R2 `jurisdiction` now uses an explicit default when
omitted.
* **Bug Fixes**
  * Broadened avatar access to allow both fs and asset-pack providers.
* Improved workspace blob upload completion validation and handling when
stored objects are missing or mismatched.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
This commit is contained in:
DarkSky
2026-07-01 22:24:10 +08:00
committed by GitHub
parent da7d438377
commit 8ebdb7452f
102 changed files with 6487 additions and 4508 deletions
@@ -0,0 +1,9 @@
pub(super) const BYOK_LOCAL_LEASE_ACTIVE_PURPOSE: &str = "copilot_byok_local_lease:active";
pub(super) const BYOK_LOCAL_LEASE_PURPOSE: &str = "copilot_byok_local_lease";
pub(super) const MAGIC_LINK_OTP_PURPOSE: &str = "magic_link_otp";
pub(super) const MAX_MAGIC_LINK_OTP_ATTEMPTS: i32 = 10;
pub(super) const WORKSPACE_INVITE_LINK_ID_PURPOSE: &str = "workspace_invite_link:id";
pub(super) const WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE: &str = "workspace_invite_link:workspace";
pub(super) const WORKSPACE_STATS_LEASE_KEY: &str = "workspace:admin-stats:refresh";
pub(super) const WORKSPACE_STATS_LOCK_NAMESPACE: i64 = 97_301;
pub(super) const WORKSPACE_STATS_REFRESH_LOCK_KEY: i64 = 1;
@@ -0,0 +1,163 @@
use napi::Result;
use sqlx::{FromRow, PgPool};
use super::{BackendRuntime, RuntimeError, RuntimeResult, napi_error, types::CoordinationLeaseGrant};
#[derive(FromRow)]
struct LeaseGrantRow {
fencing_token: i64,
}
struct CoordinationLeaseStore {
pool: PgPool,
}
impl CoordinationLeaseStore {
fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn acquire(&self, key: String, owner: String, ttl_ms: i64) -> RuntimeResult<Option<CoordinationLeaseGrant>> {
let row = sqlx::query_as::<_, LeaseGrantRow>(
r#"
INSERT INTO runtime_leases (key, owner, fencing_token, expires_at)
VALUES ($1, $2, 1, CURRENT_TIMESTAMP + ($3 * INTERVAL '1 millisecond'))
ON CONFLICT (key) DO UPDATE
SET owner = EXCLUDED.owner,
fencing_token = runtime_leases.fencing_token + 1,
expires_at = EXCLUDED.expires_at,
updated_at = CURRENT_TIMESTAMP
WHERE runtime_leases.expires_at <= CURRENT_TIMESTAMP
RETURNING fencing_token
"#,
)
.bind(&key)
.bind(&owner)
.bind(ttl_ms as f64)
.fetch_optional(&self.pool)
.await
.map_err(|err| RuntimeError::database("CoordinationLease acquire failed", err))?;
Ok(row.map(|row| CoordinationLeaseGrant {
key,
owner,
fencing_token: row.fencing_token,
}))
}
async fn release(&self, key: &str, owner: &str, fencing_token: i64) -> RuntimeResult<bool> {
let result = sqlx::query(
r#"
DELETE FROM runtime_leases
WHERE key = $1 AND owner = $2 AND fencing_token = $3
"#,
)
.bind(key)
.bind(owner)
.bind(fencing_token)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("CoordinationLease release failed", err))?;
Ok(result.rows_affected() == 1)
}
async fn renew(&self, key: &str, owner: &str, fencing_token: i64, ttl_ms: i64) -> RuntimeResult<bool> {
let result = sqlx::query(
r#"
UPDATE runtime_leases
SET expires_at = CURRENT_TIMESTAMP + ($4 * INTERVAL '1 millisecond'),
updated_at = CURRENT_TIMESTAMP
WHERE key = $1
AND owner = $2
AND fencing_token = $3
AND expires_at > CURRENT_TIMESTAMP
"#,
)
.bind(key)
.bind(owner)
.bind(fencing_token)
.bind(ttl_ms as f64)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("CoordinationLease renew failed", err))?;
Ok(result.rows_affected() == 1)
}
}
#[napi_derive::napi]
impl BackendRuntime {
pub(crate) async fn acquire_coordination_lease_inner(
&self,
key: String,
owner: String,
ttl_ms: i64,
) -> RuntimeResult<Option<CoordinationLeaseGrant>> {
if ttl_ms <= 0 {
return Err(RuntimeError::invalid_input("coordination lease ttl must be positive"));
}
if owner.is_empty() {
return Err(RuntimeError::invalid_input("coordination lease owner is required"));
}
CoordinationLeaseStore::new(self.pool().await?)
.acquire(key, owner, ttl_ms)
.await
}
pub(crate) async fn release_coordination_lease_inner(
&self,
key: String,
owner: String,
fencing_token: i64,
) -> RuntimeResult<bool> {
CoordinationLeaseStore::new(self.pool().await?)
.release(&key, &owner, fencing_token)
.await
}
#[napi]
pub async fn acquire_coordination_lease(
&self,
key: String,
owner: String,
ttl_ms: i64,
) -> Result<Option<CoordinationLeaseGrant>> {
self
.acquire_coordination_lease_inner(key, owner, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn release_coordination_lease(
&self,
key: String,
owner: String,
#[napi(ts_arg_type = "bigint | number")] fencing_token: i64,
) -> Result<bool> {
self
.release_coordination_lease_inner(key, owner, fencing_token)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn renew_coordination_lease(
&self,
key: String,
owner: String,
#[napi(ts_arg_type = "bigint | number")] fencing_token: i64,
ttl_ms: i64,
) -> Result<bool> {
if ttl_ms <= 0 {
return Err(napi_error("coordination lease ttl must be positive"));
}
CoordinationLeaseStore::new(self.pool().await?)
.renew(&key, &owner, fencing_token, ttl_ms)
.await
.map_err(napi::Error::from)
}
}
@@ -0,0 +1,410 @@
use chrono::{DateTime, Duration, Utc};
use sqlx::{FromRow, PgPool, Postgres, Row, Transaction};
use y_octo::Doc;
use super::{BackendRuntime, RuntimeError, RuntimeResult, napi_error, types::RuntimeDocCompactionResult};
#[derive(FromRow)]
struct SnapshotRow {
blob: Vec<u8>,
updated_at: DateTime<Utc>,
updated_by: Option<String>,
}
#[derive(FromRow)]
struct UpdateRow {
blob: Vec<u8>,
created_at: DateTime<Utc>,
created_by: Option<String>,
}
struct DocCompactorStore {
pool: PgPool,
}
impl DocCompactorStore {
fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn compact_doc(
&self,
workspace_id: &str,
doc_id: &str,
batch_limit: i64,
history_min_interval_ms: i64,
history_max_age_seconds: i64,
) -> RuntimeResult<(i64, bool)> {
compact_doc(
self.pool.clone(),
workspace_id,
doc_id,
batch_limit,
history_min_interval_ms,
history_max_age_seconds,
)
.await
}
}
fn is_empty_doc(bin: &[u8]) -> bool {
bin.is_empty() || (bin.len() == 1 && bin[0] == 0) || (bin.len() == 2 && bin[0] == 0 && bin[1] == 0)
}
fn apply_updates(updates: impl IntoIterator<Item = Vec<u8>>) -> RuntimeResult<Vec<u8>> {
let mut doc = Doc::default();
for update in updates {
doc
.apply_update_from_binary_v1(&update)
.map_err(|err| RuntimeError::invalid_state(format!("DocCompactor merge failed: {err}")))?;
}
doc
.encode_update_v1()
.map_err(|err| RuntimeError::invalid_state(format!("DocCompactor encode failed: {err}")))
}
fn checked_milliseconds(value: i64, field: &str) -> RuntimeResult<Duration> {
Duration::try_milliseconds(value)
.ok_or_else(|| RuntimeError::invalid_input(format!("DocCompactor {field} is too large")))
}
fn checked_seconds(value: i64, field: &str) -> RuntimeResult<Duration> {
Duration::try_seconds(value).ok_or_else(|| RuntimeError::invalid_input(format!("DocCompactor {field} is too large")))
}
async fn load_snapshot(
tx: &mut Transaction<'_, Postgres>,
workspace_id: &str,
doc_id: &str,
) -> RuntimeResult<Option<SnapshotRow>> {
sqlx::query_as::<_, SnapshotRow>(
r#"
SELECT blob, updated_at, updated_by
FROM snapshots
WHERE workspace_id = $1 AND guid = $2
FOR UPDATE
"#,
)
.bind(workspace_id)
.bind(doc_id)
.fetch_optional(&mut **tx)
.await
.map_err(|err| RuntimeError::database("DocCompactor load snapshot failed", err))
}
async fn load_updates(
tx: &mut Transaction<'_, Postgres>,
workspace_id: &str,
doc_id: &str,
batch_limit: i64,
) -> RuntimeResult<Vec<UpdateRow>> {
sqlx::query_as::<_, UpdateRow>(
r#"
SELECT blob, created_at, created_by
FROM updates
WHERE workspace_id = $1 AND guid = $2
ORDER BY created_at ASC
LIMIT $3
FOR UPDATE
"#,
)
.bind(workspace_id)
.bind(doc_id)
.bind(batch_limit)
.fetch_all(&mut **tx)
.await
.map_err(|err| RuntimeError::database("DocCompactor load updates failed", err))
}
async fn upsert_snapshot(
tx: &mut Transaction<'_, Postgres>,
workspace_id: &str,
doc_id: &str,
blob: &[u8],
timestamp: DateTime<Utc>,
editor: Option<&str>,
) -> RuntimeResult<bool> {
if is_empty_doc(blob) {
return Ok(false);
}
let row = sqlx::query(
r#"
INSERT INTO snapshots
(workspace_id, guid, blob, size, created_at, updated_at, created_by, updated_by)
VALUES
($1, $2, $3, $4, $5, $5, $6, $6)
ON CONFLICT (workspace_id, guid)
DO UPDATE SET
blob = $3,
size = $4,
updated_at = $5,
updated_by = $6
WHERE snapshots.workspace_id = $1
AND snapshots.guid = $2
AND snapshots.updated_at <= $5
RETURNING updated_at
"#,
)
.bind(workspace_id)
.bind(doc_id)
.bind(blob)
.bind(blob.len() as i64)
.bind(timestamp)
.bind(editor)
.fetch_optional(&mut **tx)
.await
.map_err(|err| RuntimeError::database("DocCompactor upsert snapshot failed", err))?;
Ok(row.is_some())
}
async fn should_create_history(
tx: &mut Transaction<'_, Postgres>,
snapshot: &SnapshotRow,
workspace_id: &str,
doc_id: &str,
history_min_interval_ms: i64,
) -> RuntimeResult<bool> {
if is_empty_doc(&snapshot.blob) {
return Ok(false);
}
let row = sqlx::query(
r#"
SELECT timestamp
FROM snapshot_histories
WHERE workspace_id = $1 AND guid = $2
ORDER BY timestamp DESC
LIMIT 1
"#,
)
.bind(workspace_id)
.bind(doc_id)
.fetch_optional(&mut **tx)
.await
.map_err(|err| RuntimeError::database("DocCompactor load latest history failed", err))?;
let Some(row) = row else {
return Ok(true);
};
let last_timestamp: DateTime<Utc> = row.get("timestamp");
if last_timestamp == snapshot.updated_at {
return Ok(false);
}
let min_interval = checked_milliseconds(history_min_interval_ms, "history interval")?;
let threshold = snapshot
.updated_at
.checked_sub_signed(min_interval)
.ok_or_else(|| RuntimeError::invalid_input("DocCompactor history interval is out of range"))?;
Ok(last_timestamp < threshold)
}
async fn create_history(
tx: &mut Transaction<'_, Postgres>,
workspace_id: &str,
doc_id: &str,
snapshot: &SnapshotRow,
max_age_seconds: i64,
) -> RuntimeResult<bool> {
if max_age_seconds <= 0 {
return Ok(false);
}
let max_age = checked_seconds(max_age_seconds, "history max age")?;
let expired_at = Utc::now()
.checked_add_signed(max_age)
.ok_or_else(|| RuntimeError::invalid_input("DocCompactor history max age is out of range"))?;
sqlx::query(
r#"
INSERT INTO snapshot_histories
(workspace_id, guid, timestamp, blob, expired_at, created_by)
VALUES
($1, $2, $3, $4, $5, $6)
ON CONFLICT (workspace_id, guid, timestamp)
DO UPDATE SET expired_at = EXCLUDED.expired_at
"#,
)
.bind(workspace_id)
.bind(doc_id)
.bind(snapshot.updated_at)
.bind(&snapshot.blob)
.bind(expired_at)
.bind(snapshot.updated_by.as_deref())
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database("DocCompactor create history failed", err))?;
Ok(true)
}
async fn delete_updates(
tx: &mut Transaction<'_, Postgres>,
workspace_id: &str,
doc_id: &str,
timestamps: &[DateTime<Utc>],
) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM updates
WHERE workspace_id = $1
AND guid = $2
AND created_at = ANY($3)
"#,
)
.bind(workspace_id)
.bind(doc_id)
.bind(timestamps)
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database("DocCompactor delete updates failed", err))?;
Ok(result.rows_affected() as i64)
}
async fn compact_doc(
pool: PgPool,
workspace_id: &str,
doc_id: &str,
batch_limit: i64,
history_min_interval_ms: i64,
history_max_age_seconds: i64,
) -> RuntimeResult<(i64, bool)> {
let mut tx = pool
.begin()
.await
.map_err(|err| RuntimeError::database("DocCompactor begin transaction failed", err))?;
let snapshot = load_snapshot(&mut tx, workspace_id, doc_id).await?;
let updates = load_updates(&mut tx, workspace_id, doc_id, batch_limit).await?;
if updates.is_empty() {
tx.commit()
.await
.map_err(|err| RuntimeError::database("DocCompactor commit transaction failed", err))?;
return Ok((0, false));
}
let last = updates.last().expect("updates is not empty");
let mut merge_inputs = Vec::with_capacity(updates.len() + usize::from(snapshot.is_some()));
if let Some(snapshot) = &snapshot {
merge_inputs.push(snapshot.blob.clone());
}
merge_inputs.extend(updates.iter().map(|update| update.blob.clone()));
let final_blob = if merge_inputs.len() == 1 {
merge_inputs.remove(0)
} else {
apply_updates(merge_inputs)?
};
let snapshot_updated = upsert_snapshot(
&mut tx,
workspace_id,
doc_id,
&final_blob,
last.created_at,
last.created_by.as_deref(),
)
.await?;
let mut history_created = false;
if snapshot_updated
&& let Some(snapshot) = &snapshot
&& should_create_history(&mut tx, snapshot, workspace_id, doc_id, history_min_interval_ms).await?
{
history_created = create_history(&mut tx, workspace_id, doc_id, snapshot, history_max_age_seconds).await?;
}
let timestamps = updates.iter().map(|update| update.created_at).collect::<Vec<_>>();
let deleted = delete_updates(&mut tx, workspace_id, doc_id, &timestamps).await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("DocCompactor commit transaction failed", err))?;
Ok((deleted, history_created))
}
#[napi_derive::napi]
impl BackendRuntime {
/// Merge pending doc updates with y-octo and persist the merged snapshot.
///
/// Do not use this for snapshots that will be sent back to yjs clients until
/// the y-octo/yjs round-trip compatibility issue is resolved.
///
/// The caller owns quota reconciliation and must pass a fresh
/// history_max_age_seconds value. The compactor intentionally does not read
/// effective_workspace_quota_states; if a future caller cannot provide a
/// fresh quota state, fail and retry after Node reconciles it.
#[napi]
#[allow(clippy::too_many_arguments)]
pub async fn compact_pending_doc_updates(
&self,
workspace_id: String,
doc_id: String,
batch_limit: i64,
history_min_interval_ms: i64,
history_max_age_seconds: i64,
owner: String,
lease_ttl_ms: i64,
) -> napi::Result<RuntimeDocCompactionResult> {
if batch_limit <= 0 {
return Err(napi_error("doc compactor batch limit must be positive"));
}
if history_min_interval_ms < 0 {
return Err(napi_error("doc compactor history interval must be non-negative"));
}
if history_max_age_seconds < 0 {
return Err(napi_error("doc compactor history max age must be non-negative"));
}
checked_milliseconds(history_min_interval_ms, "history interval")?;
if history_max_age_seconds > 0 {
let max_age = checked_seconds(history_max_age_seconds, "history max age")?;
Utc::now()
.checked_add_signed(max_age)
.ok_or_else(|| RuntimeError::invalid_input("DocCompactor history max age is out of range"))?;
}
let lease_key = format!("doc:update:{workspace_id}:{doc_id}");
let Some(lease) = self.acquire_coordination_lease(lease_key, owner, lease_ttl_ms).await? else {
return Ok(RuntimeDocCompactionResult {
lease_acquired: false,
merged: false,
workspace_id,
doc_id,
updates_merged: 0,
history_created: false,
});
};
let result = DocCompactorStore::new(self.pool().await?)
.compact_doc(
&workspace_id,
&doc_id,
batch_limit,
history_min_interval_ms,
history_max_age_seconds,
)
.await;
let released = self
.release_coordination_lease(lease.key, lease.owner, lease.fencing_token)
.await?;
if !released {
return Err(RuntimeError::invalid_state("DocCompactor failed to release coordination lease").into());
}
let (updates_merged, history_created) = result?;
Ok(RuntimeDocCompactionResult {
lease_acquired: true,
merged: updates_merged > 0,
workspace_id,
doc_id,
updates_merged,
history_created,
})
}
}
@@ -0,0 +1,162 @@
use chrono::{DateTime, Duration, Utc};
use napi::bindgen_prelude::Buffer;
use sqlx::{PgPool, Row};
use super::{BackendRuntime, RuntimeError, RuntimeResult, napi_error, types::RuntimeDocHistoryInput};
fn is_empty_doc(bin: &[u8]) -> bool {
bin.is_empty() || (bin.len() == 1 && bin[0] == 0) || (bin.len() == 2 && bin[0] == 0 && bin[1] == 0)
}
async fn latest_history_timestamp(
pool: &PgPool,
workspace_id: &str,
doc_id: &str,
) -> RuntimeResult<Option<DateTime<Utc>>> {
sqlx::query(
r#"
SELECT timestamp
FROM snapshot_histories
WHERE workspace_id = $1 AND guid = $2
ORDER BY timestamp DESC
LIMIT 1
"#,
)
.bind(workspace_id)
.bind(doc_id)
.fetch_optional(pool)
.await
.map(|row| row.map(|row| row.get("timestamp")))
.map_err(|err| RuntimeError::database("DocStorage load latest history failed", err))
}
#[napi_derive::napi]
impl BackendRuntime {
#[napi]
pub async fn upsert_doc_snapshot(
&self,
workspace_id: String,
doc_id: String,
blob: Buffer,
timestamp_ms: i64,
editor_id: Option<String>,
) -> napi::Result<bool> {
if is_empty_doc(blob.as_ref()) {
return Ok(false);
}
let timestamp = DateTime::<Utc>::from_timestamp_millis(timestamp_ms)
.ok_or_else(|| RuntimeError::invalid_input(format!("Invalid doc snapshot timestamp: {timestamp_ms}")))?;
let pool = self.pool().await?;
let row = sqlx::query(
r#"
INSERT INTO snapshots
(workspace_id, guid, blob, size, created_at, updated_at, created_by, updated_by)
VALUES
($1, $2, $3, $4, $5, $5, $6, $6)
ON CONFLICT (workspace_id, guid)
DO UPDATE SET
blob = $3,
size = $4,
updated_at = $5,
updated_by = $6
WHERE snapshots.workspace_id = $1
AND snapshots.guid = $2
AND snapshots.updated_at <= $5
RETURNING updated_at
"#,
)
.bind(&workspace_id)
.bind(&doc_id)
.bind(blob.as_ref())
.bind(blob.len() as i64)
.bind(timestamp)
.bind(editor_id.as_deref())
.fetch_optional(&pool)
.await
.map_err(|err| RuntimeError::database("DocStorage upsert snapshot failed", err))?;
Ok(row.is_some())
}
#[napi]
pub async fn create_doc_history(&self, input: RuntimeDocHistoryInput) -> napi::Result<bool> {
if input.history_min_interval_ms < 0 {
return Err(napi_error("doc history interval must be non-negative"));
}
if input.history_max_age_ms <= 0 || is_empty_doc(input.blob.as_ref()) {
return Ok(false);
}
let timestamp = DateTime::<Utc>::from_timestamp_millis(input.timestamp_ms)
.ok_or_else(|| RuntimeError::invalid_input(format!("Invalid doc history timestamp: {}", input.timestamp_ms)))?;
let pool = self.pool().await?;
let should_create = match latest_history_timestamp(&pool, &input.workspace_id, &input.doc_id).await? {
None => true,
Some(last_timestamp) if last_timestamp == timestamp => false,
Some(last_timestamp) => {
input.force || last_timestamp < timestamp - Duration::milliseconds(input.history_min_interval_ms)
}
};
if !should_create {
return Ok(false);
}
let expired_at = Utc::now() + Duration::milliseconds(input.history_max_age_ms);
sqlx::query(
r#"
INSERT INTO snapshot_histories
(workspace_id, guid, timestamp, blob, expired_at, created_by)
VALUES
($1, $2, $3, $4, $5, $6)
ON CONFLICT (workspace_id, guid, timestamp)
DO UPDATE SET expired_at = EXCLUDED.expired_at
"#,
)
.bind(&input.workspace_id)
.bind(&input.doc_id)
.bind(timestamp)
.bind(input.blob.as_ref())
.bind(expired_at)
.bind(input.editor_id.as_deref())
.execute(&pool)
.await
.map_err(|err| RuntimeError::database("DocStorage create history failed", err))?;
Ok(true)
}
#[napi]
pub async fn delete_doc_storage(&self, workspace_id: String, doc_id: String) -> napi::Result<()> {
let pool = self.pool().await?;
let mut tx = pool
.begin()
.await
.map_err(|err| RuntimeError::database("DocStorage delete begin transaction failed", err))?;
sqlx::query("DELETE FROM snapshots WHERE workspace_id = $1 AND guid = $2")
.bind(&workspace_id)
.bind(&doc_id)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("DocStorage delete snapshot failed", err))?;
sqlx::query("DELETE FROM updates WHERE workspace_id = $1 AND guid = $2")
.bind(&workspace_id)
.bind(&doc_id)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("DocStorage delete updates failed", err))?;
sqlx::query("DELETE FROM snapshot_histories WHERE workspace_id = $1 AND guid = $2")
.bind(&workspace_id)
.bind(&doc_id)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("DocStorage delete histories failed", err))?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("DocStorage delete commit failed", err))?;
Ok(())
}
}
@@ -0,0 +1,94 @@
use napi::Result;
use sqlx::PgPool;
use super::{BackendRuntime, RuntimeError, RuntimeResult, napi_error};
struct RuntimeGateStore {
pool: PgPool,
}
impl RuntimeGateStore {
fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn put_if_absent(&self, key: &str, ttl_ms: i64) -> RuntimeResult<bool> {
let mut tx = self
.pool
.begin()
.await
.map_err(|err| RuntimeError::database("RuntimeGate transaction failed", err))?;
sqlx::query("DELETE FROM runtime_gates WHERE key = $1 AND expires_at <= CURRENT_TIMESTAMP")
.bind(key)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("RuntimeGate expired cleanup failed", err))?;
let inserted = sqlx::query(
r#"
INSERT INTO runtime_gates (key, expires_at)
VALUES ($1, CURRENT_TIMESTAMP + ($2 * INTERVAL '1 millisecond'))
ON CONFLICT (key) DO NOTHING
"#,
)
.bind(key)
.bind(ttl_ms as f64)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("RuntimeGate put_if_absent failed", err))?
.rows_affected()
== 1;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeGate transaction commit failed", err))?;
Ok(inserted)
}
async fn cleanup_expired(&self, limit: i64) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM runtime_gates
WHERE key IN (
SELECT key FROM runtime_gates
WHERE expires_at <= CURRENT_TIMESTAMP
ORDER BY expires_at ASC
LIMIT $1
)
"#,
)
.bind(limit)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("RuntimeGate cleanup failed", err))?;
Ok(result.rows_affected() as i64)
}
}
#[napi_derive::napi]
impl BackendRuntime {
#[napi]
pub async fn put_runtime_gate_if_absent(&self, key: String, ttl_ms: i64) -> Result<bool> {
if ttl_ms <= 0 {
return Err(napi_error("runtime gate ttl must be positive"));
}
RuntimeGateStore::new(self.pool().await?)
.put_if_absent(&key, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn cleanup_expired_runtime_gates(&self, limit: i64) -> Result<i64> {
if limit <= 0 {
return Err(napi_error("runtime gate cleanup limit must be positive"));
}
RuntimeGateStore::new(self.pool().await?)
.cleanup_expired(limit)
.await
.map_err(napi::Error::from)
}
}
@@ -0,0 +1,82 @@
use napi::Result;
use sqlx::PgPool;
use super::{BackendRuntime, RuntimeError, RuntimeResult, napi_error};
struct HousekeepingStore {
pool: PgPool,
}
impl HousekeepingStore {
fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn cleanup_expired_user_sessions(&self, limit: i64) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM user_sessions
WHERE id IN (
SELECT id FROM user_sessions
WHERE expires_at <= CURRENT_TIMESTAMP
ORDER BY expires_at ASC
LIMIT $1
)
"#,
)
.bind(limit)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("Housekeeping user sessions cleanup failed", err))?;
Ok(result.rows_affected() as i64)
}
async fn cleanup_expired_snapshot_histories(&self, limit: i64) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM snapshot_histories
WHERE (workspace_id, guid, timestamp) IN (
SELECT workspace_id, guid, timestamp
FROM snapshot_histories
WHERE expired_at <= CURRENT_TIMESTAMP
ORDER BY expired_at ASC
LIMIT $1
)
"#,
)
.bind(limit)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("Housekeeping snapshot histories cleanup failed", err))?;
Ok(result.rows_affected() as i64)
}
}
#[napi_derive::napi]
impl BackendRuntime {
#[napi]
pub async fn cleanup_expired_user_sessions(&self, limit: i64) -> Result<i64> {
if limit <= 0 {
return Err(napi_error("user sessions cleanup limit must be positive"));
}
HousekeepingStore::new(self.pool().await?)
.cleanup_expired_user_sessions(limit)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn cleanup_expired_snapshot_histories(&self, limit: i64) -> Result<i64> {
if limit <= 0 {
return Err(napi_error("snapshot histories cleanup limit must be positive"));
}
HousekeepingStore::new(self.pool().await?)
.cleanup_expired_snapshot_histories(limit)
.await
.map_err(napi::Error::from)
}
}
@@ -0,0 +1,133 @@
mod constants;
mod coordination_lease;
mod doc_compactor;
mod doc_storage;
mod gate;
mod housekeeping;
mod runtime_state;
#[cfg(test)]
mod tests;
mod workspace_stats;
use std::{sync::RwLock, time::Duration};
use napi::Result;
use sha2::{Digest, Sha256};
use sqlx::{PgPool, Row, postgres::PgPoolOptions};
use tokio::sync::Mutex;
use self::types::BackendRuntimeHealth;
pub(crate) use super::types;
use super::{
BackendRuntimeConfig, RuntimeError, RuntimeResult, migrations::migrate_runtime_tables, napi_error, to_napi_error,
};
pub(super) fn token_hash(token: &str) -> String {
hex::encode(Sha256::digest(token.as_bytes()))
}
#[napi_derive::napi]
pub struct BackendRuntime {
config: RwLock<BackendRuntimeConfig>,
pool: Mutex<Option<PgPool>>,
}
#[napi_derive::napi]
impl BackendRuntime {
#[napi(constructor)]
pub fn new() -> Result<Self> {
Ok(Self {
config: RwLock::new(BackendRuntimeConfig::from_config_files().map_err(to_napi_error)?),
pool: Mutex::new(None),
})
}
#[napi]
pub async fn start(&self) -> Result<()> {
self.start_inner().await.map_err(to_napi_error)
}
async fn start_inner(&self) -> RuntimeResult<()> {
let mut guard = self.pool.lock().await;
if guard.is_some() {
return Ok(());
}
let database_url = self.config()?.database_url;
let pool = PgPoolOptions::new()
.max_connections(5)
.acquire_timeout(Duration::from_secs(5))
.connect(&database_url)
.await
.map_err(|err| RuntimeError::database("BackendRuntime failed to connect postgres", err))?;
sqlx::query("SELECT 1")
.execute(&pool)
.await
.map_err(|err| RuntimeError::database("BackendRuntime postgres health check failed", err))?;
let config = self.config()?.with_db_overrides(&pool).await?;
self.update_config(config)?;
*guard = Some(pool);
Ok(())
}
#[napi]
pub async fn stop(&self) -> Result<()> {
let pool = self.pool.lock().await.take();
if let Some(pool) = pool {
pool.close().await;
}
Ok(())
}
#[napi]
pub async fn health(&self) -> Result<BackendRuntimeHealth> {
let pool = self.pool.lock().await.as_ref().cloned();
let database_connected = match pool.as_ref() {
Some(pool) => sqlx::query("SELECT 1")
.fetch_one(pool)
.await
.map(|row| row.try_get::<i32, _>(0).unwrap_or(0) == 1)
.unwrap_or(false),
None => false,
};
Ok(BackendRuntimeHealth {
started: pool.is_some(),
database_connected,
})
}
#[napi]
pub async fn run_migrations(&self) -> Result<()> {
let pool = self.pool().await?;
migrate_runtime_tables(&pool).await.map_err(to_napi_error)
}
pub(crate) async fn pool(&self) -> RuntimeResult<PgPool> {
self
.pool
.lock()
.await
.as_ref()
.cloned()
.ok_or_else(|| RuntimeError::invalid_state("BackendRuntime must be started before using postgres operations"))
}
pub(crate) fn config(&self) -> RuntimeResult<BackendRuntimeConfig> {
self
.config
.read()
.map(|config| config.clone())
.map_err(|_| RuntimeError::invalid_state("BackendRuntime config lock poisoned"))
}
fn update_config(&self, config: BackendRuntimeConfig) -> RuntimeResult<()> {
*self
.config
.write()
.map_err(|_| RuntimeError::invalid_state("BackendRuntime config lock poisoned"))? = config;
Ok(())
}
}
@@ -0,0 +1,40 @@
use super::{Result, auth_challenge_purpose, dto::RuntimeStateRows};
pub(super) async fn create(
rows: &RuntimeStateRows,
purpose: &str,
token: &str,
payload: serde_json::Value,
ttl_ms: i64,
) -> Result<bool> {
rows
.insert_payload_if_absent(
&auth_challenge_purpose(purpose),
token,
None,
payload,
ttl_ms,
"RuntimeState auth challenge create",
)
.await
}
pub(super) async fn get(rows: &RuntimeStateRows, purpose: &str, token: &str) -> Result<Option<serde_json::Value>> {
rows
.active_payload(
&auth_challenge_purpose(purpose),
token,
"RuntimeState auth challenge get",
)
.await
}
pub(super) async fn consume(rows: &RuntimeStateRows, purpose: &str, token: &str) -> Result<Option<serde_json::Value>> {
rows
.consume_payload(
&auth_challenge_purpose(purpose),
token,
"RuntimeState auth challenge consume",
)
.await
}
@@ -0,0 +1,128 @@
use super::{
BYOK_LOCAL_LEASE_ACTIVE_PURPOSE, BYOK_LOCAL_LEASE_PURPOSE, Result, RuntimeByokLocalLeaseRecord, RuntimeError,
dto::{RuntimeStateInsertPayload, RuntimeStatePayloadRow, RuntimeStateRows},
};
pub(super) async fn get(rows: &RuntimeStateRows, lease_id: String) -> Result<Option<RuntimeByokLocalLeaseRecord>> {
get_lease_by_id(rows, &lease_id).await
}
pub(super) async fn create(
rows: &RuntimeStateRows,
active_key: String,
lease_id: String,
payload: serde_json::Value,
ttl_ms: i64,
) -> Result<RuntimeByokLocalLeaseRecord> {
if ttl_ms <= 0 {
return Err(RuntimeError::invalid_input("BYOK local lease ttl must be positive"));
}
let mut tx = rows.begin("RuntimeState BYOK local lease").await?;
sqlx::query("SELECT pg_advisory_xact_lock(hashtextextended($1, 0))")
.bind(&active_key)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("RuntimeState BYOK local lease active lock failed", err))?;
if let Some(active) = rows
.active_payload_with_expires_for_update_in_tx(
&mut tx,
BYOK_LOCAL_LEASE_ACTIVE_PURPOSE,
&active_key,
"RuntimeState BYOK local lease active get",
)
.await?
{
let existing_lease = match active.payload.get("leaseId").and_then(serde_json::Value::as_str) {
Some(existing_lease_id) => get_lease_by_id_in_tx(rows, &mut tx, existing_lease_id).await?,
None => None,
};
if let Some(lease) = existing_lease {
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState BYOK local lease transaction commit failed", err))?;
return Ok(lease);
}
rows
.delete_by_key_in_tx(
&mut tx,
BYOK_LOCAL_LEASE_ACTIVE_PURPOSE,
&active_key,
"RuntimeState BYOK local lease stale active delete",
)
.await?;
}
let expires_at_ms = rows
.insert_payload_returning_expires_in_tx(
&mut tx,
RuntimeStateInsertPayload {
purpose: BYOK_LOCAL_LEASE_PURPOSE,
token: &lease_id,
lookup_key: &active_key,
payload: &payload,
ttl_ms,
context: "RuntimeState BYOK local lease create",
},
)
.await?;
let active_payload = serde_json::json!({ "leaseId": lease_id });
rows
.insert_payload_returning_expires_in_tx(
&mut tx,
RuntimeStateInsertPayload {
purpose: BYOK_LOCAL_LEASE_ACTIVE_PURPOSE,
token: &active_key,
lookup_key: &active_key,
payload: &active_payload,
ttl_ms,
context: "RuntimeState BYOK local lease active create",
},
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState BYOK local lease transaction commit failed", err))?;
Ok(RuntimeByokLocalLeaseRecord {
lease_id,
payload,
expires_at_ms,
})
}
async fn get_lease_by_id(rows: &RuntimeStateRows, lease_id: &str) -> Result<Option<RuntimeByokLocalLeaseRecord>> {
rows
.active_payload_with_expires(BYOK_LOCAL_LEASE_PURPOSE, lease_id, "RuntimeState BYOK local lease get")
.await?
.map(|row| record_from_row(lease_id, row))
.transpose()
}
async fn get_lease_by_id_in_tx(
rows: &RuntimeStateRows,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
lease_id: &str,
) -> Result<Option<RuntimeByokLocalLeaseRecord>> {
rows
.active_payload_with_expires_for_update_in_tx(
tx,
BYOK_LOCAL_LEASE_PURPOSE,
lease_id,
"RuntimeState BYOK local lease get",
)
.await?
.map(|row| record_from_row(lease_id, row))
.transpose()
}
fn record_from_row(lease_id: &str, row: RuntimeStatePayloadRow) -> Result<RuntimeByokLocalLeaseRecord> {
Ok(RuntimeByokLocalLeaseRecord {
lease_id: lease_id.to_string(),
payload: row.payload,
expires_at_ms: row.expires_at_ms,
})
}
@@ -0,0 +1,455 @@
use sqlx::{PgPool, Row};
use super::{RuntimeError, RuntimeResult, token_hash};
type Result<T> = RuntimeResult<T>;
pub(super) struct RuntimeStatePayloadRow {
pub(super) payload: serde_json::Value,
pub(super) expires_at_ms: i64,
}
pub(super) struct RuntimeStateLockedRow {
pub(super) payload: serde_json::Value,
pub(super) attempts: i32,
pub(super) expires_at: chrono::DateTime<chrono::Utc>,
}
pub(super) struct RuntimeStateInsertPayload<'a> {
pub(super) purpose: &'a str,
pub(super) token: &'a str,
pub(super) lookup_key: &'a str,
pub(super) payload: &'a serde_json::Value,
pub(super) ttl_ms: i64,
pub(super) context: &'a str,
}
#[derive(Clone)]
pub(super) struct RuntimeStateRows {
pub(super) pool: PgPool,
}
impl RuntimeStateRows {
pub(super) fn new(pool: PgPool) -> Self {
Self { pool }
}
pub(super) fn pool(&self) -> &PgPool {
&self.pool
}
pub(super) async fn begin(&self, context: &str) -> Result<sqlx::Transaction<'_, sqlx::Postgres>> {
self
.pool
.begin()
.await
.map_err(|err| RuntimeError::database(format!("{context} transaction failed"), err))
}
pub(super) async fn insert_payload(
&self,
purpose: &str,
token: &str,
lookup_key: Option<&str>,
payload: serde_json::Value,
ttl_ms: i64,
context: &str,
) -> Result<()> {
sqlx::query(
r#"
INSERT INTO runtime_states (purpose, token_hash, lookup_key, payload, expires_at)
VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP + ($5 * INTERVAL '1 millisecond'))
"#,
)
.bind(purpose)
.bind(token_hash(token))
.bind(lookup_key)
.bind(payload)
.bind(ttl_ms as f64)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(())
}
pub(super) async fn insert_payload_if_absent(
&self,
purpose: &str,
token: &str,
lookup_key: Option<&str>,
payload: serde_json::Value,
ttl_ms: i64,
context: &str,
) -> Result<bool> {
let inserted = sqlx::query(
r#"
INSERT INTO runtime_states (purpose, token_hash, lookup_key, payload, expires_at)
VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP + ($5 * INTERVAL '1 millisecond'))
ON CONFLICT (purpose, token_hash) DO NOTHING
"#,
)
.bind(purpose)
.bind(token_hash(token))
.bind(lookup_key)
.bind(payload)
.bind(ttl_ms as f64)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?
.rows_affected()
== 1;
Ok(inserted)
}
pub(super) async fn upsert_payload_reset_attempts(
&self,
purpose: &str,
token: &str,
lookup_key: &str,
payload: serde_json::Value,
ttl_ms: i64,
context: &str,
) -> Result<()> {
sqlx::query(
r#"
INSERT INTO runtime_states (purpose, token_hash, lookup_key, payload, attempts, consumed_at, expires_at)
VALUES ($1, $2, $3, $4, 0, NULL, CURRENT_TIMESTAMP + ($5 * INTERVAL '1 millisecond'))
ON CONFLICT (purpose, token_hash) DO UPDATE
SET lookup_key = EXCLUDED.lookup_key,
payload = EXCLUDED.payload,
attempts = 0,
consumed_at = NULL,
expires_at = EXCLUDED.expires_at,
updated_at = CURRENT_TIMESTAMP
"#,
)
.bind(purpose)
.bind(token_hash(token))
.bind(lookup_key)
.bind(payload)
.bind(ttl_ms as f64)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(())
}
pub(super) async fn active_payload(
&self,
purpose: &str,
token: &str,
context: &str,
) -> Result<Option<serde_json::Value>> {
let row = sqlx::query(
r#"
SELECT payload
FROM runtime_states
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > CURRENT_TIMESTAMP
"#,
)
.bind(purpose)
.bind(token_hash(token))
.fetch_optional(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(row.map(|row| row.get::<serde_json::Value, _>("payload")))
}
pub(super) async fn active_payload_with_expires(
&self,
purpose: &str,
token: &str,
context: &str,
) -> Result<Option<RuntimeStatePayloadRow>> {
let row = sqlx::query(
r#"
SELECT payload, (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
FROM runtime_states
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > CURRENT_TIMESTAMP
"#,
)
.bind(purpose)
.bind(token_hash(token))
.fetch_optional(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(row.map(payload_row))
}
pub(super) async fn consume_payload(
&self,
purpose: &str,
token: &str,
context: &str,
) -> Result<Option<serde_json::Value>> {
let row = sqlx::query(
r#"
UPDATE runtime_states
SET consumed_at = CURRENT_TIMESTAMP,
updated_at = CURRENT_TIMESTAMP
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > CURRENT_TIMESTAMP
RETURNING payload
"#,
)
.bind(purpose)
.bind(token_hash(token))
.fetch_optional(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(row.map(|row| row.get::<serde_json::Value, _>("payload")))
}
pub(super) async fn consume_payload_with_expires(
&self,
purpose: &str,
token: &str,
context: &str,
) -> Result<Option<RuntimeStatePayloadRow>> {
let row = sqlx::query(
r#"
UPDATE runtime_states
SET consumed_at = CURRENT_TIMESTAMP,
updated_at = CURRENT_TIMESTAMP
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > CURRENT_TIMESTAMP
RETURNING payload, (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
"#,
)
.bind(purpose)
.bind(token_hash(token))
.fetch_optional(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(row.map(payload_row))
}
pub(super) async fn active_payload_with_expires_for_update_in_tx(
&self,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
purpose: &str,
token: &str,
context: &str,
) -> Result<Option<RuntimeStatePayloadRow>> {
let row = sqlx::query(
r#"
SELECT payload, (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
FROM runtime_states
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > clock_timestamp()
FOR UPDATE
"#,
)
.bind(purpose)
.bind(token_hash(token))
.fetch_optional(&mut **tx)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(row.map(payload_row))
}
pub(super) async fn unconsumed_row_for_update_in_tx(
&self,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
purpose: &str,
token: &str,
context: &str,
) -> Result<Option<RuntimeStateLockedRow>> {
let row = sqlx::query(
r#"
SELECT payload, attempts, expires_at
FROM runtime_states
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
FOR UPDATE
"#,
)
.bind(purpose)
.bind(token_hash(token))
.fetch_optional(&mut **tx)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(row.map(|row| RuntimeStateLockedRow {
payload: row.get("payload"),
attempts: row.get("attempts"),
expires_at: row.get("expires_at"),
}))
}
pub(super) async fn insert_payload_returning_expires_in_tx(
&self,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
input: RuntimeStateInsertPayload<'_>,
) -> Result<i64> {
let row = sqlx::query(
r#"
INSERT INTO runtime_states (purpose, token_hash, lookup_key, payload, expires_at)
VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP + ($5 * INTERVAL '1 millisecond'))
RETURNING (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
"#,
)
.bind(input.purpose)
.bind(token_hash(input.token))
.bind(input.lookup_key)
.bind(input.payload)
.bind(input.ttl_ms as f64)
.fetch_one(&mut **tx)
.await
.map_err(|err| RuntimeError::database(input.context, err))?;
Ok(row.get::<i64, _>("expires_at_ms"))
}
pub(super) async fn upsert_expired_or_consumed_payload_returning_expires_in_tx(
&self,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
input: RuntimeStateInsertPayload<'_>,
) -> Result<Option<i64>> {
let row = sqlx::query(
r#"
INSERT INTO runtime_states (purpose, token_hash, lookup_key, payload, expires_at)
VALUES ($1, $2, $3, $4, clock_timestamp() + ($5 * INTERVAL '1 millisecond'))
ON CONFLICT (purpose, token_hash) DO UPDATE
SET lookup_key = EXCLUDED.lookup_key,
payload = EXCLUDED.payload,
attempts = 0,
consumed_at = NULL,
expires_at = clock_timestamp() + ($5 * INTERVAL '1 millisecond')
WHERE runtime_states.consumed_at IS NOT NULL
OR runtime_states.expires_at <= clock_timestamp()
RETURNING (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
"#,
)
.bind(input.purpose)
.bind(token_hash(input.token))
.bind(input.lookup_key)
.bind(input.payload)
.bind(input.ttl_ms as f64)
.fetch_optional(&mut **tx)
.await
.map_err(|err| RuntimeError::database(input.context, err))?;
Ok(row.map(|row| row.get::<i64, _>("expires_at_ms")))
}
pub(super) async fn update_attempts_in_tx(
&self,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
purpose: &str,
token: &str,
attempts: i32,
context: &str,
) -> Result<()> {
sqlx::query(
r#"
UPDATE runtime_states
SET attempts = $3,
updated_at = CURRENT_TIMESTAMP
WHERE purpose = $1
AND token_hash = $2
"#,
)
.bind(purpose)
.bind(token_hash(token))
.bind(attempts)
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(())
}
pub(super) async fn delete_by_key_in_tx(
&self,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
purpose: &str,
token: &str,
context: &str,
) -> Result<()> {
sqlx::query("DELETE FROM runtime_states WHERE purpose = $1 AND token_hash = $2")
.bind(purpose)
.bind(token_hash(token))
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(())
}
pub(super) async fn cleanup_expired_or_consumed(&self, limit: i64, context: &str) -> Result<i64> {
let result = sqlx::query(
r#"
DELETE FROM runtime_states
WHERE (purpose, token_hash) IN (
SELECT purpose, token_hash FROM runtime_states
WHERE expires_at <= CURRENT_TIMESTAMP
OR consumed_at IS NOT NULL
ORDER BY expires_at ASC
LIMIT $1
)
"#,
)
.bind(limit)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(result.rows_affected() as i64)
}
pub(super) async fn cleanup_expired_by_purpose_prefix(
&self,
purpose_prefix: &str,
limit: i64,
context: &str,
) -> Result<i64> {
let result = sqlx::query(
r#"
DELETE FROM runtime_states
WHERE (purpose, token_hash) IN (
SELECT purpose, token_hash FROM runtime_states
WHERE purpose LIKE $1
AND expires_at <= CURRENT_TIMESTAMP
ORDER BY expires_at ASC
LIMIT $2
)
"#,
)
.bind(format!("{purpose_prefix}%"))
.bind(limit)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database(context, err))?;
Ok(result.rows_affected() as i64)
}
}
fn payload_row(row: sqlx::postgres::PgRow) -> RuntimeStatePayloadRow {
RuntimeStatePayloadRow {
payload: row.get("payload"),
expires_at_ms: row.get("expires_at_ms"),
}
}
@@ -0,0 +1,184 @@
use super::{
Result, RuntimeError, RuntimeWorkspaceInviteLinkRecord, WORKSPACE_INVITE_LINK_ID_PURPOSE,
WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE,
dto::{RuntimeStateInsertPayload, RuntimeStatePayloadRow, RuntimeStateRows},
};
pub(super) async fn get_by_workspace(
rows: &RuntimeStateRows,
workspace_id: String,
) -> Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
get_by_key(rows, WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE, &workspace_id).await
}
pub(super) async fn get_by_invite_id(
rows: &RuntimeStateRows,
invite_id: String,
) -> Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
get_by_key(rows, WORKSPACE_INVITE_LINK_ID_PURPOSE, &invite_id).await
}
pub(super) async fn create(
rows: &RuntimeStateRows,
workspace_id: String,
invite_id: String,
inviter_user_id: String,
ttl_ms: i64,
) -> Result<RuntimeWorkspaceInviteLinkRecord> {
if ttl_ms <= 0 {
return Err(RuntimeError::invalid_input(
"workspace invite link ttl must be positive",
));
}
let mut tx = rows.begin("RuntimeState workspace invite link").await?;
sqlx::query("SELECT pg_advisory_xact_lock(hashtextextended($1, 0))")
.bind(&workspace_id)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("RuntimeState workspace invite link active lock failed", err))?;
if let Some(existing) =
get_by_key_in_tx(rows, &mut tx, WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE, &workspace_id).await?
{
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState workspace invite link transaction commit failed", err))?;
return Ok(existing);
}
let payload = serde_json::json!({
"workspaceId": workspace_id,
"inviteId": invite_id,
"inviterUserId": inviter_user_id,
});
let Some(expires_at_ms) = rows
.upsert_expired_or_consumed_payload_returning_expires_in_tx(
&mut tx,
RuntimeStateInsertPayload {
purpose: WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE,
token: &workspace_id,
lookup_key: &workspace_id,
payload: &payload,
ttl_ms,
context: "RuntimeState workspace invite link create",
},
)
.await?
else {
let existing = get_by_key_in_tx(rows, &mut tx, WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE, &workspace_id).await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState workspace invite link transaction commit failed", err))?;
return existing
.ok_or_else(|| RuntimeError::invalid_state("RuntimeState workspace invite link active conflict missing row"));
};
rows
.insert_payload_returning_expires_in_tx(
&mut tx,
RuntimeStateInsertPayload {
purpose: WORKSPACE_INVITE_LINK_ID_PURPOSE,
token: &invite_id,
lookup_key: &invite_id,
payload: &payload,
ttl_ms,
context: "RuntimeState workspace invite link create",
},
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState workspace invite link transaction commit failed", err))?;
Ok(RuntimeWorkspaceInviteLinkRecord {
workspace_id,
invite_id,
inviter_user_id,
expires_at_ms,
})
}
pub(super) async fn revoke(rows: &RuntimeStateRows, workspace_id: String) -> Result<bool> {
let mut tx = rows.begin("RuntimeState workspace invite link").await?;
let existing = get_by_key_in_tx(rows, &mut tx, WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE, &workspace_id).await?;
let Some(existing) = existing else {
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState workspace invite link transaction commit failed", err))?;
return Ok(false);
};
rows
.delete_by_key_in_tx(
&mut tx,
WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE,
&workspace_id,
"RuntimeState workspace invite link revoke",
)
.await?;
rows
.delete_by_key_in_tx(
&mut tx,
WORKSPACE_INVITE_LINK_ID_PURPOSE,
&existing.invite_id,
"RuntimeState workspace invite link revoke",
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState workspace invite link transaction commit failed", err))?;
Ok(true)
}
async fn get_by_key(
rows: &RuntimeStateRows,
purpose: &str,
key: &str,
) -> Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
rows
.active_payload_with_expires(purpose, key, "RuntimeState workspace invite link get")
.await?
.map(record_from_row)
.transpose()
}
async fn get_by_key_in_tx(
rows: &RuntimeStateRows,
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
purpose: &str,
key: &str,
) -> Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
rows
.active_payload_with_expires_for_update_in_tx(tx, purpose, key, "RuntimeState workspace invite link get")
.await?
.map(record_from_row)
.transpose()
}
fn record_from_row(row: RuntimeStatePayloadRow) -> Result<RuntimeWorkspaceInviteLinkRecord> {
Ok(RuntimeWorkspaceInviteLinkRecord {
workspace_id: row
.payload
.get("workspaceId")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| RuntimeError::invalid_state("RuntimeState workspace invite link payload missing workspaceId"))?
.to_string(),
invite_id: row
.payload
.get("inviteId")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| RuntimeError::invalid_state("RuntimeState workspace invite link payload missing inviteId"))?
.to_string(),
inviter_user_id: row
.payload
.get("inviterUserId")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| RuntimeError::invalid_state("RuntimeState workspace invite link payload missing inviterUserId"))?
.to_string(),
expires_at_ms: row.expires_at_ms,
})
}
@@ -0,0 +1,172 @@
use super::{
MAGIC_LINK_OTP_PURPOSE, MAX_MAGIC_LINK_OTP_ATTEMPTS, Result, RuntimeError, RuntimeMagicLinkOtpConsumeResult,
dto::RuntimeStateRows,
};
impl RuntimeMagicLinkOtpConsumeResult {
fn ok(token: String) -> Self {
Self {
ok: true,
token: Some(token),
reason: None,
}
}
fn fail(reason: &'static str) -> Self {
Self {
ok: false,
token: None,
reason: Some(reason.to_string()),
}
}
}
pub(super) async fn upsert(
rows: &RuntimeStateRows,
email: String,
otp_hash: String,
token: String,
client_nonce: Option<String>,
ttl_ms: i64,
) -> Result<()> {
if ttl_ms <= 0 {
return Err(RuntimeError::invalid_input("magic link otp ttl must be positive"));
}
let payload = serde_json::json!({
"otpHash": otp_hash,
"token": token,
"clientNonce": client_nonce,
});
rows
.upsert_payload_reset_attempts(
MAGIC_LINK_OTP_PURPOSE,
&email,
&email,
payload,
ttl_ms,
"RuntimeState magic link otp upsert",
)
.await
}
pub(super) async fn consume(
rows: &RuntimeStateRows,
email: String,
otp_hash: String,
client_nonce: Option<String>,
) -> Result<RuntimeMagicLinkOtpConsumeResult> {
let mut tx = rows.begin("RuntimeState magic link otp").await?;
let row = rows
.unconsumed_row_for_update_in_tx(
&mut tx,
MAGIC_LINK_OTP_PURPOSE,
&email,
"RuntimeState magic link otp lookup",
)
.await?;
let Some(row) = row else {
tx.rollback()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction rollback failed", err))?;
return Ok(RuntimeMagicLinkOtpConsumeResult::fail("not_found"));
};
let payload = row.payload;
let attempts = row.attempts;
let expires_at = row.expires_at;
if expires_at <= chrono::Utc::now() {
rows
.delete_by_key_in_tx(
&mut tx,
MAGIC_LINK_OTP_PURPOSE,
&email,
"RuntimeState magic link otp delete",
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction commit failed", err))?;
return Ok(RuntimeMagicLinkOtpConsumeResult::fail("expired"));
}
let stored_client_nonce = payload.get("clientNonce").and_then(serde_json::Value::as_str);
if stored_client_nonce.is_some() && stored_client_nonce != client_nonce.as_deref() {
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction commit failed", err))?;
return Ok(RuntimeMagicLinkOtpConsumeResult::fail("nonce_mismatch"));
}
if attempts >= MAX_MAGIC_LINK_OTP_ATTEMPTS {
rows
.delete_by_key_in_tx(
&mut tx,
MAGIC_LINK_OTP_PURPOSE,
&email,
"RuntimeState magic link otp delete",
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction commit failed", err))?;
return Ok(RuntimeMagicLinkOtpConsumeResult::fail("locked"));
}
let stored_otp_hash = payload.get("otpHash").and_then(serde_json::Value::as_str);
if stored_otp_hash != Some(otp_hash.as_str()) {
let attempts = attempts + 1;
if attempts >= MAX_MAGIC_LINK_OTP_ATTEMPTS {
rows
.delete_by_key_in_tx(
&mut tx,
MAGIC_LINK_OTP_PURPOSE,
&email,
"RuntimeState magic link otp delete",
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction commit failed", err))?;
return Ok(RuntimeMagicLinkOtpConsumeResult::fail("locked"));
}
rows
.update_attempts_in_tx(
&mut tx,
MAGIC_LINK_OTP_PURPOSE,
&email,
attempts,
"RuntimeState magic link otp attempts update",
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction commit failed", err))?;
return Ok(RuntimeMagicLinkOtpConsumeResult::fail("invalid_otp"));
}
let token = payload
.get("token")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| RuntimeError::invalid_state("RuntimeState magic link otp payload missing token"))?
.to_string();
rows
.delete_by_key_in_tx(
&mut tx,
MAGIC_LINK_OTP_PURPOSE,
&email,
"RuntimeState magic link otp delete",
)
.await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("RuntimeState magic link otp transaction commit failed", err))?;
Ok(RuntimeMagicLinkOtpConsumeResult::ok(token))
}
@@ -0,0 +1,258 @@
use super::{BackendRuntime, RuntimeError, RuntimeResult, napi_error};
pub(super) use super::{
constants::{
BYOK_LOCAL_LEASE_ACTIVE_PURPOSE, BYOK_LOCAL_LEASE_PURPOSE, MAGIC_LINK_OTP_PURPOSE, MAX_MAGIC_LINK_OTP_ATTEMPTS,
WORKSPACE_INVITE_LINK_ID_PURPOSE, WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE,
},
token_hash,
types::{
RuntimeByokLocalLeaseRecord, RuntimeMagicLinkOtpConsumeResult, RuntimeVerificationTokenRecord,
RuntimeWorkspaceInviteLinkRecord,
},
};
mod auth_challenge;
mod byok_local_lease;
mod dto;
mod invite_link;
mod magic_link_otp;
mod store;
mod verification_token;
use store::RuntimeStateStore;
pub(super) type Result<T> = RuntimeResult<T>;
pub(super) fn auth_challenge_purpose(purpose: &str) -> String {
format!("auth_challenge:{purpose}")
}
pub(super) fn verification_token_purpose(token_type: i32) -> String {
format!("verification_token:{token_type}")
}
#[napi_derive::napi]
impl BackendRuntime {
#[napi]
pub async fn create_auth_challenge(
&self,
purpose: String,
token: String,
payload: serde_json::Value,
ttl_ms: i64,
) -> napi::Result<bool> {
if ttl_ms <= 0 {
return Err(napi_error("auth challenge ttl must be positive"));
}
RuntimeStateStore::new(self.pool().await?)
.create_auth_challenge(&purpose, &token, payload, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn get_auth_challenge(&self, purpose: String, token: String) -> napi::Result<Option<serde_json::Value>> {
RuntimeStateStore::new(self.pool().await?)
.get_auth_challenge(&purpose, &token)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn consume_auth_challenge(
&self,
purpose: String,
token: String,
) -> napi::Result<Option<serde_json::Value>> {
RuntimeStateStore::new(self.pool().await?)
.consume_auth_challenge(&purpose, &token)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn create_verification_token(
&self,
token_type: i32,
credential: Option<String>,
ttl_ms: i64,
) -> napi::Result<String> {
if ttl_ms <= 0 {
return Err(napi_error("verification token ttl must be positive"));
}
RuntimeStateStore::new(self.pool().await?)
.create_verification_token(token_type, credential, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn get_verification_token(
&self,
token_type: i32,
token: String,
keep: Option<bool>,
) -> napi::Result<Option<RuntimeVerificationTokenRecord>> {
let keep = keep.unwrap_or(false);
RuntimeStateStore::new(self.pool().await?)
.get_verification_token(token_type, token, keep)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn verify_verification_token(
&self,
token_type: i32,
token: String,
credential: Option<String>,
keep: Option<bool>,
) -> napi::Result<Option<RuntimeVerificationTokenRecord>> {
let keep = keep.unwrap_or(false);
RuntimeStateStore::new(self.pool().await?)
.verify_verification_token(token_type, token, credential, keep)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn cleanup_expired_verification_tokens(&self, limit: i64) -> napi::Result<i64> {
if limit <= 0 {
return Err(napi_error("verification token cleanup limit must be positive"));
}
RuntimeStateStore::new(self.pool().await?)
.cleanup_expired_verification_tokens(limit)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn upsert_magic_link_otp(
&self,
email: String,
otp_hash: String,
token: String,
client_nonce: Option<String>,
ttl_ms: i64,
) -> napi::Result<()> {
RuntimeStateStore::new(self.pool().await?)
.upsert_magic_link_otp(email, otp_hash, token, client_nonce, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn consume_magic_link_otp(
&self,
email: String,
otp_hash: String,
client_nonce: Option<String>,
) -> napi::Result<RuntimeMagicLinkOtpConsumeResult> {
RuntimeStateStore::new(self.pool().await?)
.consume_magic_link_otp(email, otp_hash, client_nonce)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn create_workspace_invite_link(
&self,
workspace_id: String,
invite_id: String,
inviter_user_id: String,
ttl_ms: i64,
) -> napi::Result<RuntimeWorkspaceInviteLinkRecord> {
RuntimeStateStore::new(self.pool().await?)
.create_workspace_invite_link(workspace_id, invite_id, inviter_user_id, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn get_workspace_invite_link(
&self,
workspace_id: String,
) -> napi::Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
RuntimeStateStore::new(self.pool().await?)
.get_workspace_invite_link(workspace_id)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn get_workspace_invite_link_by_id(
&self,
invite_id: String,
) -> napi::Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
RuntimeStateStore::new(self.pool().await?)
.get_workspace_invite_link_by_id(invite_id)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn revoke_workspace_invite_link(&self, workspace_id: String) -> napi::Result<bool> {
RuntimeStateStore::new(self.pool().await?)
.revoke_workspace_invite_link(workspace_id)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn create_byok_local_lease(
&self,
active_key: String,
lease_id: String,
payload: serde_json::Value,
ttl_ms: i64,
) -> napi::Result<RuntimeByokLocalLeaseRecord> {
RuntimeStateStore::new(self.pool().await?)
.create_byok_local_lease(active_key, lease_id, payload, ttl_ms)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn get_byok_local_lease(&self, lease_id: String) -> napi::Result<Option<RuntimeByokLocalLeaseRecord>> {
RuntimeStateStore::new(self.pool().await?)
.get_byok_local_lease(lease_id)
.await
.map_err(napi::Error::from)
}
#[napi]
pub async fn cleanup_expired_runtime_states(&self, limit: i64) -> napi::Result<i64> {
if limit <= 0 {
return Err(napi_error("runtime state cleanup limit must be positive"));
}
RuntimeStateStore::new(self.pool().await?)
.cleanup_expired_runtime_states(limit)
.await
.map_err(napi::Error::from)
}
}
#[cfg(test)]
mod tests {
use super::{
MAGIC_LINK_OTP_PURPOSE, WORKSPACE_INVITE_LINK_ID_PURPOSE, WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE, token_hash,
};
#[test]
fn magic_link_otp_uses_scoped_purpose_and_email_hash() {
assert_eq!(MAGIC_LINK_OTP_PURPOSE, "magic_link_otp");
assert_ne!(token_hash("user@affine.test"), "user@affine.test");
assert_eq!(token_hash("user@affine.test"), token_hash("user@affine.test"));
assert_ne!(token_hash("user@affine.test"), token_hash("other@affine.test"));
}
#[test]
fn workspace_invite_link_uses_scoped_purposes_and_hashes() {
assert_eq!(
WORKSPACE_INVITE_LINK_WORKSPACE_PURPOSE,
"workspace_invite_link:workspace"
);
assert_eq!(WORKSPACE_INVITE_LINK_ID_PURPOSE, "workspace_invite_link:id");
assert_ne!(token_hash("workspace-id"), "workspace-id");
assert_ne!(token_hash("invite-id"), "invite-id");
}
}
@@ -0,0 +1,138 @@
use sqlx::PgPool;
use super::{
Result, RuntimeByokLocalLeaseRecord, RuntimeMagicLinkOtpConsumeResult, RuntimeVerificationTokenRecord,
RuntimeWorkspaceInviteLinkRecord, auth_challenge, byok_local_lease, dto::RuntimeStateRows, invite_link,
magic_link_otp, verification_token,
};
pub(super) struct RuntimeStateStore {
rows: RuntimeStateRows,
}
impl RuntimeStateStore {
pub(super) fn new(pool: PgPool) -> Self {
Self {
rows: RuntimeStateRows::new(pool),
}
}
pub(super) async fn create_auth_challenge(
&self,
purpose: &str,
token: &str,
payload: serde_json::Value,
ttl_ms: i64,
) -> Result<bool> {
auth_challenge::create(&self.rows, purpose, token, payload, ttl_ms).await
}
pub(super) async fn get_auth_challenge(&self, purpose: &str, token: &str) -> Result<Option<serde_json::Value>> {
auth_challenge::get(&self.rows, purpose, token).await
}
pub(super) async fn consume_auth_challenge(&self, purpose: &str, token: &str) -> Result<Option<serde_json::Value>> {
auth_challenge::consume(&self.rows, purpose, token).await
}
pub(super) async fn create_verification_token(
&self,
token_type: i32,
credential: Option<String>,
ttl_ms: i64,
) -> Result<String> {
verification_token::create(&self.rows, token_type, credential, ttl_ms).await
}
pub(super) async fn get_verification_token(
&self,
token_type: i32,
token: String,
keep: bool,
) -> Result<Option<RuntimeVerificationTokenRecord>> {
verification_token::get(&self.rows, token_type, token, keep).await
}
pub(super) async fn verify_verification_token(
&self,
token_type: i32,
token: String,
credential: Option<String>,
keep: bool,
) -> Result<Option<RuntimeVerificationTokenRecord>> {
verification_token::verify(&self.rows, token_type, token, credential, keep).await
}
pub(super) async fn cleanup_expired_verification_tokens(&self, limit: i64) -> Result<i64> {
verification_token::cleanup_expired(&self.rows, limit).await
}
pub(super) async fn cleanup_expired_runtime_states(&self, limit: i64) -> Result<i64> {
self
.rows
.cleanup_expired_or_consumed(limit, "RuntimeState cleanup")
.await
}
pub(super) async fn upsert_magic_link_otp(
&self,
email: String,
otp_hash: String,
token: String,
client_nonce: Option<String>,
ttl_ms: i64,
) -> Result<()> {
magic_link_otp::upsert(&self.rows, email, otp_hash, token, client_nonce, ttl_ms).await
}
pub(super) async fn consume_magic_link_otp(
&self,
email: String,
otp_hash: String,
client_nonce: Option<String>,
) -> Result<RuntimeMagicLinkOtpConsumeResult> {
magic_link_otp::consume(&self.rows, email, otp_hash, client_nonce).await
}
pub(super) async fn create_workspace_invite_link(
&self,
workspace_id: String,
invite_id: String,
inviter_user_id: String,
ttl_ms: i64,
) -> Result<RuntimeWorkspaceInviteLinkRecord> {
invite_link::create(&self.rows, workspace_id, invite_id, inviter_user_id, ttl_ms).await
}
pub(super) async fn get_workspace_invite_link(
&self,
workspace_id: String,
) -> Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
invite_link::get_by_workspace(&self.rows, workspace_id).await
}
pub(super) async fn get_workspace_invite_link_by_id(
&self,
invite_id: String,
) -> Result<Option<RuntimeWorkspaceInviteLinkRecord>> {
invite_link::get_by_invite_id(&self.rows, invite_id).await
}
pub(super) async fn revoke_workspace_invite_link(&self, workspace_id: String) -> Result<bool> {
invite_link::revoke(&self.rows, workspace_id).await
}
pub(super) async fn create_byok_local_lease(
&self,
active_key: String,
lease_id: String,
payload: serde_json::Value,
ttl_ms: i64,
) -> Result<RuntimeByokLocalLeaseRecord> {
byok_local_lease::create(&self.rows, active_key, lease_id, payload, ttl_ms).await
}
pub(super) async fn get_byok_local_lease(&self, lease_id: String) -> Result<Option<RuntimeByokLocalLeaseRecord>> {
byok_local_lease::get(&self.rows, lease_id).await
}
}
@@ -0,0 +1,149 @@
use sqlx::{PgPool, Row};
use uuid::Uuid;
use super::{
Result, RuntimeError, RuntimeVerificationTokenRecord,
dto::{RuntimeStatePayloadRow, RuntimeStateRows},
token_hash, verification_token_purpose,
};
pub(super) async fn create(
rows: &RuntimeStateRows,
token_type: i32,
credential: Option<String>,
ttl_ms: i64,
) -> Result<String> {
let token = Uuid::new_v4().to_string();
let payload = serde_json::json!({ "credential": credential });
rows
.insert_payload(
&verification_token_purpose(token_type),
&token,
credential.as_deref(),
payload,
ttl_ms,
"RuntimeState verification token create",
)
.await?;
Ok(token)
}
pub(super) async fn get(
rows: &RuntimeStateRows,
token_type: i32,
token: String,
keep: bool,
) -> Result<Option<RuntimeVerificationTokenRecord>> {
let purpose = verification_token_purpose(token_type);
let row = if keep {
rows
.active_payload_with_expires(&purpose, &token, "RuntimeState verification token get")
.await?
} else {
rows
.consume_payload_with_expires(&purpose, &token, "RuntimeState verification token get")
.await?
};
Ok(row.map(|row| record_from_row(token_type, token, row)))
}
pub(super) async fn verify(
rows: &RuntimeStateRows,
token_type: i32,
token: String,
credential: Option<String>,
keep: bool,
) -> Result<Option<RuntimeVerificationTokenRecord>> {
let purpose = verification_token_purpose(token_type);
let row = if keep {
active_payload_with_credential(rows.pool(), &purpose, &token, credential.as_deref()).await
} else {
consume_payload_with_credential(rows.pool(), &purpose, &token, credential.as_deref()).await
}
.map_err(|err| RuntimeError::database("RuntimeState verification token verify failed", err))?;
Ok(row.map(|row| record_from_row(token_type, token, row)))
}
pub(super) async fn cleanup_expired(rows: &RuntimeStateRows, limit: i64) -> Result<i64> {
rows
.cleanup_expired_by_purpose_prefix("verification_token:", limit, "RuntimeState verification token cleanup")
.await
}
async fn active_payload_with_credential(
pool: &PgPool,
purpose: &str,
token: &str,
credential: Option<&str>,
) -> sqlx::Result<Option<RuntimeStatePayloadRow>> {
let row = sqlx::query(
r#"
SELECT payload, (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
FROM runtime_states
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > CURRENT_TIMESTAMP
AND (payload->>'credential' IS NULL OR payload->>'credential' = $3)
"#,
)
.bind(purpose)
.bind(token_hash(token))
.bind(credential)
.fetch_optional(pool)
.await?;
Ok(row.map(payload_row))
}
async fn consume_payload_with_credential(
pool: &PgPool,
purpose: &str,
token: &str,
credential: Option<&str>,
) -> sqlx::Result<Option<RuntimeStatePayloadRow>> {
let row = sqlx::query(
r#"
UPDATE runtime_states
SET consumed_at = CURRENT_TIMESTAMP,
updated_at = CURRENT_TIMESTAMP
WHERE purpose = $1
AND token_hash = $2
AND consumed_at IS NULL
AND expires_at > CURRENT_TIMESTAMP
AND (payload->>'credential' IS NULL OR payload->>'credential' = $3)
RETURNING payload, (EXTRACT(EPOCH FROM expires_at) * 1000)::BIGINT AS expires_at_ms
"#,
)
.bind(purpose)
.bind(token_hash(token))
.bind(credential)
.fetch_optional(pool)
.await?;
Ok(row.map(payload_row))
}
fn payload_row(row: sqlx::postgres::PgRow) -> RuntimeStatePayloadRow {
RuntimeStatePayloadRow {
payload: row.get("payload"),
expires_at_ms: row.get("expires_at_ms"),
}
}
fn record_from_row(token_type: i32, token: String, row: RuntimeStatePayloadRow) -> RuntimeVerificationTokenRecord {
RuntimeVerificationTokenRecord {
token_type,
token,
credential: row
.payload
.get("credential")
.and_then(serde_json::Value::as_str)
.map(ToString::to_string),
expires_at_ms: row.expires_at_ms,
}
}
@@ -0,0 +1,85 @@
WITH targets AS (
SELECT UNNEST($1::varchar[]) AS workspace_id
),
snapshot_stats AS (
SELECT workspace_id,
COUNT(*) AS snapshot_count,
COALESCE(SUM(COALESCE(size, octet_length(blob))), 0) AS snapshot_size
FROM snapshots
WHERE workspace_id IN (SELECT workspace_id FROM targets)
GROUP BY workspace_id
),
blob_stats AS (
SELECT workspace_id,
COUNT(*) FILTER (WHERE deleted_at IS NULL AND status = 'completed') AS blob_count,
COALESCE(SUM(size) FILTER (WHERE deleted_at IS NULL AND status = 'completed'), 0) AS blob_size
FROM blobs
WHERE workspace_id IN (SELECT workspace_id FROM targets)
GROUP BY workspace_id
),
member_stats AS (
SELECT workspace_id, COUNT(*) AS member_count
FROM workspace_user_permissions
WHERE workspace_id IN (SELECT workspace_id FROM targets)
GROUP BY workspace_id
),
public_page_stats AS (
SELECT workspace_id, COUNT(*) AS public_page_count
FROM workspace_pages
WHERE public = TRUE AND workspace_id IN (SELECT workspace_id FROM targets)
GROUP BY workspace_id
),
feature_stats AS (
SELECT workspace_id,
ARRAY_AGG(DISTINCT name ORDER BY name) FILTER (WHERE activated) AS features
FROM workspace_features
WHERE workspace_id IN (SELECT workspace_id FROM targets)
GROUP BY workspace_id
),
aggregated AS (
SELECT t.workspace_id,
COALESCE(ss.snapshot_count, 0) AS snapshot_count,
COALESCE(ss.snapshot_size, 0) AS snapshot_size,
COALESCE(bs.blob_count, 0) AS blob_count,
COALESCE(bs.blob_size, 0) AS blob_size,
COALESCE(ms.member_count, 0) AS member_count,
COALESCE(pp.public_page_count, 0) AS public_page_count,
COALESCE(fs.features, ARRAY[]::text[]) AS features
FROM targets t
LEFT JOIN snapshot_stats ss ON ss.workspace_id = t.workspace_id
LEFT JOIN blob_stats bs ON bs.workspace_id = t.workspace_id
LEFT JOIN member_stats ms ON ms.workspace_id = t.workspace_id
LEFT JOIN public_page_stats pp ON pp.workspace_id = t.workspace_id
LEFT JOIN feature_stats fs ON fs.workspace_id = t.workspace_id
)
INSERT INTO workspace_admin_stats (
workspace_id,
snapshot_count,
snapshot_size,
blob_count,
blob_size,
member_count,
public_page_count,
features,
updated_at
)
SELECT
workspace_id,
snapshot_count,
snapshot_size,
blob_count,
blob_size,
member_count,
public_page_count,
features,
NOW()
FROM aggregated
ON CONFLICT (workspace_id) DO UPDATE SET
snapshot_count = EXCLUDED.snapshot_count,
snapshot_size = EXCLUDED.snapshot_size,
blob_count = EXCLUDED.blob_count,
blob_size = EXCLUDED.blob_size,
member_count = EXCLUDED.member_count,
public_page_count = EXCLUDED.public_page_count,
features = EXCLUDED.features,
updated_at = EXCLUDED.updated_at
@@ -0,0 +1,436 @@
use anyhow::{Context, Result as AnyResult, anyhow};
use super::{
super::migrations::{RUNTIME_MIGRATIONS, migrate_runtime_tables},
runtime_state::*,
*,
};
static PG_TEST_LOCK: std::sync::OnceLock<tokio::sync::Mutex<()>> = std::sync::OnceLock::new();
const TEST_VERIFICATION_TOKEN_TYPE: i32 = 99_999;
fn pg_test_lock() -> &'static tokio::sync::Mutex<()> {
PG_TEST_LOCK.get_or_init(|| tokio::sync::Mutex::new(()))
}
#[test]
fn migrations_include_runtime_tables_without_worker_heartbeats() {
assert!(RUNTIME_MIGRATIONS.contains("runtime_states"));
assert!(RUNTIME_MIGRATIONS.contains("runtime_gates"));
assert!(RUNTIME_MIGRATIONS.contains("runtime_leases"));
assert!(RUNTIME_MIGRATIONS.contains("blob_reconciliation_runs"));
assert!(RUNTIME_MIGRATIONS.contains("blob_reconciliation_checkpoints"));
assert!(RUNTIME_MIGRATIONS.contains("doc_blob_refs"));
assert!(RUNTIME_MIGRATIONS.contains("blob_cleanup_candidates"));
assert!(!RUNTIME_MIGRATIONS.contains("runtime_worker_heartbeats"));
}
#[test]
fn auth_challenge_state_uses_scoped_purpose_and_token_hash() {
assert_eq!(auth_challenge_purpose("oauth_state"), "auth_challenge:oauth_state");
assert_ne!(token_hash("plain-token"), "plain-token");
assert_eq!(token_hash("plain-token"), token_hash("plain-token"));
assert_ne!(token_hash("plain-token"), token_hash("other-token"));
}
#[test]
fn verification_token_state_uses_typed_purpose_and_token_hash() {
assert_eq!(verification_token_purpose(0), "verification_token:0");
assert_ne!(token_hash("verification-token"), "verification-token");
assert_eq!(token_hash("verification-token"), token_hash("verification-token"));
assert_ne!(token_hash("verification-token"), token_hash("other-token"));
}
async fn runtime_from_database_url() -> AnyResult<Option<BackendRuntime>> {
let Ok(database_url) = std::env::var("DATABASE_URL") else {
return Ok(None);
};
let pool = PgPoolOptions::new()
.max_connections(5)
.connect(&database_url)
.await
.context("connect postgres for backend runtime tests")?;
migrate_runtime_tables(&pool)
.await
.map_err(|err| anyhow!(err.to_string()))?;
sqlx::query(
r#"
DELETE FROM runtime_states
WHERE purpose LIKE 'rust_test:%'
OR purpose LIKE 'auth_challenge:rust_test:%'
OR purpose = 'verification_token:99999'
"#,
)
.execute(&pool)
.await
.context("cleanup runtime_states for backend runtime tests")?;
sqlx::query("DELETE FROM runtime_gates WHERE key LIKE 'rust-test:%'")
.execute(&pool)
.await
.context("cleanup runtime_gates for backend runtime tests")?;
sqlx::query("DELETE FROM runtime_leases WHERE key LIKE 'rust-test:%'")
.execute(&pool)
.await
.context("cleanup runtime_leases for backend runtime tests")?;
Ok(Some(BackendRuntime {
config: std::sync::RwLock::new(BackendRuntimeConfig { database_url }),
pool: Mutex::new(Some(pool)),
}))
}
#[tokio::test]
async fn runtime_gate_sql_semantics_are_atomic_and_ttl_bound() {
let _guard = pg_test_lock().lock().await;
let Some(runtime) = runtime_from_database_url().await.unwrap() else {
eprintln!("skipping postgres integration test: DATABASE_URL is not set");
return;
};
struct Case {
key: &'static str,
first_ttl_ms: i64,
wait_ms: Option<u64>,
second_expected: bool,
}
for case in [
Case {
key: "rust-test:gate:same-key",
first_ttl_ms: 30_000,
wait_ms: None,
second_expected: false,
},
Case {
key: "rust-test:gate:expired-key",
first_ttl_ms: 1,
wait_ms: Some(20),
second_expected: true,
},
] {
assert!(
runtime
.put_runtime_gate_if_absent(case.key.to_string(), case.first_ttl_ms)
.await
.unwrap()
);
if let Some(wait_ms) = case.wait_ms {
tokio::time::sleep(Duration::from_millis(wait_ms)).await;
}
assert_eq!(
runtime
.put_runtime_gate_if_absent(case.key.to_string(), 30_000)
.await
.unwrap(),
case.second_expected,
"{}",
case.key
);
}
let mut tasks = Vec::new();
for _ in 0..16 {
let runtime = BackendRuntime {
config: std::sync::RwLock::new(runtime.config().unwrap()),
pool: Mutex::new(Some(runtime.pool().await.unwrap())),
};
tasks.push(tokio::spawn(async move {
runtime
.put_runtime_gate_if_absent("rust-test:gate:concurrent".to_string(), 30_000)
.await
.unwrap()
}));
}
let mut successful = 0;
for task in tasks {
if task.await.unwrap() {
successful += 1;
}
}
assert_eq!(successful, 1);
assert!(
runtime
.put_runtime_gate_if_absent("rust-test:gate:cleanup".to_string(), 1)
.await
.unwrap()
);
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(runtime.cleanup_expired_runtime_gates(100).await.unwrap(), 1);
assert_eq!(runtime.cleanup_expired_runtime_gates(100).await.unwrap(), 0);
}
#[tokio::test]
async fn coordination_lease_sql_semantics_are_fenced_and_ttl_bound() {
let _guard = pg_test_lock().lock().await;
let Some(runtime) = runtime_from_database_url().await.unwrap() else {
eprintln!("skipping postgres integration test: DATABASE_URL is not set");
return;
};
let lease = runtime
.acquire_coordination_lease("rust-test:lease:basic".to_string(), "owner-1".to_string(), 30_000)
.await
.unwrap()
.expect("first owner should acquire lease");
assert_eq!(lease.fencing_token, 1);
assert!(
!runtime
.release_coordination_lease(lease.key.clone(), "owner-2".to_string(), lease.fencing_token)
.await
.unwrap()
);
assert!(
runtime
.release_coordination_lease(lease.key.clone(), lease.owner.clone(), lease.fencing_token)
.await
.unwrap()
);
let mut tasks = Vec::new();
for index in 0..16 {
let runtime = BackendRuntime {
config: std::sync::RwLock::new(runtime.config().unwrap()),
pool: Mutex::new(Some(runtime.pool().await.unwrap())),
};
tasks.push(tokio::spawn(async move {
runtime
.acquire_coordination_lease(
"rust-test:lease:concurrent".to_string(),
format!("owner-{index}"),
30_000,
)
.await
.unwrap()
.is_some()
}));
}
let mut successful = 0;
for task in tasks {
if task.await.unwrap() {
successful += 1;
}
}
assert_eq!(successful, 1);
let stale = runtime
.acquire_coordination_lease("rust-test:lease:stale".to_string(), "owner-1".to_string(), 1)
.await
.unwrap()
.expect("stale lease owner should acquire");
tokio::time::sleep(Duration::from_millis(20)).await;
let takeover = runtime
.acquire_coordination_lease("rust-test:lease:stale".to_string(), "owner-2".to_string(), 30_000)
.await
.unwrap()
.expect("expired lease should be taken over");
assert_eq!(takeover.fencing_token, stale.fencing_token + 1);
assert!(
!runtime
.release_coordination_lease(stale.key.clone(), stale.owner.clone(), stale.fencing_token)
.await
.unwrap()
);
let renew = runtime
.acquire_coordination_lease("rust-test:lease:renew".to_string(), "owner-1".to_string(), 30_000)
.await
.unwrap()
.expect("renew lease owner should acquire");
assert!(
!runtime
.renew_coordination_lease(renew.key.clone(), "owner-2".to_string(), renew.fencing_token, 30_000)
.await
.unwrap()
);
assert!(
!runtime
.renew_coordination_lease(renew.key.clone(), renew.owner.clone(), renew.fencing_token + 1, 30_000)
.await
.unwrap()
);
assert!(
runtime
.renew_coordination_lease(renew.key.clone(), renew.owner.clone(), renew.fencing_token, 30_000)
.await
.unwrap()
);
}
#[tokio::test]
async fn runtime_state_cleanup_deletes_expired_and_consumed_rows() {
let _guard = pg_test_lock().lock().await;
let Some(runtime) = runtime_from_database_url().await.unwrap() else {
eprintln!("skipping postgres integration test: DATABASE_URL is not set");
return;
};
assert!(
runtime
.create_auth_challenge(
"rust_test:cleanup".to_string(),
"expired".to_string(),
serde_json::json!({}),
1
)
.await
.unwrap()
);
assert!(
runtime
.create_auth_challenge(
"rust_test:cleanup".to_string(),
"consumed".to_string(),
serde_json::json!({}),
30_000,
)
.await
.unwrap()
);
assert!(
runtime
.consume_auth_challenge("rust_test:cleanup".to_string(), "consumed".to_string())
.await
.unwrap()
.is_some()
);
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(runtime.cleanup_expired_runtime_states(100).await.unwrap(), 2);
assert_eq!(runtime.cleanup_expired_runtime_states(100).await.unwrap(), 0);
}
#[tokio::test]
async fn verification_token_sql_state_machine_handles_keep_verify_and_cleanup() {
let _guard = pg_test_lock().lock().await;
let Some(runtime) = runtime_from_database_url().await.unwrap() else {
eprintln!("skipping postgres integration test: DATABASE_URL is not set");
return;
};
let mismatch_token = runtime
.create_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
Some("user@affine.test".to_string()),
30_000,
)
.await
.unwrap();
assert!(
runtime
.verify_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
mismatch_token.clone(),
Some("wrong@affine.test".to_string()),
None,
)
.await
.unwrap()
.is_none()
);
assert!(
runtime
.verify_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
mismatch_token.clone(),
Some("user@affine.test".to_string()),
None,
)
.await
.unwrap()
.is_some()
);
assert!(
runtime
.verify_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
mismatch_token.clone(),
Some("user@affine.test".to_string()),
None,
)
.await
.unwrap()
.is_none()
);
let keep_token = runtime
.create_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
Some("keep@affine.test".to_string()),
30_000,
)
.await
.unwrap();
assert!(
runtime
.get_verification_token(TEST_VERIFICATION_TOKEN_TYPE, keep_token.clone(), Some(true))
.await
.unwrap()
.is_some()
);
assert!(
runtime
.get_verification_token(TEST_VERIFICATION_TOKEN_TYPE, keep_token.clone(), None)
.await
.unwrap()
.is_some()
);
assert!(
runtime
.get_verification_token(TEST_VERIFICATION_TOKEN_TYPE, keep_token.clone(), None)
.await
.unwrap()
.is_none()
);
let concurrent_token = runtime
.create_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
Some("concurrent@affine.test".to_string()),
30_000,
)
.await
.unwrap();
let mut tasks = Vec::new();
for _ in 0..16 {
let runtime = BackendRuntime {
config: std::sync::RwLock::new(runtime.config().unwrap()),
pool: Mutex::new(Some(runtime.pool().await.unwrap())),
};
let token = concurrent_token.clone();
tasks.push(tokio::spawn(async move {
runtime
.verify_verification_token(
TEST_VERIFICATION_TOKEN_TYPE,
token,
Some("concurrent@affine.test".to_string()),
None,
)
.await
.unwrap()
.is_some()
}));
}
let mut successful = 0;
for task in tasks {
if task.await.unwrap() {
successful += 1;
}
}
assert_eq!(successful, 1);
let expired_token = runtime
.create_verification_token(TEST_VERIFICATION_TOKEN_TYPE, Some("expired@affine.test".to_string()), 1)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(
runtime
.get_verification_token(TEST_VERIFICATION_TOKEN_TYPE, expired_token.clone(), None)
.await
.unwrap()
.is_none()
);
assert_eq!(runtime.cleanup_expired_verification_tokens(100).await.unwrap(), 1);
assert_eq!(runtime.cleanup_expired_verification_tokens(100).await.unwrap(), 0);
}
@@ -0,0 +1,530 @@
use sqlx::{FromRow, PgPool, Postgres, Row, Transaction};
use tokio::time::{Duration as TokioDuration, sleep};
use super::{
BackendRuntime, RuntimeError, RuntimeResult,
constants::{WORKSPACE_STATS_LEASE_KEY, WORKSPACE_STATS_LOCK_NAMESPACE, WORKSPACE_STATS_REFRESH_LOCK_KEY},
napi_error,
types::{
CoordinationLeaseGrant, RuntimeWorkspaceStatsDailyRecalibrationResult, RuntimeWorkspaceStatsRecalibrationResult,
RuntimeWorkspaceStatsRefreshResult, RuntimeWorkspaceStatsSnapshotResult,
},
};
const UPSERT_WORKSPACE_ADMIN_STATS_SQL: &str = include_str!("sql/upsert_workspace_admin_stats.sql");
#[napi_derive::napi]
impl BackendRuntime {
#[napi]
pub async fn refresh_workspace_admin_stats_dirty(
&self,
batch_limit: i64,
owner: String,
lease_ttl_ms: i64,
) -> napi::Result<RuntimeWorkspaceStatsRefreshResult> {
if batch_limit <= 0 {
return Err(napi_error("workspace stats dirty refresh limit must be positive"));
}
let Some(lease) = self
.acquire_coordination_lease_inner(WORKSPACE_STATS_LEASE_KEY.to_string(), owner, lease_ttl_ms)
.await?
else {
return Ok(RuntimeWorkspaceStatsRefreshResult {
processed: 0,
backlog: 0,
skipped: true,
});
};
let result = async {
WorkspaceStatsStore::new(self.pool().await?)
.refresh_dirty(batch_limit)
.await
}
.await;
release_workspace_stats_lease(self, lease).await?;
Ok(result?)
}
#[napi]
pub async fn recalibrate_workspace_admin_stats(
&self,
last_sid: i64,
batch_limit: i64,
owner: String,
lease_ttl_ms: i64,
) -> napi::Result<RuntimeWorkspaceStatsRecalibrationResult> {
if batch_limit <= 0 {
return Err(napi_error("workspace stats recalibration limit must be positive"));
}
let Some(lease) = self
.acquire_coordination_lease_inner(WORKSPACE_STATS_LEASE_KEY.to_string(), owner, lease_ttl_ms)
.await?
else {
return Ok(RuntimeWorkspaceStatsRecalibrationResult {
processed: 0,
last_sid,
skipped: true,
});
};
let result = async {
WorkspaceStatsStore::new(self.pool().await?)
.recalibrate(last_sid, batch_limit)
.await
}
.await;
release_workspace_stats_lease(self, lease).await?;
Ok(result?)
}
#[napi]
pub async fn write_workspace_admin_stats_daily_snapshot(
&self,
owner: String,
lease_ttl_ms: i64,
) -> napi::Result<RuntimeWorkspaceStatsSnapshotResult> {
let Some(lease) = self
.acquire_coordination_lease_inner(WORKSPACE_STATS_LEASE_KEY.to_string(), owner, lease_ttl_ms)
.await?
else {
return Ok(RuntimeWorkspaceStatsSnapshotResult {
snapshotted: 0,
skipped: true,
});
};
let result = async {
WorkspaceStatsStore::new(self.pool().await?)
.write_daily_snapshot()
.await
}
.await;
release_workspace_stats_lease(self, lease).await?;
Ok(result?)
}
#[napi]
pub async fn recalibrate_workspace_admin_stats_daily(
&self,
batch_limit: i64,
owner: String,
lease_ttl_ms: i64,
lock_retry_times: i64,
lock_retry_delay_ms: i64,
) -> napi::Result<RuntimeWorkspaceStatsDailyRecalibrationResult> {
if batch_limit <= 0 {
return Err(napi_error("workspace stats daily recalibration limit must be positive"));
}
if lock_retry_times <= 0 {
return Err(napi_error(
"workspace stats daily recalibration retry times must be positive",
));
}
if lock_retry_delay_ms < 0 {
return Err(napi_error(
"workspace stats daily recalibration retry delay must be non-negative",
));
}
let Some(lease) = acquire_workspace_stats_lease_with_retry(
self,
owner.clone(),
lease_ttl_ms,
lock_retry_times,
lock_retry_delay_ms,
)
.await?
else {
return Ok(RuntimeWorkspaceStatsDailyRecalibrationResult {
processed: 0,
last_sid: 0,
snapshotted: 0,
skipped: true,
});
};
let result: RuntimeResult<RuntimeWorkspaceStatsDailyRecalibrationResult> = async {
let store = WorkspaceStatsStore::new(self.pool().await?);
let mut processed = 0;
let mut last_sid = 0;
loop {
let batch = retry_workspace_stats_operation(lock_retry_times, lock_retry_delay_ms, || {
store.recalibrate(last_sid, batch_limit)
})
.await?;
if batch.skipped {
return Ok(RuntimeWorkspaceStatsDailyRecalibrationResult {
processed,
last_sid,
snapshotted: 0,
skipped: true,
});
}
if batch.processed == 0 {
break;
}
processed += batch.processed;
last_sid = batch.last_sid;
if batch.processed < batch_limit {
break;
}
}
let snapshot =
retry_workspace_stats_operation(lock_retry_times, lock_retry_delay_ms, || store.write_daily_snapshot()).await?;
Ok(RuntimeWorkspaceStatsDailyRecalibrationResult {
processed,
last_sid,
snapshotted: snapshot.snapshotted,
skipped: snapshot.skipped,
})
}
.await;
release_workspace_stats_lease(self, lease).await?;
Ok(result?)
}
}
#[derive(FromRow)]
struct WorkspaceSid {
id: String,
sid: i32,
}
struct WorkspaceStatsStore {
pool: PgPool,
}
impl WorkspaceStatsStore {
fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn refresh_dirty(&self, batch_limit: i64) -> RuntimeResult<RuntimeWorkspaceStatsRefreshResult> {
let mut tx = self
.pool
.begin()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats dirty refresh transaction failed", err))?;
if !try_transaction_lock(&mut tx).await? {
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats dirty refresh commit failed", err))?;
return Ok(RuntimeWorkspaceStatsRefreshResult {
processed: 0,
backlog: 0,
skipped: true,
});
}
let backlog = count_dirty(&mut tx).await?;
let dirty = load_dirty(&mut tx, batch_limit).await?;
if dirty.is_empty() {
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats dirty refresh commit failed", err))?;
return Ok(RuntimeWorkspaceStatsRefreshResult {
processed: 0,
backlog,
skipped: false,
});
}
upsert_stats(&mut tx, &dirty).await?;
clear_dirty(&mut tx, &dirty).await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats dirty refresh commit failed", err))?;
Ok(RuntimeWorkspaceStatsRefreshResult {
processed: dirty.len() as i64,
backlog,
skipped: false,
})
}
async fn recalibrate(
&self,
last_sid: i64,
batch_limit: i64,
) -> RuntimeResult<RuntimeWorkspaceStatsRecalibrationResult> {
let mut tx = self
.pool
.begin()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats recalibration transaction failed", err))?;
if !try_transaction_lock(&mut tx).await? {
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats recalibration commit failed", err))?;
return Ok(RuntimeWorkspaceStatsRecalibrationResult {
processed: 0,
last_sid,
skipped: true,
});
}
let workspaces = fetch_workspace_batch(&mut tx, last_sid, batch_limit).await?;
if workspaces.is_empty() {
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats recalibration commit failed", err))?;
return Ok(RuntimeWorkspaceStatsRecalibrationResult {
processed: 0,
last_sid,
skipped: false,
});
}
let ids = workspaces
.iter()
.map(|workspace| workspace.id.clone())
.collect::<Vec<_>>();
let next_sid = workspaces
.last()
.map(|workspace| workspace.sid as i64)
.unwrap_or(last_sid);
upsert_stats(&mut tx, &ids).await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats recalibration commit failed", err))?;
Ok(RuntimeWorkspaceStatsRecalibrationResult {
processed: ids.len() as i64,
last_sid: next_sid,
skipped: false,
})
}
async fn write_daily_snapshot(&self) -> RuntimeResult<RuntimeWorkspaceStatsSnapshotResult> {
let mut tx = self
.pool
.begin()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats daily snapshot transaction failed", err))?;
if !try_transaction_lock(&mut tx).await? {
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats daily snapshot commit failed", err))?;
return Ok(RuntimeWorkspaceStatsSnapshotResult {
snapshotted: 0,
skipped: true,
});
}
let snapshotted = write_daily_snapshot(&mut tx).await?;
tx.commit()
.await
.map_err(|err| RuntimeError::database("WorkspaceStats daily snapshot commit failed", err))?;
Ok(RuntimeWorkspaceStatsSnapshotResult {
snapshotted,
skipped: false,
})
}
}
async fn release_workspace_stats_lease(runtime: &BackendRuntime, lease: CoordinationLeaseGrant) -> RuntimeResult<()> {
let _ = runtime
.release_coordination_lease_inner(lease.key, lease.owner, lease.fencing_token)
.await?;
Ok(())
}
async fn acquire_workspace_stats_lease_with_retry(
runtime: &BackendRuntime,
owner: String,
lease_ttl_ms: i64,
retry_times: i64,
retry_delay_ms: i64,
) -> RuntimeResult<Option<CoordinationLeaseGrant>> {
for attempt in 0..retry_times {
let lease = runtime
.acquire_coordination_lease_inner(WORKSPACE_STATS_LEASE_KEY.to_string(), owner.clone(), lease_ttl_ms)
.await?;
if lease.is_some() {
return Ok(lease);
}
if attempt < retry_times - 1 && retry_delay_ms > 0 {
sleep(TokioDuration::from_millis(retry_delay_ms as u64)).await;
}
}
Ok(None)
}
async fn retry_workspace_stats_operation<T, F, Fut>(
retry_times: i64,
retry_delay_ms: i64,
mut operation: F,
) -> RuntimeResult<T>
where
T: WorkspaceStatsSkippable,
F: FnMut() -> Fut,
Fut: std::future::Future<Output = RuntimeResult<T>>,
{
for attempt in 0..retry_times {
let result = operation().await?;
if !result.skipped() || attempt == retry_times - 1 {
return Ok(result);
}
if retry_delay_ms > 0 {
sleep(TokioDuration::from_millis(retry_delay_ms as u64)).await;
}
}
unreachable!("workspace stats retry loop validates retry_times > 0")
}
trait WorkspaceStatsSkippable {
fn skipped(&self) -> bool;
}
impl WorkspaceStatsSkippable for RuntimeWorkspaceStatsRecalibrationResult {
fn skipped(&self) -> bool {
self.skipped
}
}
impl WorkspaceStatsSkippable for RuntimeWorkspaceStatsSnapshotResult {
fn skipped(&self) -> bool {
self.skipped
}
}
async fn try_transaction_lock(tx: &mut Transaction<'_, Postgres>) -> RuntimeResult<bool> {
let row = sqlx::query(
r#"
SELECT pg_try_advisory_xact_lock(($1::bigint << 32) + $2::bigint) AS locked
"#,
)
.bind(WORKSPACE_STATS_LOCK_NAMESPACE)
.bind(WORKSPACE_STATS_REFRESH_LOCK_KEY)
.fetch_one(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats transaction lock failed", err))?;
Ok(row.get::<bool, _>("locked"))
}
async fn load_dirty(tx: &mut Transaction<'_, Postgres>, limit: i64) -> RuntimeResult<Vec<String>> {
let rows = sqlx::query(
r#"
SELECT workspace_id
FROM workspace_admin_stats_dirty
ORDER BY updated_at ASC
LIMIT $1
FOR UPDATE SKIP LOCKED
"#,
)
.bind(limit)
.fetch_all(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats load dirty workspaces failed", err))?;
Ok(rows.into_iter().map(|row| row.get("workspace_id")).collect())
}
async fn count_dirty(tx: &mut Transaction<'_, Postgres>) -> RuntimeResult<i64> {
let row = sqlx::query("SELECT COUNT(*) AS total FROM workspace_admin_stats_dirty")
.fetch_one(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats count dirty workspaces failed", err))?;
Ok(row.get::<i64, _>("total"))
}
async fn clear_dirty(tx: &mut Transaction<'_, Postgres>, workspace_ids: &[String]) -> RuntimeResult<()> {
sqlx::query(
r#"
DELETE FROM workspace_admin_stats_dirty
WHERE workspace_id = ANY($1::varchar[])
"#,
)
.bind(workspace_ids)
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats clear dirty workspaces failed", err))?;
Ok(())
}
async fn upsert_stats(tx: &mut Transaction<'_, Postgres>, workspace_ids: &[String]) -> RuntimeResult<()> {
if workspace_ids.is_empty() {
return Ok(());
}
sqlx::query(UPSERT_WORKSPACE_ADMIN_STATS_SQL)
.bind(workspace_ids)
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats upsert stats failed", err))?;
Ok(())
}
async fn fetch_workspace_batch(
tx: &mut Transaction<'_, Postgres>,
last_sid: i64,
limit: i64,
) -> RuntimeResult<Vec<WorkspaceSid>> {
sqlx::query_as::<_, WorkspaceSid>(
r#"
SELECT id, sid
FROM workspaces
WHERE sid > $1
ORDER BY sid
LIMIT $2
"#,
)
.bind(last_sid)
.bind(limit)
.fetch_all(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats fetch workspace batch failed", err))
}
async fn write_daily_snapshot(tx: &mut Transaction<'_, Postgres>) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
INSERT INTO workspace_admin_stats_daily (
workspace_id,
date,
snapshot_size,
blob_size,
member_count,
updated_at
)
SELECT
workspace_id,
CURRENT_DATE,
snapshot_size,
blob_size,
member_count,
NOW()
FROM workspace_admin_stats
ON CONFLICT (workspace_id, date)
DO UPDATE SET
snapshot_size = EXCLUDED.snapshot_size,
blob_size = EXCLUDED.blob_size,
member_count = EXCLUDED.member_count,
updated_at = EXCLUDED.updated_at
"#,
)
.execute(&mut **tx)
.await
.map_err(|err| RuntimeError::database("WorkspaceStats daily snapshot failed", err))?;
Ok(result.rows_affected() as i64)
}
@@ -0,0 +1,200 @@
use std::{
env, fs,
path::{Path, PathBuf},
};
use serde::Deserialize;
use serde_json::Map;
use sqlx::{PgPool, Row};
use super::{RuntimeError, RuntimeResult};
#[derive(Clone, Debug)]
pub(crate) struct BackendRuntimeConfig {
pub(crate) database_url: String,
}
impl BackendRuntimeConfig {
pub(crate) fn from_config_files() -> RuntimeResult<Self> {
let app_config = app_config_from_config_files()?;
let database_url = database_url_from_env()
.or(app_config.database_url())
.unwrap_or_else(|| "postgresql://localhost:5432/affine".to_string());
Ok(Self { database_url })
}
pub(crate) async fn with_db_overrides(&self, pool: &PgPool) -> RuntimeResult<Self> {
let mut app_config = app_config_from_config_files()?;
app_config.apply_file_config(load_app_config_overrides_from_db(pool).await?);
Ok(Self {
// The DB override is loaded after this connection already exists, so it
// must not rewrite the active datasource URL.
database_url: self.database_url.clone(),
})
}
}
#[derive(Debug, Default, Deserialize)]
struct AppConfigFile {
db: Option<DbConfigFile>,
}
#[derive(Debug, Default, Deserialize)]
#[serde(rename_all = "camelCase")]
struct DbConfigFile {
datasource_url: Option<String>,
}
impl AppConfigFile {
fn database_url(&self) -> Option<String> {
self
.db
.as_ref()
.and_then(|db| db.datasource_url.clone())
.and_then(non_empty_string)
}
}
fn database_url_from_env() -> Option<String> {
env::var("DATABASE_URL").ok().and_then(non_empty_string)
}
fn non_empty_string(value: String) -> Option<String> {
if value.trim().is_empty() { None } else { Some(value) }
}
fn app_config_from_config_files() -> RuntimeResult<AppConfigFile> {
let mut merged = AppConfigFile::default();
for path in config_json_paths() {
if !path.exists() {
continue;
}
let raw = fs::read_to_string(&path).map_err(|err| RuntimeError::io("failed to read config file", err))?;
let config: AppConfigFile =
serde_json::from_str(&raw).map_err(|err| RuntimeError::json("failed to parse config file", err))?;
merged.apply_file_config(config);
}
Ok(merged)
}
impl AppConfigFile {
fn apply_file_config(&mut self, config: AppConfigFile) {
if config.db.is_some() {
self.db = config.db;
}
}
}
async fn load_app_config_overrides_from_db(pool: &PgPool) -> RuntimeResult<AppConfigFile> {
let rows = match sqlx::query("SELECT id, value FROM app_configs").fetch_all(pool).await {
Ok(rows) => rows,
Err(sqlx::Error::Database(err)) if err.code().as_deref() == Some("42P01") => return Ok(AppConfigFile::default()),
Err(err) => return Err(RuntimeError::database("failed to load app config overrides", err)),
};
app_config_from_flat_overrides(rows.into_iter().map(|row| {
let id: String = row.get("id");
let value: serde_json::Value = row.get("value");
(id, value)
}))
}
fn app_config_from_flat_overrides<I, S>(rows: I) -> RuntimeResult<AppConfigFile>
where
I: IntoIterator<Item = (S, serde_json::Value)>,
S: AsRef<str>,
{
let mut root = Map::new();
for (path, value) in rows {
let Some((module, key)) = path.as_ref().split_once('.') else {
continue;
};
root
.entry(module.to_string())
.or_insert_with(|| serde_json::Value::Object(Map::new()));
if let Some(serde_json::Value::Object(module_object)) = root.get_mut(module) {
module_object.insert(key.to_string(), value);
}
}
serde_json::from_value(serde_json::Value::Object(root))
.map_err(|err| RuntimeError::json("invalid app config overrides", err))
}
pub(super) fn config_json_paths() -> Vec<PathBuf> {
let mut paths = Vec::new();
if let Ok(exe) = env::current_exe()
&& let Some(dir) = exe.parent()
{
paths.push(config_in(dir));
}
if let Ok(cwd) = env::current_dir() {
paths.push(config_in(&cwd));
}
dedupe_paths(paths)
}
fn config_in(dir: &Path) -> PathBuf {
dir.join("config.json")
}
fn dedupe_paths(paths: Vec<PathBuf>) -> Vec<PathBuf> {
let mut deduped = Vec::new();
for path in paths {
if !deduped.contains(&path) {
deduped.push(path);
}
}
deduped
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn config_paths_are_limited_to_executable_dir_and_cwd() {
let paths = config_json_paths();
assert!(!paths.is_empty());
assert!(paths.len() <= 2);
assert!(
paths
.iter()
.all(|path| path.file_name().is_some_and(|name| name == "config.json"))
);
assert!(paths.iter().all(|path| !path.to_string_lossy().contains(".affine")));
assert!(
paths
.iter()
.all(|path| !path.to_string_lossy().contains("packages/backend/server"))
);
}
#[test]
fn blank_database_urls_are_ignored() {
assert_eq!(non_empty_string("".to_string()), None);
assert_eq!(non_empty_string(" ".to_string()), None);
assert_eq!(
non_empty_string("postgresql://affine:affine@localhost:5432/affine".to_string()),
Some("postgresql://affine:affine@localhost:5432/affine".to_string())
);
}
#[test]
fn ignores_storage_app_config_values() {
let app_config = app_config_from_flat_overrides([
(
"storages.blob.storage",
serde_json::json!({"provider": "cloudflare-r2"}),
),
("db.datasourceUrl", serde_json::json!("postgresql://example/runtime")),
])
.unwrap();
assert_eq!(
app_config.database_url().as_deref(),
Some("postgresql://example/runtime")
);
}
}
@@ -0,0 +1,126 @@
use napi::{Error, Status};
use super::storage_runtime::object_storage::error::ObjectStorageError;
pub(crate) type RuntimeResult<T> = std::result::Result<T, RuntimeError>;
#[derive(Debug, thiserror::Error)]
pub(crate) enum RuntimeError {
#[error("{0}")]
Config(String),
#[error("{0}")]
InvalidInput(String),
#[error("{0}")]
InvalidState(String),
#[error("{context}: {source}")]
Database {
context: String,
#[source]
source: sqlx::Error,
},
#[error("{context}: {source}")]
Io {
context: String,
#[source]
source: std::io::Error,
},
#[error("{context}: {source}")]
Json {
context: String,
#[source]
source: serde_json::Error,
},
#[error("{context}: {source}")]
Time {
context: String,
#[source]
source: std::time::SystemTimeError,
},
#[error(transparent)]
ObjectStorage(#[from] ObjectStorageError),
#[error("{0}")]
NapiBoundary(String),
}
impl RuntimeError {
pub(crate) fn config(message: impl Into<String>) -> Self {
Self::Config(message.into())
}
pub(crate) fn invalid_input(message: impl Into<String>) -> Self {
Self::InvalidInput(message.into())
}
pub(crate) fn invalid_state(message: impl Into<String>) -> Self {
Self::InvalidState(message.into())
}
pub(crate) fn database(context: impl Into<String>, source: sqlx::Error) -> Self {
Self::Database {
context: context.into(),
source,
}
}
pub(crate) fn io(context: impl Into<String>, source: std::io::Error) -> Self {
Self::Io {
context: context.into(),
source,
}
}
pub(crate) fn json(context: impl Into<String>, source: serde_json::Error) -> Self {
Self::Json {
context: context.into(),
source,
}
}
pub(crate) fn is_object_missing(&self) -> bool {
match self {
Self::ObjectStorage(error) => error.is_not_found(),
Self::Io { source, .. } => source.kind() == std::io::ErrorKind::NotFound,
Self::InvalidState(message)
| Self::InvalidInput(message)
| Self::Config(message)
| Self::NapiBoundary(message) => {
message.contains("NoSuchKey") || message.contains("NotFound") || message.contains("not found")
}
_ => false,
}
}
}
pub(crate) fn to_napi_error(error: RuntimeError) -> Error {
Error::new(Status::GenericFailure, error.to_string())
}
impl From<RuntimeError> for Error {
fn from(error: RuntimeError) -> Self {
to_napi_error(error)
}
}
impl From<ObjectStorageError> for Error {
fn from(error: ObjectStorageError) -> Self {
to_napi_error(RuntimeError::from(error))
}
}
impl From<Error> for RuntimeError {
fn from(error: Error) -> Self {
Self::NapiBoundary(error.to_string())
}
}
pub(crate) fn napi_error(message: impl Into<String>) -> Error {
Error::new(Status::GenericFailure, message.into())
}
@@ -0,0 +1,20 @@
use sqlx::PgPool;
use super::{RuntimeError, RuntimeResult};
pub(crate) const RUNTIME_MIGRATIONS: &str = include_str!("sql/runtime_migrations.sql");
pub(crate) async fn migrate_runtime_tables(pool: &PgPool) -> RuntimeResult<()> {
for statement in RUNTIME_MIGRATIONS
.split(';')
.map(str::trim)
.filter(|statement| !statement.is_empty())
{
sqlx::query(statement)
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Runtime migration failed", err))?;
}
Ok(())
}
@@ -0,0 +1,10 @@
pub mod backend_runtime;
pub mod storage_runtime;
pub(crate) mod config;
pub(crate) mod error;
pub(crate) mod migrations;
pub(crate) mod types;
pub(crate) use config::BackendRuntimeConfig;
pub(crate) use error::{RuntimeError, RuntimeResult, napi_error, to_napi_error};
@@ -0,0 +1,112 @@
CREATE TABLE IF NOT EXISTS runtime_states (
purpose TEXT NOT NULL,
token_hash TEXT NOT NULL,
lookup_key TEXT,
payload JSONB NOT NULL,
attempts INTEGER NOT NULL DEFAULT 0,
consumed_at TIMESTAMPTZ(3),
expires_at TIMESTAMPTZ(3) NOT NULL,
created_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY (purpose, token_hash)
);
CREATE INDEX IF NOT EXISTS runtime_states_lookup_idx
ON runtime_states (purpose, lookup_key)
WHERE lookup_key IS NOT NULL AND consumed_at IS NULL;
CREATE INDEX IF NOT EXISTS runtime_states_expires_at_idx
ON runtime_states (expires_at);
CREATE TABLE IF NOT EXISTS runtime_gates (
key TEXT PRIMARY KEY,
expires_at TIMESTAMPTZ(3) NOT NULL,
created_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS runtime_gates_expires_at_idx
ON runtime_gates (expires_at);
CREATE TABLE IF NOT EXISTS runtime_leases (
key TEXT PRIMARY KEY,
owner TEXT NOT NULL,
fencing_token BIGINT NOT NULL,
expires_at TIMESTAMPTZ(3) NOT NULL,
created_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS runtime_leases_expires_at_idx
ON runtime_leases (expires_at);
CREATE TABLE IF NOT EXISTS blob_reconciliation_runs (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
kind TEXT NOT NULL,
mode TEXT NOT NULL,
status TEXT NOT NULL,
workspace_id TEXT,
started_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
finished_at TIMESTAMPTZ(3),
cursor JSONB NOT NULL DEFAULT '{}',
scanned INTEGER NOT NULL DEFAULT 0,
changed INTEGER NOT NULL DEFAULT 0,
failed INTEGER NOT NULL DEFAULT 0,
metadata JSONB NOT NULL DEFAULT '{}'
);
CREATE INDEX IF NOT EXISTS blob_reconciliation_runs_workspace_idx
ON blob_reconciliation_runs (workspace_id, started_at DESC);
CREATE TABLE IF NOT EXISTS blob_reconciliation_checkpoints (
kind TEXT NOT NULL,
scope TEXT NOT NULL,
status TEXT NOT NULL,
cursor JSONB NOT NULL DEFAULT '{}',
last_key TEXT,
last_sid INTEGER,
completed_at TIMESTAMPTZ(3),
updated_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
metadata JSONB NOT NULL DEFAULT '{}',
PRIMARY KEY (kind, scope)
);
CREATE INDEX IF NOT EXISTS blob_reconciliation_checkpoints_status_idx
ON blob_reconciliation_checkpoints (kind, status, updated_at DESC);
CREATE TABLE IF NOT EXISTS doc_blob_refs (
workspace_id TEXT NOT NULL,
doc_id TEXT NOT NULL,
blob_key TEXT NOT NULL,
block_id TEXT NOT NULL,
flavour TEXT NOT NULL,
snapshot_updated_at TIMESTAMPTZ(3) NOT NULL,
indexed_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
parser_version INTEGER NOT NULL,
status TEXT NOT NULL DEFAULT 'fresh',
error TEXT,
PRIMARY KEY (workspace_id, doc_id, blob_key, block_id)
);
CREATE INDEX IF NOT EXISTS doc_blob_refs_workspace_blob_idx
ON doc_blob_refs (workspace_id, blob_key);
CREATE INDEX IF NOT EXISTS doc_blob_refs_workspace_status_idx
ON doc_blob_refs (workspace_id, status);
CREATE TABLE IF NOT EXISTS blob_cleanup_candidates (
workspace_id TEXT NOT NULL,
blob_key TEXT NOT NULL,
reason TEXT NOT NULL,
status TEXT NOT NULL,
object_size BIGINT NOT NULL,
object_last_modified TIMESTAMPTZ(3),
planned_at TIMESTAMPTZ(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
executed_at TIMESTAMPTZ(3),
run_id UUID NOT NULL,
evidence JSONB NOT NULL DEFAULT '{}',
error TEXT,
PRIMARY KEY (workspace_id, blob_key)
);
CREATE INDEX IF NOT EXISTS blob_cleanup_candidates_run_idx
ON blob_cleanup_candidates (run_id, status);
@@ -0,0 +1,358 @@
use std::{
io::{BufReader, Cursor},
path::PathBuf,
time::SystemTime,
};
use assetpack_core::{
Codec, FileHint, FileTransformConfig, Hash32, ObjectKind, Pipeline, PipelineConfig, SqliteStore, TransformRegistry,
TransformSelector, build_recipe, pack::ObjectRecord, parse_recipe_checked,
};
use sqlx::Row;
use super::{
FsStorageConfig, MAX_BLOB_SIZE, ObjectGetResult, ObjectListEntry, ObjectMetadata, ObjectPutMetadata, RuntimeError,
RuntimeResult, fs_bucket_path, normalize_storage_key, system_time_ms,
};
pub(super) async fn put(
config: &FsStorageConfig,
scope: &str,
key: &str,
body: Vec<u8>,
metadata: ObjectPutMetadata,
) -> RuntimeResult<ObjectMetadata> {
normalize_storage_key(key)?;
let metadata = metadata.complete_for_body(&body);
let content_length = metadata.content_length.unwrap_or(body.len() as i64);
if content_length != body.len() as i64 {
return Err(RuntimeError::invalid_input(
"Assetpack contentLength does not match body length",
));
}
if !(0..=MAX_BLOB_SIZE).contains(&content_length) {
return Err(RuntimeError::invalid_input(
"Assetpack contentLength exceeds supported blob size",
));
}
let store = open_store(config).await?;
let transform_config = FileTransformConfig::default();
let bucket_path = fs_bucket_path(config);
let selector = TransformSelector::new(
transform_config.clone(),
transform_config.resolved_temp_dir(&bucket_path),
assetpack_transform_precomp2::default_specs(),
);
let original_hash = Hash32::sha3_256(&body);
let hint = FileHint {
size: body.len() as u64,
extension: extension_from_key(key),
head: Some(body.iter().take(4096).copied().collect()),
};
let plan = Pipeline::new(PipelineConfig::default())
.run(body, &hint, original_hash, Some(&selector))
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack pipeline failed: {err}")))?;
let chunks = plan
.chunks
.iter()
.map(|chunk| (chunk.hash, chunk.raw_len))
.collect::<Vec<_>>();
let recipe = build_recipe(
plan.original_size,
&chunks,
plan.original_hash,
plan.transform_id,
plan.transform_version,
);
let recipe_hash = Hash32::sha3_256(&recipe);
let mut objects = Vec::with_capacity(plan.chunks.len() + 1);
for chunk in plan.chunks {
let Some(payload) = chunk.payload else {
return Err(RuntimeError::invalid_state(
"Assetpack pipeline unexpectedly discarded chunk payload",
));
};
objects.push(ObjectRecord {
hash: chunk.hash,
kind: ObjectKind::Chunk,
size: chunk.raw_len as u64,
codec: chunk.codec,
content: payload,
});
}
objects.push(ObjectRecord {
hash: recipe_hash,
kind: ObjectKind::Recipe,
size: recipe.len() as u64,
codec: Codec::Raw,
content: recipe,
});
let mut tx = store
.begin_write_tx()
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack begin write failed: {err}")))?;
store
.put_objects_batch_tx(&mut tx, &objects)
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack object write failed: {err}")))?;
store
.put_file_recipe_cache_batch_tx(&mut tx, &[(original_hash, recipe_hash)])
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack recipe cache write failed: {err}")))?;
let object_metadata = ObjectMetadata {
content_type: metadata
.content_type
.unwrap_or_else(|| "application/octet-stream".to_string()),
content_length,
last_modified_ms: system_time_ms(SystemTime::now())?,
checksum_crc32: metadata.checksum_crc32,
};
sqlx::query(
r#"
INSERT INTO storage_assetpack_blobs
(scope, key, recipe_hash, content_type, content_length, checksum_crc32, last_modified_ms)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)
ON CONFLICT (scope, key)
DO UPDATE SET
recipe_hash = excluded.recipe_hash,
content_type = excluded.content_type,
content_length = excluded.content_length,
checksum_crc32 = excluded.checksum_crc32,
last_modified_ms = excluded.last_modified_ms
"#,
)
.bind(scope)
.bind(key)
.bind(recipe_hash.to_hex())
.bind(&object_metadata.content_type)
.bind(object_metadata.content_length)
.bind(&object_metadata.checksum_crc32)
.bind(object_metadata.last_modified_ms)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("Assetpack manifest write failed", err))?;
tx.commit()
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack commit failed: {err}")))?;
Ok(object_metadata)
}
pub(super) async fn head(config: &FsStorageConfig, scope: &str, key: &str) -> RuntimeResult<Option<ObjectMetadata>> {
normalize_storage_key(key)?;
let store = open_store(config).await?;
manifest_row(&store, scope, key)
.await
.map(|row| row.map(|row| row.metadata))
}
pub(super) async fn get(config: &FsStorageConfig, scope: &str, key: &str) -> RuntimeResult<Option<ObjectGetResult>> {
normalize_storage_key(key)?;
let store = open_store(config).await?;
let Some(row) = manifest_row(&store, scope, key).await? else {
return Ok(None);
};
let recipe_hash = Hash32::from_hex(&row.recipe_hash)
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack manifest recipe hash is invalid: {err}")))?;
let Some(recipe_object) = store
.get_object(&recipe_hash)
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack recipe read failed: {err}")))?
else {
return Err(RuntimeError::invalid_state(format!(
"Assetpack recipe object is missing for {key}"
)));
};
let recipe = parse_recipe_checked(&recipe_object.content, &recipe_hash)
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack recipe parse failed: {err}")))?;
let mut stored_stream = Vec::with_capacity(recipe.stored_stream_size as usize);
for (chunk_hash, expected_len) in &recipe.chunks {
let Some(chunk) = store
.get_object(chunk_hash)
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack chunk read failed: {err}")))?
else {
return Err(RuntimeError::invalid_state(format!(
"Assetpack chunk is missing for {key}: {chunk_hash}"
)));
};
if chunk.kind != ObjectKind::Chunk || chunk.size != *expected_len as u64 {
return Err(RuntimeError::invalid_state(format!(
"Assetpack chunk metadata mismatch for {key}: {chunk_hash}"
)));
}
stored_stream.extend_from_slice(&chunk.content);
}
let body = decode_stored_stream(recipe.transform_id, stored_stream)?;
if body.len() as u64 != recipe.original_file_size || Hash32::sha3_256(&body) != recipe.original_file_hash {
return Err(RuntimeError::invalid_state(format!(
"Assetpack reconstructed body failed integrity check for {key}"
)));
}
Ok(Some(ObjectGetResult {
body,
metadata: row.metadata,
}))
}
pub(super) async fn list(
config: &FsStorageConfig,
scope: &str,
prefix: Option<String>,
) -> RuntimeResult<Vec<ObjectListEntry>> {
let prefix = prefix
.map(|prefix| super::normalize_storage_prefix(&prefix))
.transpose()?
.unwrap_or_default();
let store = open_store(config).await?;
let rows = sqlx::query(
r#"
SELECT key, content_length, last_modified_ms
FROM storage_assetpack_blobs
WHERE scope = ?1 AND key LIKE ?2 ESCAPE '\'
ORDER BY key ASC
"#,
)
.bind(scope)
.bind(format!("{}%", escape_sqlite_like(&prefix)))
.fetch_all(store.pool())
.await
.map_err(|err| RuntimeError::database("Assetpack manifest list failed", err))?;
rows
.into_iter()
.map(|row| {
Ok(ObjectListEntry {
key: row.get("key"),
content_length: row.get::<i64, _>("content_length"),
last_modified_ms: row.get("last_modified_ms"),
})
})
.collect()
}
fn escape_sqlite_like(value: &str) -> String {
let mut escaped = String::with_capacity(value.len());
for ch in value.chars() {
match ch {
'%' | '_' | '\\' => {
escaped.push('\\');
escaped.push(ch);
}
_ => escaped.push(ch),
}
}
escaped
}
pub(super) async fn delete(config: &FsStorageConfig, scope: &str, key: &str) -> RuntimeResult<()> {
normalize_storage_key(key)?;
let store = open_store(config).await?;
sqlx::query("DELETE FROM storage_assetpack_blobs WHERE scope = ?1 AND key = ?2")
.bind(scope)
.bind(key)
.execute(store.pool())
.await
.map_err(|err| RuntimeError::database("Assetpack manifest delete failed", err))?;
Ok(())
}
async fn open_store(config: &FsStorageConfig) -> RuntimeResult<SqliteStore> {
let store = SqliteStore::open(store_path(config))
.await
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack store open failed: {err}")))?;
ensure_manifest_schema(&store).await?;
Ok(store)
}
fn store_path(config: &FsStorageConfig) -> PathBuf {
fs_bucket_path(config).join("assetpack.sqlite")
}
fn extension_from_key(key: &str) -> Option<String> {
key
.rsplit_once('.')
.and_then(|(_, extension)| (!extension.is_empty()).then(|| extension.to_ascii_lowercase()))
}
struct ManifestRow {
recipe_hash: String,
metadata: ObjectMetadata,
}
async fn ensure_manifest_schema(store: &SqliteStore) -> RuntimeResult<()> {
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS storage_assetpack_blobs (
scope TEXT NOT NULL,
key TEXT NOT NULL,
recipe_hash TEXT NOT NULL,
content_type TEXT NOT NULL,
content_length INTEGER NOT NULL,
checksum_crc32 TEXT,
last_modified_ms INTEGER NOT NULL,
PRIMARY KEY (scope, key)
)
"#,
)
.execute(store.pool())
.await
.map_err(|err| RuntimeError::database("Assetpack manifest schema create failed", err))?;
sqlx::query(
"CREATE INDEX IF NOT EXISTS storage_assetpack_blobs_scope_prefix_idx ON storage_assetpack_blobs (scope, key)",
)
.execute(store.pool())
.await
.map_err(|err| RuntimeError::database("Assetpack manifest index create failed", err))?;
Ok(())
}
async fn manifest_row(store: &SqliteStore, scope: &str, key: &str) -> RuntimeResult<Option<ManifestRow>> {
let row = sqlx::query(
r#"
SELECT recipe_hash, content_type, content_length, checksum_crc32, last_modified_ms
FROM storage_assetpack_blobs
WHERE scope = ?1 AND key = ?2
"#,
)
.bind(scope)
.bind(key)
.fetch_optional(store.pool())
.await
.map_err(|err| RuntimeError::database("Assetpack manifest read failed", err))?;
row
.map(|row| {
Ok(ManifestRow {
recipe_hash: row.get("recipe_hash"),
metadata: ObjectMetadata {
content_type: row.get("content_type"),
content_length: row.get("content_length"),
checksum_crc32: row.get("checksum_crc32"),
last_modified_ms: row.get("last_modified_ms"),
},
})
})
.transpose()
}
fn decode_stored_stream(transform_id: u16, stored_stream: Vec<u8>) -> RuntimeResult<Vec<u8>> {
let transform_config = FileTransformConfig::default();
let registry = TransformRegistry::new(&transform_config, assetpack_transform_precomp2::default_specs());
let transform = registry
.get(transform_id)
.ok_or_else(|| RuntimeError::invalid_state(format!("Assetpack transform is not registered: {transform_id}")))?;
let mut out = Vec::new();
transform
.decode(&mut BufReader::new(Cursor::new(stored_stream)), &mut out)
.map_err(|err| RuntimeError::invalid_state(format!("Assetpack transform decode failed: {err}")))?;
Ok(out)
}
@@ -0,0 +1,610 @@
use chrono::{DateTime, Duration, Utc};
use sqlx::{FromRow, PgPool};
use super::{
RuntimeBlobCleanupExecuteResult, RuntimeBlobCleanupPlanResult, RuntimeError, RuntimeResult, StorageRuntime,
napi_error,
};
#[derive(FromRow)]
struct BlobCandidateRow {
workspace_id: String,
key: String,
size: i32,
}
#[derive(FromRow)]
struct MarkedCandidateRow {
workspace_id: String,
blob_key: String,
}
fn push_workspace_once(workspace_ids: &mut Vec<String>, workspace_id: &str) {
if !workspace_ids.iter().any(|id| id == workspace_id) {
workspace_ids.push(workspace_id.to_string());
}
}
async fn checkpoint_completed(pool: &PgPool, kind: &str, scope: &str) -> RuntimeResult<bool> {
sqlx::query_scalar::<_, bool>(
"SELECT EXISTS(SELECT 1 FROM blob_reconciliation_checkpoints WHERE kind = $1 AND scope = $2 AND status = \
'completed')",
)
.bind(kind)
.bind(scope)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup checkpoint check failed", err))
}
async fn projection_is_stale(pool: &PgPool, workspace_id: &str) -> RuntimeResult<bool> {
let checkpoint_fresh = checkpoint_completed(pool, "doc_blob_refs", workspace_id).await?;
let has_stale_rows = sqlx::query_scalar::<_, bool>(
"SELECT EXISTS(SELECT 1 FROM doc_blob_refs WHERE workspace_id = $1 AND status <> 'fresh')",
)
.bind(workspace_id)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup projection freshness check failed", err))?;
Ok(!checkpoint_fresh || has_stale_rows)
}
async fn stale_projection_workspaces(pool: &PgPool, workspace_id: &str) -> RuntimeResult<Vec<String>> {
if projection_is_stale(pool, workspace_id).await? {
Ok(vec![workspace_id.to_string()])
} else {
Ok(Vec::new())
}
}
async fn metadata_backfill_is_complete(pool: &PgPool, workspace_id: &str) -> RuntimeResult<bool> {
checkpoint_completed(pool, "blob_metadata_backfill", workspace_id).await
}
async fn has_doc_ref(pool: &PgPool, workspace_id: &str, key: &str) -> RuntimeResult<bool> {
sqlx::query_scalar::<_, bool>(
"SELECT EXISTS(SELECT 1 FROM doc_blob_refs WHERE workspace_id = $1 AND blob_key = $2 AND status = 'fresh')",
)
.bind(workspace_id)
.bind(key)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup doc ref check failed", err))
}
async fn has_other_ref(pool: &PgPool, workspace_id: &str, key: &str) -> RuntimeResult<bool> {
let required_ref = sqlx::query_scalar::<_, bool>(
r#"
SELECT EXISTS(SELECT 1 FROM workspaces WHERE id = $1 AND avatar_key = $2)
OR EXISTS(SELECT 1 FROM ai_transcript_tasks WHERE workspace_id = $1 AND blob_id = $2)
OR EXISTS(SELECT 1 FROM ai_jobs WHERE workspace_id = $1 AND blob_id = $2)
OR EXISTS(
SELECT 1
FROM ai_contexts c
JOIN ai_sessions_metadata s ON s.id = c.session_id
WHERE s.workspace_id = $1
AND jsonb_path_exists(
c.config::jsonb,
'$.** ? (@ == $blobKey)',
jsonb_build_object('blobKey', to_jsonb($2::text))
)
)
"#,
)
.bind(workspace_id)
.bind(key)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup protected ref check failed", err))?;
if required_ref {
return Ok(true);
}
if table_exists(pool, "ai_workspace_files").await?
&& sqlx::query_scalar::<_, bool>(
"SELECT EXISTS(SELECT 1 FROM ai_workspace_files WHERE workspace_id = $1 AND blob_id = $2)",
)
.bind(workspace_id)
.bind(key)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup workspace file ref check failed", err))?
{
return Ok(true);
}
if table_exists(pool, "ai_workspace_blob_embeddings").await?
&& sqlx::query_scalar::<_, bool>(
"SELECT EXISTS(SELECT 1 FROM ai_workspace_blob_embeddings WHERE workspace_id = $1 AND blob_id = $2)",
)
.bind(workspace_id)
.bind(key)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup workspace blob embedding ref check failed", err))?
{
return Ok(true);
}
Ok(false)
}
async fn table_exists(pool: &PgPool, table: &str) -> RuntimeResult<bool> {
sqlx::query_scalar::<_, bool>("SELECT to_regclass($1) IS NOT NULL")
.bind(format!("public.{table}"))
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup table existence check failed", err))
}
async fn load_completed_blobs(
pool: &PgPool,
workspace_id: &str,
after_key: Option<&str>,
limit: i64,
) -> RuntimeResult<Vec<BlobCandidateRow>> {
sqlx::query_as::<_, BlobCandidateRow>(
r#"
SELECT workspace_id, key, size
FROM blobs
WHERE workspace_id = $1
AND status = 'completed'
AND deleted_at IS NULL
AND ($2::text IS NULL OR key > $2)
ORDER BY key ASC
LIMIT $3
"#,
)
.bind(workspace_id)
.bind(after_key)
.bind(limit)
.fetch_all(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup load completed blobs failed", err))
}
async fn load_plan_cursor(pool: &PgPool, workspace_id: &str) -> RuntimeResult<Option<String>> {
let row = sqlx::query_as::<_, (String, serde_json::Value)>(
"SELECT status, cursor FROM blob_reconciliation_checkpoints WHERE kind = 'blob_cleanup_plan' AND scope = $1",
)
.bind(workspace_id)
.fetch_optional(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup plan checkpoint load failed", err))?;
let Some((status, cursor)) = row else {
return Ok(None);
};
if status == "completed" {
return Ok(None);
}
Ok({
cursor
.get("lastBlobKey")
.and_then(|value| value.as_str())
.map(ToString::to_string)
})
}
async fn upsert_plan_checkpoint(
pool: &PgPool,
workspace_id: &str,
last_blob_key: Option<&str>,
completed: bool,
) -> RuntimeResult<()> {
let status = if completed { "completed" } else { "running" };
sqlx::query(
r#"
INSERT INTO blob_reconciliation_checkpoints
(kind, scope, status, cursor, last_key, completed_at)
VALUES ('blob_cleanup_plan', $1, $2, $3, $4, CASE WHEN $5 THEN CURRENT_TIMESTAMP ELSE NULL END)
ON CONFLICT (kind, scope) DO UPDATE
SET status = EXCLUDED.status,
cursor = EXCLUDED.cursor,
last_key = COALESCE(EXCLUDED.last_key, blob_reconciliation_checkpoints.last_key),
completed_at = CASE WHEN $5 THEN CURRENT_TIMESTAMP ELSE NULL END,
updated_at = CURRENT_TIMESTAMP
"#,
)
.bind(workspace_id)
.bind(status)
.bind(serde_json::json!({ "lastBlobKey": last_blob_key }))
.bind(last_blob_key)
.bind(completed)
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup plan checkpoint write failed", err))?;
Ok(())
}
async fn create_run(pool: &PgPool, workspace_id: &str) -> RuntimeResult<String> {
sqlx::query_scalar::<_, String>(
r#"
INSERT INTO blob_reconciliation_runs (kind, mode, status, workspace_id)
VALUES ('blob_cleanup_plan', 'mark_only', 'running', $1)
RETURNING id::text
"#,
)
.bind(workspace_id)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup create run failed", err))
}
async fn finish_run(
pool: &PgPool,
run_id: &str,
workspace_id: &str,
result: &RuntimeBlobCleanupPlanResult,
stale_projection_workspaces: Vec<String>,
) -> RuntimeResult<()> {
let candidate_bytes = sqlx::query_scalar::<_, Option<i64>>(
"SELECT SUM(object_size)::bigint FROM blob_cleanup_candidates WHERE run_id = $1::uuid AND status = 'marked'",
)
.bind(run_id)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup candidate bytes audit failed", err))?
.unwrap_or(0);
sqlx::query(
r#"
UPDATE blob_reconciliation_runs
SET status = 'finished',
finished_at = CURRENT_TIMESTAMP,
scanned = $2,
changed = $3,
metadata = $4
WHERE id = $1::uuid
"#,
)
.bind(run_id)
.bind(result.scanned_blobs as i32)
.bind(result.candidates_marked as i32)
.bind(serde_json::json!({
"protectedByDocRefs": result.protected_by_doc_refs,
"protectedByMetadata": result.protected_by_metadata,
"protectedByOtherRefs": result.protected_by_other_refs,
"topWorkspaceCandidateBytes": [{
"workspaceId": workspace_id,
"candidateBytes": candidate_bytes,
}],
"staleOrFailedProjectionWorkspaces": stale_projection_workspaces,
}))
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup finish run failed", err))?;
Ok(())
}
async fn mark_candidate_status(
pool: &PgPool,
run_id: &str,
workspace_id: &str,
blob_key: &str,
status: &str,
evidence: serde_json::Value,
error: Option<&str>,
) -> RuntimeResult<()> {
sqlx::query(
r#"
UPDATE blob_cleanup_candidates
SET status = $3,
executed_at = CURRENT_TIMESTAMP,
evidence = evidence || $4,
error = $5
WHERE workspace_id = $1 AND blob_key = $2 AND run_id = $6::uuid
"#,
)
.bind(workspace_id)
.bind(blob_key)
.bind(status)
.bind(evidence)
.bind(error)
.bind(run_id)
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup mark candidate status failed", err))?;
Ok(())
}
async fn finish_execute_run(
pool: &PgPool,
run_id: &str,
result: &RuntimeBlobCleanupExecuteResult,
) -> RuntimeResult<()> {
sqlx::query(
r#"
UPDATE blob_reconciliation_runs
SET status = 'finished',
finished_at = CURRENT_TIMESTAMP,
scanned = $2,
changed = $3,
failed = $4,
metadata = metadata || $5
WHERE id = $1::uuid
"#,
)
.bind(run_id)
.bind(result.scanned_candidates as i32)
.bind(result.deleted_metadata as i32)
.bind(result.failed as i32)
.bind(serde_json::json!({
"deletedObjects": result.deleted_objects,
"deletedMetadata": result.deleted_metadata,
"skippedStillReferenced": result.skipped_still_referenced,
"failed": result.failed,
}))
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup execute run finish failed", err))?;
Ok(())
}
async fn mark_candidate(
pool: &PgPool,
run_id: &str,
row: &BlobCandidateRow,
object_size: i64,
object_last_modified: DateTime<Utc>,
) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
INSERT INTO blob_cleanup_candidates
(workspace_id, blob_key, reason, status, object_size, object_last_modified, run_id, evidence)
VALUES ($1, $2, 'unreferenced_completed_blob', 'marked', $3, $4, $5::uuid, $6)
ON CONFLICT (workspace_id, blob_key) DO UPDATE
SET reason = EXCLUDED.reason,
status = 'marked',
object_size = EXCLUDED.object_size,
object_last_modified = EXCLUDED.object_last_modified,
planned_at = CURRENT_TIMESTAMP,
run_id = EXCLUDED.run_id,
evidence = EXCLUDED.evidence,
error = NULL
"#,
)
.bind(&row.workspace_id)
.bind(&row.key)
.bind(object_size)
.bind(object_last_modified)
.bind(run_id)
.bind(serde_json::json!({ "metadataSize": row.size }))
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup mark candidate failed", err))?;
Ok(result.rows_affected() as i64)
}
async fn load_marked_candidates(pool: &PgPool, run_id: &str, limit: i64) -> RuntimeResult<Vec<MarkedCandidateRow>> {
sqlx::query_as::<_, MarkedCandidateRow>(
r#"
SELECT workspace_id, blob_key
FROM blob_cleanup_candidates
WHERE run_id = $1::uuid AND status IN ('marked', 'failed')
ORDER BY CASE WHEN status = 'marked' THEN 0 ELSE 1 END, planned_at ASC
LIMIT $2
"#,
)
.bind(run_id)
.bind(limit)
.fetch_all(pool)
.await
.map_err(|err| RuntimeError::database("Blob cleanup load marked candidates failed", err))
}
#[napi_derive::napi]
impl StorageRuntime {
#[napi]
pub async fn plan_unreferenced_workspace_blobs(
&self,
workspace_id: String,
grace_period_days: i64,
limit: i64,
) -> napi::Result<RuntimeBlobCleanupPlanResult> {
if limit <= 0 {
return Err(napi_error("blob cleanup plan limit must be positive"));
}
if grace_period_days < 0 {
return Err(napi_error("blob cleanup grace period must be non-negative"));
}
let pool = self.pool().await?;
let run_id = create_run(&pool, &workspace_id).await?;
let mut result = RuntimeBlobCleanupPlanResult {
run_id: Some(run_id.clone()),
scanned_blobs: 0,
candidates_marked: 0,
protected_by_doc_refs: 0,
protected_by_metadata: 0,
protected_by_other_refs: 0,
next_cursor: None,
};
let cursor = load_plan_cursor(&pool, &workspace_id).await?;
let stale_projection_workspaces = stale_projection_workspaces(&pool, &workspace_id).await?;
if !metadata_backfill_is_complete(&pool, &workspace_id).await? || !stale_projection_workspaces.is_empty() {
result.protected_by_metadata = load_completed_blobs(&pool, &workspace_id, cursor.as_deref(), limit)
.await?
.len() as i64;
finish_run(&pool, &run_id, &workspace_id, &result, stale_projection_workspaces).await?;
return Ok(result);
}
let min_last_modified = Utc::now() - Duration::days(grace_period_days);
let rows = load_completed_blobs(&pool, &workspace_id, cursor.as_deref(), limit).await?;
let has_more = rows.len() == limit as usize;
let mut last_blob_key = None;
for row in rows {
result.scanned_blobs += 1;
last_blob_key = Some(row.key.clone());
if has_doc_ref(&pool, &row.workspace_id, &row.key).await? {
result.protected_by_doc_refs += 1;
continue;
}
if has_other_ref(&pool, &row.workspace_id, &row.key).await? {
result.protected_by_other_refs += 1;
continue;
}
let object_key = format!("{}/{}", row.workspace_id, row.key);
let Some(metadata) = self.object_storage_head(object_key).await? else {
result.protected_by_metadata += 1;
continue;
};
let last_modified = DateTime::<Utc>::from_timestamp_millis(metadata.last_modified_ms)
.ok_or_else(|| RuntimeError::invalid_state("blob cleanup object last modified is invalid"))?;
if metadata.content_length != row.size as i64 || last_modified > min_last_modified {
result.protected_by_metadata += 1;
continue;
}
result.candidates_marked += mark_candidate(&pool, &run_id, &row, metadata.content_length, last_modified).await?;
}
if has_more {
result.next_cursor = last_blob_key.clone();
}
upsert_plan_checkpoint(&pool, &workspace_id, last_blob_key.as_deref(), !has_more).await?;
finish_run(&pool, &run_id, &workspace_id, &result, Vec::new()).await?;
Ok(result)
}
#[napi]
pub async fn execute_blob_cleanup_candidates(
&self,
run_id: String,
grace_period_days: i64,
limit: i64,
) -> napi::Result<RuntimeBlobCleanupExecuteResult> {
if limit <= 0 {
return Err(napi_error("blob cleanup execute limit must be positive"));
}
if grace_period_days < 0 {
return Err(napi_error("blob cleanup grace period must be non-negative"));
}
let pool = self.pool().await?;
let min_last_modified = Utc::now() - Duration::days(grace_period_days);
let rows = load_marked_candidates(&pool, &run_id, limit).await?;
let mut result = RuntimeBlobCleanupExecuteResult {
scanned_candidates: rows.len() as i64,
deleted_objects: 0,
deleted_metadata: 0,
skipped_still_referenced: 0,
failed: 0,
workspace_ids: Vec::new(),
};
for row in rows {
if projection_is_stale(&pool, &row.workspace_id).await?
|| has_doc_ref(&pool, &row.workspace_id, &row.blob_key).await?
|| has_other_ref(&pool, &row.workspace_id, &row.blob_key).await?
{
result.skipped_still_referenced += 1;
mark_candidate_status(
&pool,
&run_id,
&row.workspace_id,
&row.blob_key,
"skipped",
serde_json::json!({ "skipReason": "referenced_or_projection_stale" }),
None,
)
.await?;
continue;
}
let object_key = format!("{}/{}", row.workspace_id, row.blob_key);
let mut object_was_missing = false;
let metadata = match self.object_storage_head(object_key.clone()).await {
Ok(metadata) => metadata,
Err(err) => {
result.failed += 1;
mark_candidate_status(
&pool,
&run_id,
&row.workspace_id,
&row.blob_key,
"failed",
serde_json::json!({ "failure": "object_head_failed" }),
Some(&err.to_string()),
)
.await?;
continue;
}
};
if let Some(metadata) = metadata {
let last_modified = DateTime::<Utc>::from_timestamp_millis(metadata.last_modified_ms)
.ok_or_else(|| RuntimeError::invalid_state("blob cleanup execute object last modified is invalid"))?;
if last_modified > min_last_modified {
result.skipped_still_referenced += 1;
mark_candidate_status(
&pool,
&run_id,
&row.workspace_id,
&row.blob_key,
"skipped",
serde_json::json!({ "skipReason": "object_inside_grace_period" }),
None,
)
.await?;
continue;
}
if let Err(err) = self.object_storage_delete(object_key).await {
result.failed += 1;
mark_candidate_status(
&pool,
&run_id,
&row.workspace_id,
&row.blob_key,
"failed",
serde_json::json!({ "failure": "object_delete_failed" }),
Some(&err.to_string()),
)
.await?;
continue;
}
result.deleted_objects += 1;
} else {
object_was_missing = true;
}
let deleted_metadata =
match sqlx::query("DELETE FROM blobs WHERE workspace_id = $1 AND key = $2 AND deleted_at IS NULL")
.bind(&row.workspace_id)
.bind(&row.blob_key)
.execute(&pool)
.await
{
Ok(result) => result.rows_affected() as i64,
Err(err) => {
result.failed += 1;
mark_candidate_status(
&pool,
&run_id,
&row.workspace_id,
&row.blob_key,
"failed",
serde_json::json!({ "failure": "metadata_delete_failed" }),
Some(&err.to_string()),
)
.await?;
continue;
}
};
result.deleted_metadata += deleted_metadata;
push_workspace_once(&mut result.workspace_ids, &row.workspace_id);
mark_candidate_status(
&pool,
&run_id,
&row.workspace_id,
&row.blob_key,
"executed",
serde_json::json!({
"deletedMetadata": deleted_metadata,
"objectMissingBeforeDelete": object_was_missing,
}),
None,
)
.await?;
}
finish_execute_run(&pool, &run_id, &result).await?;
Ok(result)
}
}
@@ -0,0 +1,182 @@
use chrono::{DateTime, Utc};
use napi::Result;
use sqlx::{FromRow, PgPool};
use super::{RuntimeBlobCleanupResult, RuntimeError, RuntimeResult, StorageRuntime, napi_error};
#[derive(FromRow)]
struct BlobRow {
workspace_id: String,
key: String,
upload_id: Option<String>,
}
struct BlobReclaimerStore {
pool: PgPool,
}
impl BlobReclaimerStore {
fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn load_expired_pending(&self, cutoff: DateTime<Utc>, limit: i64) -> RuntimeResult<Vec<BlobRow>> {
sqlx::query_as::<_, BlobRow>(
r#"
SELECT workspace_id, key, upload_id
FROM blobs
WHERE status = 'pending'
AND deleted_at IS NULL
AND created_at < $1
ORDER BY created_at ASC
LIMIT $2
"#,
)
.bind(cutoff)
.bind(limit)
.fetch_all(&self.pool)
.await
.map_err(|err| RuntimeError::database("BlobReclaimer load pending blobs failed", err))
}
async fn load_deleted(&self, workspace_id: &str, limit: i64) -> RuntimeResult<Vec<BlobRow>> {
sqlx::query_as::<_, BlobRow>(
r#"
SELECT workspace_id, key, upload_id
FROM blobs
WHERE workspace_id = $1
AND deleted_at IS NOT NULL
ORDER BY deleted_at ASC
LIMIT $2
"#,
)
.bind(workspace_id)
.bind(limit)
.fetch_all(&self.pool)
.await
.map_err(|err| RuntimeError::database("BlobReclaimer load deleted blobs failed", err))
}
async fn delete_pending_metadata(&self, workspace_id: &str, key: &str) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM blobs
WHERE workspace_id = $1 AND key = $2
AND status = 'pending'
AND deleted_at IS NULL
"#,
)
.bind(workspace_id)
.bind(key)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("BlobReclaimer delete pending blob metadata failed", err))?;
Ok(result.rows_affected() as i64)
}
async fn delete_released_metadata(&self, workspace_id: &str, key: &str) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM blobs
WHERE workspace_id = $1 AND key = $2
AND deleted_at IS NOT NULL
"#,
)
.bind(workspace_id)
.bind(key)
.execute(&self.pool)
.await
.map_err(|err| RuntimeError::database("BlobReclaimer delete blob metadata failed", err))?;
Ok(result.rows_affected() as i64)
}
}
async fn delete_object_idempotent(runtime: &StorageRuntime, key: &str) -> RuntimeResult<()> {
match runtime.object_storage_delete_object(key).await {
Ok(()) => Ok(()),
Err(err) if err.is_object_missing() => Ok(()),
Err(err) => Err(err),
}
}
async fn abort_upload_idempotent(runtime: &StorageRuntime, key: &str, upload_id: &str) -> RuntimeResult<()> {
match runtime.object_storage_abort_upload(key, upload_id).await {
Ok(()) => Ok(()),
Err(err) if err.is_object_missing() => Ok(()),
Err(err) => Err(err),
}
}
fn push_workspace_once(workspace_ids: &mut Vec<String>, workspace_id: &str) {
if !workspace_ids.iter().any(|id| id == workspace_id) {
workspace_ids.push(workspace_id.to_string());
}
}
#[napi_derive::napi]
impl StorageRuntime {
#[napi]
pub async fn cleanup_expired_pending_blobs(&self, cutoff_ms: i64, limit: i64) -> Result<RuntimeBlobCleanupResult> {
if limit <= 0 {
return Err(napi_error("pending blob cleanup limit must be positive"));
}
let cutoff = DateTime::<Utc>::from_timestamp_millis(cutoff_ms)
.ok_or_else(|| RuntimeError::invalid_input("pending blob cleanup cutoff is invalid"))?;
let store = BlobReclaimerStore::new(self.pool().await?);
let rows = store.load_expired_pending(cutoff, limit).await?;
let mut deleted = 0;
let mut aborted_multipart = 0;
let mut workspace_ids = Vec::new();
for row in &rows {
let object_key = format!("{}/{}", row.workspace_id, row.key);
if let Some(upload_id) = row.upload_id.as_deref() {
abort_upload_idempotent(self, &object_key, upload_id).await?;
aborted_multipart += 1;
}
delete_object_idempotent(self, &object_key).await?;
let affected = store.delete_pending_metadata(&row.workspace_id, &row.key).await?;
if affected > 0 {
deleted += affected;
push_workspace_once(&mut workspace_ids, &row.workspace_id);
}
}
Ok(RuntimeBlobCleanupResult {
scanned: rows.len() as i64,
deleted,
aborted_multipart,
workspace_ids,
})
}
#[napi]
pub async fn release_deleted_blobs(&self, workspace_id: String, limit: i64) -> Result<RuntimeBlobCleanupResult> {
if limit <= 0 {
return Err(napi_error("deleted blob release limit must be positive"));
}
let store = BlobReclaimerStore::new(self.pool().await?);
let rows = store.load_deleted(&workspace_id, limit).await?;
let mut deleted = 0;
let mut workspace_ids = Vec::new();
for row in &rows {
let object_key = format!("{}/{}", row.workspace_id, row.key);
delete_object_idempotent(self, &object_key).await?;
let affected = store.delete_released_metadata(&row.workspace_id, &row.key).await?;
if affected > 0 {
deleted += affected;
push_workspace_once(&mut workspace_ids, &row.workspace_id);
}
}
Ok(RuntimeBlobCleanupResult {
scanned: rows.len() as i64,
deleted,
aborted_multipart: 0,
workspace_ids,
})
}
}
@@ -0,0 +1,273 @@
use chrono::{DateTime, Utc};
use sqlx::{FromRow, PgPool};
use super::{
RuntimeBlobMetadataBackfillResult, RuntimeError, RuntimeObjectMetadata, RuntimeResult, StorageRuntime, napi_error,
};
async fn workspace_exists(pool: &PgPool, workspace_id: &str) -> RuntimeResult<bool> {
sqlx::query_scalar::<_, bool>("SELECT EXISTS(SELECT 1 FROM workspaces WHERE id = $1)")
.bind(workspace_id)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob metadata backfill workspace check failed", err))
}
async fn blob_exists(pool: &PgPool, workspace_id: &str, key: &str) -> RuntimeResult<bool> {
sqlx::query_scalar::<_, bool>("SELECT EXISTS(SELECT 1 FROM blobs WHERE workspace_id = $1 AND key = $2)")
.bind(workspace_id)
.bind(key)
.fetch_one(pool)
.await
.map_err(|err| RuntimeError::database("Blob metadata backfill blob check failed", err))
}
async fn upsert_blob_metadata(
pool: &PgPool,
workspace_id: &str,
key: &str,
metadata: RuntimeObjectMetadata,
) -> RuntimeResult<i64> {
let last_modified = DateTime::<Utc>::from_timestamp_millis(metadata.last_modified_ms)
.ok_or_else(|| RuntimeError::invalid_state("Blob metadata backfill object last modified is invalid"))?;
let result = sqlx::query(
r#"
INSERT INTO blobs (workspace_id, key, size, mime, status, upload_id, created_at, deleted_at)
VALUES ($1, $2, $3, $4, 'completed', NULL, $5, NULL)
ON CONFLICT (workspace_id, key) DO UPDATE
SET size = EXCLUDED.size,
mime = EXCLUDED.mime,
status = 'completed',
upload_id = NULL,
deleted_at = NULL
WHERE blobs.deleted_at IS NULL
"#,
)
.bind(workspace_id)
.bind(key)
.bind(metadata.content_length as i32)
.bind(metadata.content_type)
.bind(last_modified)
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob metadata backfill upsert failed", err))?;
Ok(result.rows_affected() as i64)
}
fn split_workspace_blob_key(full_key: &str) -> Option<(&str, &str)> {
let (workspace_id, key) = full_key.split_once('/')?;
if workspace_id.is_empty() || key.is_empty() || key.contains('/') {
return None;
}
Some((workspace_id, key))
}
fn checkpoint_scope(workspace_id: Option<&str>) -> String {
workspace_id.unwrap_or("__all__").to_string()
}
#[derive(FromRow)]
struct BackfillCheckpoint {
last_key: Option<String>,
cursor: serde_json::Value,
}
impl BackfillCheckpoint {
fn continuation_token(&self) -> Option<String> {
self
.cursor
.get("continuationToken")
.and_then(|value| value.as_str())
.map(ToString::to_string)
}
}
async fn load_checkpoint(pool: &PgPool, scope: &str) -> RuntimeResult<Option<BackfillCheckpoint>> {
sqlx::query_as::<_, BackfillCheckpoint>(
"SELECT last_key, cursor FROM blob_reconciliation_checkpoints WHERE kind = 'blob_metadata_backfill' AND scope = $1",
)
.bind(scope)
.fetch_optional(pool)
.await
.map_err(|err| RuntimeError::database("Blob metadata backfill checkpoint load failed", err))
}
async fn upsert_checkpoint(
pool: &PgPool,
scope: &str,
last_key: Option<&str>,
continuation_token: Option<&str>,
completed: bool,
) -> RuntimeResult<()> {
let status = if completed { "completed" } else { "running" };
sqlx::query(
r#"
INSERT INTO blob_reconciliation_checkpoints
(kind, scope, status, cursor, last_key, completed_at, metadata)
VALUES ('blob_metadata_backfill', $1, $2, $3, $4, CASE WHEN $5 THEN CURRENT_TIMESTAMP ELSE NULL END, $6)
ON CONFLICT (kind, scope) DO UPDATE
SET status = EXCLUDED.status,
cursor = EXCLUDED.cursor,
last_key = COALESCE(EXCLUDED.last_key, blob_reconciliation_checkpoints.last_key),
completed_at = CASE WHEN $5 THEN CURRENT_TIMESTAMP ELSE NULL END,
updated_at = CURRENT_TIMESTAMP,
metadata = EXCLUDED.metadata
"#,
)
.bind(scope)
.bind(status)
.bind(serde_json::json!({
"lastKey": last_key,
"continuationToken": continuation_token,
}))
.bind(last_key)
.bind(completed)
.bind(serde_json::json!({
"quotaReportingReconciliationRequired": true,
}))
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Blob metadata backfill checkpoint write failed", err))?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn blob_metadata_backfill_splits_workspace_blob_keys() {
assert_eq!(
split_workspace_blob_key("workspace/blob-key"),
Some(("workspace", "blob-key"))
);
assert_eq!(split_workspace_blob_key("workspace/nested/blob-key"), None);
assert_eq!(split_workspace_blob_key("workspace/"), None);
assert_eq!(split_workspace_blob_key("blob-key"), None);
}
#[test]
fn blob_metadata_backfill_checkpoint_scope_is_explicit() {
assert_eq!(checkpoint_scope(Some("workspace")), "workspace");
assert_eq!(checkpoint_scope(None), "__all__");
}
}
fn push_workspace_once(workspace_ids: &mut Vec<String>, workspace_id: &str) {
if !workspace_ids.iter().any(|id| id == workspace_id) {
workspace_ids.push(workspace_id.to_string());
}
}
fn checked_list_page_limit(limit: i64) -> RuntimeResult<i32> {
i32::try_from(limit).map_err(|_| RuntimeError::invalid_input("blob metadata backfill limit exceeds i32::MAX"))
}
#[napi_derive::napi]
impl StorageRuntime {
#[napi]
pub async fn backfill_missing_blob_metadata(
&self,
workspace_id: Option<String>,
limit: i64,
) -> napi::Result<RuntimeBlobMetadataBackfillResult> {
if limit <= 0 {
return Err(napi_error("blob metadata backfill limit must be positive"));
}
let page_limit = checked_list_page_limit(limit)?;
let pool = self.pool().await?;
let prefix = workspace_id.as_ref().map(|id| format!("{id}/"));
let scope = checkpoint_scope(workspace_id.as_deref());
let checkpoint = load_checkpoint(&pool, &scope).await?;
let page = self
.object_storage_list_page(
prefix,
checkpoint.as_ref().and_then(BackfillCheckpoint::continuation_token),
checkpoint.as_ref().and_then(|checkpoint| checkpoint.last_key.clone()),
page_limit,
)
.await?;
let has_more = page.next_continuation_token.is_some();
let mut result = RuntimeBlobMetadataBackfillResult {
scanned_objects: 0,
headed_objects: 0,
upserted_metadata: 0,
skipped_existing: 0,
skipped_workspace_missing: 0,
failed: 0,
next_cursor: None,
workspace_ids: Vec::new(),
};
let mut last_scanned_key = None;
for object in &page.entries {
result.scanned_objects += 1;
last_scanned_key = Some(object.key.clone());
let Some((object_workspace_id, key)) = split_workspace_blob_key(&object.key) else {
result.failed += 1;
continue;
};
if workspace_id.as_deref().is_some_and(|id| id != object_workspace_id) {
result.failed += 1;
continue;
}
if !workspace_exists(&pool, object_workspace_id).await? {
result.skipped_workspace_missing += 1;
continue;
}
if blob_exists(&pool, object_workspace_id, key).await? {
result.skipped_existing += 1;
continue;
}
result.headed_objects += 1;
let Some(metadata) = self.object_storage_head(object.key.clone()).await? else {
result.failed += 1;
continue;
};
let affected = upsert_blob_metadata(&pool, object_workspace_id, key, metadata).await?;
if affected > 0 {
result.upserted_metadata += affected;
push_workspace_once(&mut result.workspace_ids, object_workspace_id);
}
}
if has_more {
result.next_cursor = last_scanned_key.clone();
}
upsert_checkpoint(
&pool,
&scope,
last_scanned_key.as_deref(),
page.next_continuation_token.as_deref(),
!has_more,
)
.await?;
sqlx::query(
r#"
INSERT INTO blob_reconciliation_runs
(kind, mode, status, workspace_id, finished_at, scanned, changed, failed, metadata)
VALUES ('blob_metadata_backfill', 'execute', 'finished', $1, CURRENT_TIMESTAMP, $2, $3, $4, $5)
"#,
)
.bind(workspace_id)
.bind(result.scanned_objects as i32)
.bind(result.upserted_metadata as i32)
.bind(result.failed as i32)
.bind(serde_json::json!({
"headedObjects": result.headed_objects,
"skippedExisting": result.skipped_existing,
"skippedWorkspaceMissing": result.skipped_workspace_missing,
"checkpointScope": scope,
"nextCursor": result.next_cursor,
"quotaReportingReconciliationRequired": true,
}))
.execute(&pool)
.await
.map_err(|err| RuntimeError::database("Blob metadata backfill run record failed", err))?;
Ok(result)
}
}
@@ -0,0 +1,427 @@
use affine_common::doc_parser;
use chrono::{DateTime, Utc};
use sqlx::{FromRow, PgPool};
use y_octo::Doc;
use super::{RuntimeDocBlobRefsResult, RuntimeError, RuntimeResult, StorageRuntime, napi_error};
const PARSER_VERSION: i32 = 1;
#[derive(FromRow)]
struct SnapshotRow {
workspace_id: String,
doc_id: String,
blob: Vec<u8>,
updated_at: DateTime<Utc>,
}
#[derive(FromRow)]
struct UpdateRow {
blob: Vec<u8>,
created_at: DateTime<Utc>,
}
struct ExtractedRef {
blob_key: String,
block_id: String,
flavour: String,
}
async fn load_snapshot(pool: &PgPool, workspace_id: &str, doc_id: &str) -> RuntimeResult<Option<SnapshotRow>> {
sqlx::query_as::<_, SnapshotRow>(
r#"
SELECT workspace_id, guid AS doc_id, blob, updated_at
FROM snapshots
WHERE workspace_id = $1 AND guid = $2
"#,
)
.bind(workspace_id)
.bind(doc_id)
.fetch_optional(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs load snapshot failed", err))
}
async fn load_updates(pool: &PgPool, workspace_id: &str, doc_id: &str) -> RuntimeResult<Vec<UpdateRow>> {
sqlx::query_as::<_, UpdateRow>(
r#"
SELECT blob, created_at
FROM updates
WHERE workspace_id = $1 AND guid = $2
ORDER BY created_at ASC
"#,
)
.bind(workspace_id)
.bind(doc_id)
.fetch_all(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs load updates failed", err))
}
fn apply_doc_updates(updates: impl IntoIterator<Item = Vec<u8>>) -> RuntimeResult<Vec<u8>> {
let mut doc = Doc::default();
for update in updates {
doc
.apply_update_from_binary_v1(&update)
.map_err(|err| RuntimeError::invalid_state(format!("Doc blob refs merge failed: {err}")))?;
}
doc
.encode_update_v1()
.map_err(|err| RuntimeError::invalid_state(format!("Doc blob refs encode failed: {err}")))
}
async fn load_current_doc(pool: &PgPool, workspace_id: &str, doc_id: &str) -> RuntimeResult<Option<SnapshotRow>> {
let snapshot = load_snapshot(pool, workspace_id, doc_id).await?;
let updates = load_updates(pool, workspace_id, doc_id).await?;
if snapshot.is_none() && updates.is_empty() {
return Ok(None);
}
let mut merge_inputs = Vec::with_capacity(updates.len() + usize::from(snapshot.is_some()));
let mut updated_at = snapshot
.as_ref()
.map(|snapshot| snapshot.updated_at)
.unwrap_or_else(Utc::now);
if let Some(snapshot) = snapshot {
merge_inputs.push(snapshot.blob);
}
for update in updates {
updated_at = update.created_at;
merge_inputs.push(update.blob);
}
Ok(Some(SnapshotRow {
workspace_id: workspace_id.to_string(),
doc_id: doc_id.to_string(),
blob: apply_doc_updates(merge_inputs)?,
updated_at,
}))
}
async fn load_workspace_doc_ids(pool: &PgPool, workspace_id: &str) -> RuntimeResult<Vec<String>> {
let Some(root) = load_current_doc(pool, workspace_id, workspace_id).await? else {
return Ok(Vec::new());
};
let ids = doc_parser::get_doc_ids_from_binary(root.blob, false)
.map_err(|err| RuntimeError::invalid_state(format!("Doc blob refs root doc parse failed: {err}")))?;
let mut ids = ids;
ids.sort();
Ok(ids)
}
async fn upsert_projection_checkpoint(
pool: &PgPool,
workspace_id: &str,
result: &RuntimeDocBlobRefsResult,
) -> RuntimeResult<()> {
let completed = result.next_cursor.is_none();
let status = if completed && result.failed_docs == 0 {
"completed"
} else if result.failed_docs > 0 {
"failed"
} else {
"running"
};
sqlx::query(
r#"
INSERT INTO blob_reconciliation_checkpoints
(kind, scope, status, cursor, completed_at, metadata)
VALUES ('doc_blob_refs', $1, $2, $3, CASE WHEN $4 THEN CURRENT_TIMESTAMP ELSE NULL END, $5)
ON CONFLICT (kind, scope) DO UPDATE
SET status = EXCLUDED.status,
cursor = EXCLUDED.cursor,
completed_at = CASE WHEN $4 THEN CURRENT_TIMESTAMP ELSE NULL END,
updated_at = CURRENT_TIMESTAMP,
metadata = EXCLUDED.metadata
"#,
)
.bind(workspace_id)
.bind(status)
.bind(serde_json::json!({ "lastDocId": result.next_cursor }))
.bind(completed && result.failed_docs == 0)
.bind(serde_json::json!({
"parserVersion": PARSER_VERSION,
}))
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs checkpoint write failed", err))?;
Ok(())
}
async fn upsert_projection_failure_checkpoint(pool: &PgPool, workspace_id: &str, error: &str) -> RuntimeResult<()> {
sqlx::query(
r#"
INSERT INTO blob_reconciliation_checkpoints
(kind, scope, status, cursor, completed_at, metadata)
VALUES ('doc_blob_refs', $1, 'failed', '{}', NULL, $2)
ON CONFLICT (kind, scope) DO UPDATE
SET status = 'failed',
cursor = '{}',
completed_at = NULL,
updated_at = CURRENT_TIMESTAMP,
metadata = EXCLUDED.metadata
"#,
)
.bind(workspace_id)
.bind(serde_json::json!({
"parserVersion": PARSER_VERSION,
"error": error,
}))
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs failure checkpoint write failed", err))?;
Ok(())
}
async fn load_projection_cursor(pool: &PgPool, workspace_id: &str) -> RuntimeResult<Option<String>> {
let cursor = sqlx::query_scalar::<_, serde_json::Value>(
"SELECT cursor FROM blob_reconciliation_checkpoints WHERE kind = 'doc_blob_refs' AND scope = $1",
)
.bind(workspace_id)
.fetch_optional(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs checkpoint load failed", err))?;
Ok(cursor.and_then(|cursor| {
cursor
.get("lastDocId")
.and_then(|value| value.as_str())
.map(ToString::to_string)
}))
}
async fn purge_removed_doc_refs(pool: &PgPool, workspace_id: &str, current_doc_ids: &[String]) -> RuntimeResult<i64> {
let result = sqlx::query(
r#"
DELETE FROM doc_blob_refs
WHERE workspace_id = $1
AND NOT (doc_id = ANY($2))
"#,
)
.bind(workspace_id)
.bind(current_doc_ids)
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs purge removed docs failed", err))?;
Ok(result.rows_affected() as i64)
}
fn extract_refs(snapshot: &SnapshotRow) -> RuntimeResult<Vec<ExtractedRef>> {
let parsed = doc_parser::parse_doc_from_binary(snapshot.blob.clone(), snapshot.doc_id.clone())
.map_err(|err| RuntimeError::invalid_state(format!("Doc blob refs parse failed: {err}")))?;
let mut refs = Vec::new();
for block in parsed.blocks {
let Some(blob_keys) = block.blob else {
continue;
};
for blob_key in blob_keys {
refs.push(ExtractedRef {
blob_key,
block_id: block.block_id.clone(),
flavour: block.flavour.clone(),
});
}
}
Ok(refs)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn doc_blob_refs_extracts_image_refs() {
let doc_id = "doc-blob-ref-test".to_string();
let blob =
doc_parser::build_full_doc("Doc", "![Alt](blob://image-blob-key)", &doc_id).expect("doc fixture should build");
let snapshot = SnapshotRow {
workspace_id: "workspace".to_string(),
doc_id,
blob,
updated_at: Utc::now(),
};
let refs = extract_refs(&snapshot).expect("refs should parse");
assert!(
refs
.iter()
.any(|reference| { reference.blob_key == "image-blob-key" && reference.flavour == "affine:image" })
);
}
}
async fn replace_doc_refs(pool: &PgPool, snapshot: &SnapshotRow, refs: Vec<ExtractedRef>) -> RuntimeResult<(i64, i64)> {
let mut tx = pool
.begin()
.await
.map_err(|err| RuntimeError::database("Doc blob refs transaction failed", err))?;
let deleted = sqlx::query("DELETE FROM doc_blob_refs WHERE workspace_id = $1 AND doc_id = $2")
.bind(&snapshot.workspace_id)
.bind(&snapshot.doc_id)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("Doc blob refs delete failed", err))?
.rows_affected() as i64;
let mut written = 0;
for reference in refs {
let affected = sqlx::query(
r#"
INSERT INTO doc_blob_refs
(workspace_id, doc_id, blob_key, block_id, flavour, snapshot_updated_at, parser_version, status)
VALUES ($1, $2, $3, $4, $5, $6, $7, 'fresh')
ON CONFLICT (workspace_id, doc_id, blob_key, block_id) DO UPDATE
SET flavour = EXCLUDED.flavour,
snapshot_updated_at = EXCLUDED.snapshot_updated_at,
indexed_at = CURRENT_TIMESTAMP,
parser_version = EXCLUDED.parser_version,
status = 'fresh',
error = NULL
"#,
)
.bind(&snapshot.workspace_id)
.bind(&snapshot.doc_id)
.bind(reference.blob_key)
.bind(reference.block_id)
.bind(reference.flavour)
.bind(snapshot.updated_at)
.bind(PARSER_VERSION)
.execute(&mut *tx)
.await
.map_err(|err| RuntimeError::database("Doc blob refs insert failed", err))?
.rows_affected() as i64;
written += affected;
}
tx.commit()
.await
.map_err(|err| RuntimeError::database("Doc blob refs transaction commit failed", err))?;
Ok((written, deleted))
}
async fn mark_doc_failed(pool: &PgPool, workspace_id: &str, doc_id: &str, error: &str) -> RuntimeResult<()> {
sqlx::query(
r#"
INSERT INTO doc_blob_refs
(workspace_id, doc_id, blob_key, block_id, flavour, snapshot_updated_at, parser_version, status, error)
VALUES ($1, $2, '__parse_failed__', '__parse_failed__', '__parse_failed__', CURRENT_TIMESTAMP, $3, 'failed', $4)
ON CONFLICT (workspace_id, doc_id, blob_key, block_id) DO UPDATE
SET indexed_at = CURRENT_TIMESTAMP,
status = 'failed',
error = EXCLUDED.error
"#,
)
.bind(workspace_id)
.bind(doc_id)
.bind(PARSER_VERSION)
.bind(error)
.execute(pool)
.await
.map_err(|err| RuntimeError::database("Doc blob refs mark failure failed", err))?;
Ok(())
}
async fn rebuild_doc_blob_refs_inner(
runtime: &StorageRuntime,
workspace_id: String,
doc_id: String,
) -> RuntimeResult<RuntimeDocBlobRefsResult> {
let pool = runtime.pool().await?;
let mut result = RuntimeDocBlobRefsResult {
scanned_docs: 1,
parsed_docs: 0,
refs_written: 0,
refs_deleted: 0,
failed_docs: 0,
next_cursor: None,
};
let Some(snapshot) = load_current_doc(&pool, &workspace_id, &doc_id).await? else {
result.failed_docs = 1;
mark_doc_failed(&pool, &workspace_id, &doc_id, "snapshot_missing").await?;
return Ok(result);
};
match extract_refs(&snapshot) {
Ok(refs) => {
let (written, deleted) = replace_doc_refs(&pool, &snapshot, refs).await?;
result.parsed_docs = 1;
result.refs_written = written;
result.refs_deleted = deleted;
}
Err(err) => {
result.failed_docs = 1;
mark_doc_failed(&pool, &workspace_id, &doc_id, &err.to_string()).await?;
}
}
Ok(result)
}
#[napi_derive::napi]
impl StorageRuntime {
#[napi]
pub async fn rebuild_doc_blob_refs(
&self,
workspace_id: String,
doc_id: String,
) -> napi::Result<RuntimeDocBlobRefsResult> {
Ok(rebuild_doc_blob_refs_inner(self, workspace_id, doc_id).await?)
}
#[napi]
pub async fn rebuild_workspace_doc_blob_refs(
&self,
workspace_id: String,
limit: i64,
) -> napi::Result<RuntimeDocBlobRefsResult> {
if limit <= 0 {
return Err(napi_error("doc blob refs rebuild limit must be positive"));
}
let pool = self.pool().await?;
let doc_ids = match load_workspace_doc_ids(&pool, &workspace_id).await {
Ok(doc_ids) => doc_ids,
Err(err) => {
upsert_projection_failure_checkpoint(&pool, &workspace_id, &err.to_string()).await?;
return Err(err.into());
}
};
let cursor = load_projection_cursor(&pool, &workspace_id).await?;
let current_doc_ids = doc_ids.clone();
let doc_ids = doc_ids
.into_iter()
.filter(|doc_id| cursor.as_ref().is_none_or(|cursor| doc_id > cursor))
.collect::<Vec<_>>();
let has_more = doc_ids.len() > limit as usize;
let mut total = RuntimeDocBlobRefsResult {
scanned_docs: 0,
parsed_docs: 0,
refs_written: 0,
refs_deleted: 0,
failed_docs: 0,
next_cursor: None,
};
let mut last_doc_id = None;
for doc_id in doc_ids.into_iter().take(limit as usize) {
last_doc_id = Some(doc_id.clone());
let result = rebuild_doc_blob_refs_inner(self, workspace_id.clone(), doc_id).await?;
total.scanned_docs += result.scanned_docs;
total.parsed_docs += result.parsed_docs;
total.refs_written += result.refs_written;
total.refs_deleted += result.refs_deleted;
total.failed_docs += result.failed_docs;
}
if has_more {
total.next_cursor = last_doc_id;
} else if total.failed_docs == 0 {
total.refs_deleted += purge_removed_doc_refs(&pool, &workspace_id, &current_doc_ids).await?;
}
upsert_projection_checkpoint(&pool, &workspace_id, &total).await?;
Ok(total)
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,850 @@
use std::{
collections::HashMap,
future::Future,
pin::Pin,
time::{Duration, SystemTime},
};
use chrono::{DateTime, FixedOffset};
use reqwest::{
Client as ReqwestClient, Method, StatusCode,
header::{CONTENT_LENGTH, CONTENT_TYPE, ETAG, HeaderMap, HeaderName, HeaderValue, LAST_MODIFIED},
};
use rusty_s3::{
Bucket, Credentials,
actions::{
AbortMultipartUpload, CompleteMultipartUpload, CreateMultipartUpload, DeleteObject, GetObject, HeadObject,
ListObjectsV2, ListParts, PutObject, S3Action, UploadPart,
},
};
use url::Url;
use super::{
error::{ObjectStorageError, ObjectStorageResult},
types::{
MultipartUploadInitResult, MultipartUploadPart, ObjectGetResult, ObjectListEntry, ObjectListPage, ObjectMetadata,
ObjectPutMetadata, PresignedObjectRequest, completed_multipart_parts, trim_etag,
},
};
const DEFAULT_REQUEST_TIMEOUT_MS: u64 = 30_000;
const MAX_MULTIPART_PART_NUMBER: i32 = 10_000;
const MAX_RESPONSE_BODY_BYTES: usize = i32::MAX as usize;
type StorageHttpFuture<'a> = Pin<Box<dyn Future<Output = ObjectStorageResult<StorageHttpResponse>> + Send + 'a>>;
#[derive(Clone)]
struct StorageHttpRequest {
method: Method,
url: Url,
headers: HashMap<String, String>,
body: Option<Vec<u8>>,
max_response_body_bytes: usize,
}
struct StorageHttpResponse {
status: StatusCode,
headers: HeaderMap,
body: Vec<u8>,
}
trait StorageHttpClient: Clone + Send + Sync + 'static {
fn execute(&self, request: StorageHttpRequest) -> StorageHttpFuture<'_>;
}
#[derive(Clone)]
struct ReqwestStorageHttpClient {
client: ReqwestClient,
}
impl ReqwestStorageHttpClient {
fn new(request_timeout_ms: Option<u64>) -> ObjectStorageResult<Self> {
let builder = ReqwestClient::builder().timeout(Duration::from_millis(
request_timeout_ms.unwrap_or(DEFAULT_REQUEST_TIMEOUT_MS),
));
Ok(Self {
client: builder.build().map_err(ObjectStorageError::HttpClientBuild)?,
})
}
}
impl StorageHttpClient for ReqwestStorageHttpClient {
fn execute(&self, request: StorageHttpRequest) -> StorageHttpFuture<'_> {
Box::pin(async move {
let mut builder = self.client.request(request.method, request.url);
for (key, value) in request.headers {
let name =
HeaderName::from_bytes(key.as_bytes()).map_err(|err| ObjectStorageError::InvalidHeader(err.to_string()))?;
let value = HeaderValue::from_str(&value).map_err(|err| ObjectStorageError::InvalidHeader(err.to_string()))?;
builder = builder.header(name, value);
}
if let Some(body) = request.body {
builder = builder.body(body);
}
let mut response = builder.send().await.map_err(ObjectStorageError::HttpRequest)?;
let status = response.status();
let headers = response.headers().clone();
if response
.content_length()
.is_some_and(|length| length > request.max_response_body_bytes as u64)
{
return Err(ObjectStorageError::BodyTooLarge {
limit: request.max_response_body_bytes,
});
}
let mut body = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(ObjectStorageError::HttpRequest)? {
if body.len() + chunk.len() > request.max_response_body_bytes {
return Err(ObjectStorageError::BodyTooLarge {
limit: request.max_response_body_bytes,
});
}
body.extend_from_slice(&chunk);
}
Ok(StorageHttpResponse { status, headers, body })
})
}
}
#[derive(Clone)]
pub(crate) struct ObjectStorageClient {
bucket: Bucket,
credentials: Credentials,
http: ReqwestStorageHttpClient,
presign_expires_in_seconds: u64,
presign_sign_content_type_for_put: bool,
}
impl ObjectStorageClient {
pub(crate) fn new(
bucket: Bucket,
credentials: Credentials,
request_timeout_ms: Option<u64>,
presign_expires_in_seconds: u64,
presign_sign_content_type_for_put: bool,
) -> ObjectStorageResult<Self> {
Ok(Self {
bucket,
credentials,
http: ReqwestStorageHttpClient::new(request_timeout_ms)?,
presign_expires_in_seconds,
presign_sign_content_type_for_put,
})
}
pub(crate) async fn put(
&self,
key: &str,
body: Vec<u8>,
metadata: ObjectPutMetadata,
) -> ObjectStorageResult<ObjectMetadata> {
let metadata = metadata.complete_for_body(&body);
let object_metadata = metadata.clone().into_object_metadata(
crate::utils::system_time_millis(SystemTime::now())
.map(|millis| millis as i64)
.map_err(|err| ObjectStorageError::InvalidInput(format!("system time before unix epoch: {err}")))?,
);
let mut headers = HashMap::from([
("content-type".to_string(), object_metadata.content_type.clone()),
("content-length".to_string(), object_metadata.content_length.to_string()),
]);
if let Some(checksum) = object_metadata.checksum_crc32.clone() {
headers.insert("x-amz-checksum-crc32".to_string(), checksum);
}
let mut action = PutObject::new(&self.bucket, Some(&self.credentials), key);
insert_action_headers(&mut action, &headers);
let response = self
.http
.execute(StorageHttpRequest {
method: Method::PUT,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers,
body: Some(body),
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage put failed for {key}"), source))?;
ensure_success_status(&response, &format!("ObjectStorage put failed for {key}"))?;
Ok(object_metadata)
}
pub(crate) async fn presign_put(
&self,
key: &str,
metadata: ObjectPutMetadata,
) -> ObjectStorageResult<PresignedObjectRequest> {
let content_type = metadata
.content_type
.unwrap_or_else(|| "application/octet-stream".to_string());
let mut headers = HashMap::new();
headers.insert("Content-Type".to_string(), content_type.clone());
if let Some(content_length) = metadata.content_length {
headers.insert("Content-Length".to_string(), content_length.to_string());
}
let mut action = PutObject::new(&self.bucket, Some(&self.credentials), key);
if self.presign_sign_content_type_for_put {
action.headers_mut().insert("content-type", content_type);
}
if let Some(content_length) = metadata.content_length {
action
.headers_mut()
.insert("content-length", content_length.to_string());
}
Ok(PresignedObjectRequest {
url: action.sign(expires_in(self.presign_expires_in_seconds)).to_string(),
headers,
expires_at_ms: expires_at_ms(self.presign_expires_in_seconds)?,
})
}
pub(crate) async fn presign_get(&self, key: &str) -> ObjectStorageResult<PresignedObjectRequest> {
let action = GetObject::new(&self.bucket, Some(&self.credentials), key);
Ok(PresignedObjectRequest {
url: action.sign(expires_in(self.presign_expires_in_seconds)).to_string(),
headers: HashMap::new(),
expires_at_ms: expires_at_ms(self.presign_expires_in_seconds)?,
})
}
pub(crate) async fn create_multipart_upload(
&self,
key: &str,
metadata: ObjectPutMetadata,
) -> ObjectStorageResult<Option<MultipartUploadInitResult>> {
let mut action = CreateMultipartUpload::new(&self.bucket, Some(&self.credentials), key);
if let Some(content_type) = metadata.content_type {
action.headers_mut().insert("content-type", content_type.clone());
let headers = HashMap::from([("content-type".to_string(), content_type)]);
return self.create_multipart_upload_with_headers(key, action, headers).await;
}
self
.create_multipart_upload_with_headers(key, action, HashMap::new())
.await
}
async fn create_multipart_upload_with_headers(
&self,
key: &str,
action: CreateMultipartUpload<'_>,
headers: HashMap<String, String>,
) -> ObjectStorageResult<Option<MultipartUploadInitResult>> {
let response = self
.http
.execute(StorageHttpRequest {
method: Method::POST,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers,
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| {
operation_error(
format!("ObjectStorage create multipart upload failed for {key}"),
source,
)
})?;
let body = ensure_success_text(
response,
format!("ObjectStorage create multipart upload failed for {key}"),
)?;
let parsed = CreateMultipartUpload::parse_response(&body).map_err(|source| ObjectStorageError::InvalidXml {
context: format!("ObjectStorage parse multipart upload response failed for {key}"),
source,
})?;
Ok(Some(MultipartUploadInitResult {
upload_id: parsed.upload_id().to_string(),
expires_at_ms: expires_at_ms(self.presign_expires_in_seconds)?,
}))
}
pub(crate) async fn presign_upload_part(
&self,
key: &str,
upload_id: &str,
part_number: i32,
) -> ObjectStorageResult<PresignedObjectRequest> {
let part_number = checked_part_number(part_number)?;
let action = UploadPart::new(&self.bucket, Some(&self.credentials), key, part_number, upload_id);
Ok(PresignedObjectRequest {
url: action.sign(expires_in(self.presign_expires_in_seconds)).to_string(),
headers: HashMap::new(),
expires_at_ms: expires_at_ms(self.presign_expires_in_seconds)?,
})
}
pub(crate) async fn upload_part(
&self,
key: &str,
upload_id: &str,
part_number: i32,
body: Vec<u8>,
content_length: Option<i64>,
) -> ObjectStorageResult<Option<String>> {
let part_number = checked_part_number(part_number)?;
let mut action = UploadPart::new(&self.bucket, Some(&self.credentials), key, part_number, upload_id);
let mut headers = HashMap::new();
if let Some(content_length) = content_length {
action
.headers_mut()
.insert("content-length", content_length.to_string());
headers.insert("content-length".to_string(), content_length.to_string());
}
let response = self
.http
.execute(StorageHttpRequest {
method: Method::PUT,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers,
body: Some(body),
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage upload multipart part failed for {key}"), source))?;
ensure_success_status(
&response,
&format!("ObjectStorage upload multipart part failed for {key}"),
)?;
Ok(response_header(&response.headers, ETAG).as_deref().map(trim_etag))
}
pub(crate) async fn list_multipart_upload_parts(
&self,
key: &str,
upload_id: &str,
) -> ObjectStorageResult<Vec<MultipartUploadPart>> {
let mut parts = Vec::new();
let mut marker = None;
loop {
let mut action = ListParts::new(&self.bucket, Some(&self.credentials), key, upload_id);
if let Some(marker) = marker {
action.set_part_number_marker(marker);
}
let response = self
.http
.execute(StorageHttpRequest {
method: Method::GET,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| {
operation_error(
format!("ObjectStorage list multipart upload parts failed for {key}"),
source,
)
})?;
if response.status == StatusCode::NOT_FOUND && is_not_found_body(&response.body) {
return Ok(Vec::new());
}
let body = ensure_success_text(
response,
format!("ObjectStorage list multipart upload parts failed for {key}"),
)?;
let parsed = ListParts::parse_response(&body).map_err(|source| ObjectStorageError::InvalidXml {
context: format!("ObjectStorage parse multipart parts failed for {key}"),
source,
})?;
parts.extend(parsed.parts.into_iter().map(|part| MultipartUploadPart {
part_number: i32::from(part.number),
etag: trim_etag(&part.etag),
}));
let Some(next_marker) = parsed.next_part_number_marker else {
break;
};
marker = Some(next_marker);
}
Ok(parts)
}
pub(crate) async fn complete_multipart_upload(
&self,
key: &str,
upload_id: &str,
parts: Vec<MultipartUploadPart>,
) -> ObjectStorageResult<()> {
let ordered_parts = completed_multipart_parts(parts);
validate_completed_parts(&ordered_parts)?;
let etags = ordered_parts.iter().map(|part| part.etag.as_str());
let action = CompleteMultipartUpload::new(&self.bucket, Some(&self.credentials), key, upload_id, etags);
let url = action.sign(expires_in(self.presign_expires_in_seconds));
let body = complete_multipart_body(&ordered_parts);
let response = self
.http
.execute(StorageHttpRequest {
method: Method::POST,
url,
headers: HashMap::from([("content-type".to_string(), "application/xml".to_string())]),
body: Some(body.into_bytes()),
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| {
operation_error(
format!("ObjectStorage complete multipart upload failed for {key}"),
source,
)
})?;
ensure_success_status(
&response,
&format!("ObjectStorage complete multipart upload failed for {key}"),
)?;
Ok(())
}
pub(crate) async fn abort_multipart_upload(&self, key: &str, upload_id: &str) -> ObjectStorageResult<()> {
let action = AbortMultipartUpload::new(&self.bucket, Some(&self.credentials), key, upload_id);
let response = self
.http
.execute(StorageHttpRequest {
method: Method::DELETE,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage abort multipart upload failed for {key}"), source))?;
ensure_success_status(
&response,
&format!("ObjectStorage abort multipart upload failed for {key}"),
)?;
Ok(())
}
pub(crate) async fn head(&self, key: &str) -> ObjectStorageResult<Option<ObjectMetadata>> {
let action = HeadObject::new(&self.bucket, Some(&self.credentials), key);
let response = self
.http
.execute(StorageHttpRequest {
method: Method::HEAD,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage head failed for {key}"), source))?;
if response.status == StatusCode::NOT_FOUND {
let get_action = GetObject::new(&self.bucket, Some(&self.credentials), key);
let get_response = self
.http
.execute(StorageHttpRequest {
method: Method::GET,
url: get_action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage head missing check failed for {key}"), source))?;
if get_response.status == StatusCode::NOT_FOUND && is_not_found_body(&get_response.body) {
return Ok(None);
}
ensure_success_status(
&get_response,
&format!("ObjectStorage head missing check failed for {key}"),
)?;
return Ok(Some(metadata_from_headers(&get_response.headers)));
}
ensure_success_status(&response, &format!("ObjectStorage head failed for {key}"))?;
Ok(Some(metadata_from_headers(&response.headers)))
}
pub(crate) async fn get(&self, key: &str) -> ObjectStorageResult<Option<ObjectGetResult>> {
let action = GetObject::new(&self.bucket, Some(&self.credentials), key);
let response = self
.http
.execute(StorageHttpRequest {
method: Method::GET,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage get failed for {key}"), source))?;
if response.status == StatusCode::NOT_FOUND && is_not_found_body(&response.body) {
return Ok(None);
}
ensure_success_status(&response, &format!("ObjectStorage get failed for {key}"))?;
let metadata = metadata_from_headers(&response.headers);
Ok(Some(ObjectGetResult {
body: response.body,
metadata,
}))
}
pub(crate) async fn list(&self, prefix: Option<String>) -> ObjectStorageResult<Vec<ObjectListEntry>> {
let mut entries = Vec::new();
let mut token = None;
loop {
let page = self.list_page(prefix.clone(), token, None, 1000).await?;
entries.extend(page.entries);
if let Some(next_token) = page.next_continuation_token {
token = Some(next_token);
} else {
break;
}
}
Ok(entries)
}
pub(crate) async fn list_page(
&self,
prefix: Option<String>,
continuation_token: Option<String>,
start_after: Option<String>,
max_keys: i32,
) -> ObjectStorageResult<ObjectListPage> {
let max_keys = usize::try_from(max_keys)
.map_err(|_| ObjectStorageError::InvalidInput("maxKeys must be positive".to_string()))?;
let mut action = ListObjectsV2::new(&self.bucket, Some(&self.credentials));
action.with_max_keys(max_keys);
if let Some(prefix) = &prefix {
action.with_prefix(prefix.clone());
}
if let Some(continuation_token) = &continuation_token {
action.with_continuation_token(continuation_token.clone());
} else if let Some(start_after) = &start_after {
action.with_start_after(start_after.clone());
}
let response = self
.http
.execute(StorageHttpRequest {
method: Method::GET,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error("ObjectStorage list page failed", source))?;
let body = ensure_success_text(response, "ObjectStorage list page failed".to_string())?;
let parsed = ListObjectsV2::parse_response(&body).map_err(|source| ObjectStorageError::InvalidXml {
context: "ObjectStorage parse list response failed".to_string(),
source,
})?;
Ok(ObjectListPage {
entries: parsed
.contents
.into_iter()
.map(|object| ObjectListEntry {
key: object.key,
content_length: i64::try_from(object.size).unwrap_or(i64::MAX),
last_modified_ms: parse_rfc3339_ms(&object.last_modified),
})
.collect(),
next_continuation_token: parsed.next_continuation_token,
})
}
pub(crate) async fn delete(&self, key: &str) -> ObjectStorageResult<()> {
let action = DeleteObject::new(&self.bucket, Some(&self.credentials), key);
let response = self
.http
.execute(StorageHttpRequest {
method: Method::DELETE,
url: action.sign(expires_in(self.presign_expires_in_seconds)),
headers: HashMap::new(),
body: None,
max_response_body_bytes: MAX_RESPONSE_BODY_BYTES,
})
.await
.map_err(|source| operation_error(format!("ObjectStorage delete failed for {key}"), source))?;
ensure_success_status(&response, &format!("ObjectStorage delete failed for {key}"))?;
Ok(())
}
}
fn insert_action_headers<'a, T: S3Action<'a>>(action: &mut T, headers: &HashMap<String, String>) {
for (key, value) in headers {
action.headers_mut().insert(key.clone(), value.clone());
}
}
fn operation_error(context: impl Into<String>, source: ObjectStorageError) -> ObjectStorageError {
ObjectStorageError::Operation {
context: context.into(),
source: Box::new(source),
}
}
fn ensure_success_text(response: StorageHttpResponse, context: String) -> ObjectStorageResult<String> {
ensure_success_status(&response, &context)?;
String::from_utf8(response.body).map_err(|source| ObjectStorageError::InvalidUtf8 { context, source })
}
fn ensure_success_status(response: &StorageHttpResponse, context: &str) -> ObjectStorageResult<()> {
if response.status.is_success() {
return Ok(());
}
let body = String::from_utf8_lossy(&response.body);
Err(ObjectStorageError::HttpStatus {
context: context.to_string(),
status: response.status,
body: body.to_string(),
})
}
fn is_not_found_body(body: &[u8]) -> bool {
let body = String::from_utf8_lossy(body);
body.contains("<Code>NoSuchKey</Code>")
|| body.contains("<Code>NotFound</Code>")
|| body.contains("<Code>NoSuchUpload</Code>")
}
fn metadata_from_headers(headers: &HeaderMap) -> ObjectMetadata {
ObjectMetadata {
content_type: response_header(headers, CONTENT_TYPE).unwrap_or_else(|| "application/octet-stream".to_string()),
content_length: response_header(headers, CONTENT_LENGTH)
.and_then(|value| value.parse::<i64>().ok())
.unwrap_or(0),
last_modified_ms: response_header(headers, LAST_MODIFIED)
.and_then(|value| DateTime::<FixedOffset>::parse_from_rfc2822(&value).ok())
.map(|value| value.timestamp_millis())
.unwrap_or(0),
checksum_crc32: response_header_name(headers, "x-amz-checksum-crc32"),
}
}
fn response_header(headers: &HeaderMap, name: HeaderName) -> Option<String> {
headers
.get(name)
.and_then(|value| value.to_str().ok())
.map(ToString::to_string)
}
fn response_header_name(headers: &HeaderMap, name: &str) -> Option<String> {
headers
.get(name)
.and_then(|value| value.to_str().ok())
.map(ToString::to_string)
}
fn checked_part_number(part_number: i32) -> ObjectStorageResult<u16> {
if !(1..=MAX_MULTIPART_PART_NUMBER).contains(&part_number) {
return Err(ObjectStorageError::InvalidInput(
"multipart part number must be between 1 and 10000".to_string(),
));
}
Ok(part_number as u16)
}
fn validate_completed_parts(parts: &[MultipartUploadPart]) -> ObjectStorageResult<()> {
for part in parts {
checked_part_number(part.part_number)?;
if part.etag.is_empty() {
return Err(ObjectStorageError::InvalidInput(
"multipart part etag is required".to_string(),
));
}
}
Ok(())
}
fn complete_multipart_body(parts: &[MultipartUploadPart]) -> String {
let mut body = String::from("<CompleteMultipartUpload>");
for part in parts {
body.push_str("<Part><ETag>");
body.push_str(&xml_escape(&part.etag));
body.push_str("</ETag><PartNumber>");
body.push_str(&part.part_number.to_string());
body.push_str("</PartNumber></Part>");
}
body.push_str("</CompleteMultipartUpload>");
body
}
fn xml_escape(value: &str) -> String {
value.replace('&', "&amp;").replace('<', "&lt;").replace('>', "&gt;")
}
fn expires_in(seconds: u64) -> Duration {
Duration::from_secs(seconds)
}
fn expires_at_ms(expires_in_seconds: u64) -> ObjectStorageResult<i64> {
let expires_at = SystemTime::now()
.checked_add(Duration::from_secs(expires_in_seconds))
.ok_or_else(|| ObjectStorageError::InvalidInput("presign expiration overflow".to_string()))?;
crate::utils::system_time_millis(expires_at)
.map(|millis| millis as i64)
.map_err(|err| ObjectStorageError::InvalidInput(format!("system time before unix epoch: {err}")))
}
fn parse_rfc3339_ms(value: &str) -> i64 {
DateTime::parse_from_rfc3339(value)
.map(|value| value.timestamp_millis())
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use reqwest::header::HeaderValue;
use super::*;
#[test]
fn metadata_from_headers_uses_s3_defaults_and_checksum() {
let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("text/plain"));
headers.insert(CONTENT_LENGTH, HeaderValue::from_static("42"));
headers.insert(LAST_MODIFIED, HeaderValue::from_static("Wed, 21 Oct 2015 07:28:00 GMT"));
headers.insert("x-amz-checksum-crc32", HeaderValue::from_static("checksum"));
let metadata = metadata_from_headers(&headers);
assert_eq!(metadata.content_type, "text/plain");
assert_eq!(metadata.content_length, 42);
assert_eq!(metadata.last_modified_ms, 1_445_412_480_000);
assert_eq!(metadata.checksum_crc32.as_deref(), Some("checksum"));
let defaults = metadata_from_headers(&HeaderMap::new());
assert_eq!(defaults.content_type, "application/octet-stream");
assert_eq!(defaults.content_length, 0);
assert_eq!(defaults.last_modified_ms, 0);
assert!(defaults.checksum_crc32.is_none());
}
#[test]
fn not_found_body_accepts_object_missing_codes_only() {
for body in [
"<Error><Code>NoSuchKey</Code></Error>",
"<Error><Code>NotFound</Code></Error>",
"<Error><Code>NoSuchUpload</Code></Error>",
] {
assert!(is_not_found_body(body.as_bytes()), "{body}");
}
assert!(!is_not_found_body(b""));
assert!(!is_not_found_body(b"<Error><Code>AccessDenied</Code></Error>"));
}
#[test]
fn list_parts_xml_handles_array_single_part_and_pagination() {
let xml = r#"<?xml version="1.0" encoding="UTF-8"?>
<ListPartsResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/">
<Bucket>test</Bucket>
<Key>key</Key>
<UploadId>upload-id</UploadId>
<PartNumberMarker>0</PartNumberMarker>
<NextPartNumberMarker>3</NextPartNumberMarker>
<MaxParts>2</MaxParts>
<IsTruncated>true</IsTruncated>
<Part>
<PartNumber>1</PartNumber>
<LastModified>2010-11-10T20:48:34.000Z</LastModified>
<ETag>"etag-1"</ETag>
<Size>10485760</Size>
</Part>
<Part>
<PartNumber>2</PartNumber>
<LastModified>2010-11-10T20:48:33.000Z</LastModified>
<ETag>etag-2</ETag>
<Size>10485760</Size>
</Part>
</ListPartsResult>"#;
let parsed = ListParts::parse_response(xml).unwrap();
let parts = parsed
.parts
.into_iter()
.map(|part| MultipartUploadPart {
part_number: i32::from(part.number),
etag: trim_etag(&part.etag),
})
.collect::<Vec<_>>();
assert_eq!(
parts,
vec![
MultipartUploadPart {
part_number: 1,
etag: "etag-1".to_string()
},
MultipartUploadPart {
part_number: 2,
etag: "etag-2".to_string()
}
]
);
assert_eq!(parsed.next_part_number_marker, Some(3));
let single = r#"<?xml version="1.0" encoding="UTF-8"?>
<ListPartsResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/">
<Bucket>test</Bucket>
<Key>key</Key>
<UploadId>upload-id</UploadId>
<MaxParts>1</MaxParts>
<IsTruncated>false</IsTruncated>
<Part>
<PartNumber>5</PartNumber>
<LastModified>2010-11-10T20:48:34.000Z</LastModified>
<ETag>"etag-5"</ETag>
<Size>10485760</Size>
</Part>
</ListPartsResult>"#;
let parsed = ListParts::parse_response(single).unwrap();
assert_eq!(parsed.parts.len(), 1);
assert_eq!(parsed.parts[0].number, 5);
assert_eq!(trim_etag(&parsed.parts[0].etag), "etag-5");
assert_eq!(parsed.next_part_number_marker, None);
}
#[test]
fn complete_multipart_body_orders_and_escapes_parts() {
let mut parts = completed_multipart_parts(vec![
MultipartUploadPart {
part_number: 2,
etag: "b&c".to_string(),
},
MultipartUploadPart {
part_number: 1,
etag: "a<tag>".to_string(),
},
]);
validate_completed_parts(&parts).unwrap();
let body = complete_multipart_body(&parts);
assert_eq!(
body,
"<CompleteMultipartUpload><Part><ETag>a&lt;tag&gt;</ETag><PartNumber>1</PartNumber></Part><Part><ETag>b&amp;c</\
ETag><PartNumber>2</PartNumber></Part></CompleteMultipartUpload>"
);
parts[0].etag.clear();
assert!(validate_completed_parts(&parts).is_err());
assert!(
validate_completed_parts(&[MultipartUploadPart {
part_number: -1,
etag: "etag".to_string(),
}])
.is_err()
);
assert!(
validate_completed_parts(&[MultipartUploadPart {
part_number: 0,
etag: "etag".to_string(),
}])
.is_err()
);
assert!(
validate_completed_parts(&[MultipartUploadPart {
part_number: 10_001,
etag: "etag".to_string(),
}])
.is_err()
);
}
#[test]
fn parse_rfc3339_ms_returns_zero_for_invalid_values() {
assert_eq!(parse_rfc3339_ms("2024-01-02T03:04:05Z"), 1_704_164_645_000);
assert_eq!(parse_rfc3339_ms("not a date"), 0);
}
}
@@ -0,0 +1,210 @@
use rusty_s3::{Bucket, Credentials, UrlStyle};
use serde::Deserialize;
use url::Url;
use super::{
client::ObjectStorageClient,
error::{ObjectStorageError, ObjectStorageResult},
types::StorageProviderConfig,
};
#[derive(Clone, Debug)]
pub(crate) struct ObjectStorageConfig {
pub(crate) provider: String,
pub(crate) bucket: String,
pub(crate) endpoint: Option<String>,
pub(crate) region: Option<String>,
pub(crate) access_key_id: Option<String>,
pub(crate) secret_access_key: Option<String>,
pub(crate) session_token: Option<String>,
pub(crate) force_path_style: bool,
pub(crate) request_timeout_ms: Option<u64>,
pub(crate) min_part_size: Option<u64>,
pub(crate) presign_expires_in_seconds: Option<u64>,
pub(crate) presign_sign_content_type_for_put: Option<bool>,
pub(crate) use_presigned_url: bool,
pub(crate) proxy_upload: bool,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct S3ConfigFile {
endpoint: Option<String>,
region: Option<String>,
credentials: Option<S3CredentialsConfigFile>,
force_path_style: Option<bool>,
request_timeout_ms: Option<u64>,
min_part_size: Option<u64>,
presign: Option<S3PresignConfigFile>,
#[serde(rename = "usePresignedURL")]
use_presigned_url: Option<UsePresignedUrlConfigFile>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct R2ConfigFile {
account_id: String,
jurisdiction: Option<String>,
region: Option<String>,
credentials: Option<S3CredentialsConfigFile>,
request_timeout_ms: Option<u64>,
min_part_size: Option<u64>,
presign: Option<S3PresignConfigFile>,
#[serde(rename = "usePresignedURL")]
use_presigned_url: Option<UsePresignedUrlConfigFile>,
}
#[derive(Debug, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
struct S3CredentialsConfigFile {
access_key_id: Option<String>,
secret_access_key: Option<String>,
session_token: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct S3PresignConfigFile {
expires_in_seconds: Option<u64>,
sign_content_type_for_put: Option<bool>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct UsePresignedUrlConfigFile {
enabled: bool,
url_prefix: Option<String>,
sign_key: Option<String>,
}
impl ObjectStorageConfig {
pub(crate) fn from_provider_config(storage: Option<StorageProviderConfig>) -> ObjectStorageResult<Option<Self>> {
let Some(storage) = storage else {
return Ok(None);
};
match storage.provider.as_str() {
"aws-s3" => Self::from_s3_config(storage),
"cloudflare-r2" => Self::from_r2_config(storage),
"fs" => Ok(None),
provider => Err(ObjectStorageError::Config(format!(
"unsupported blob storage provider for StorageRuntime: {provider}"
))),
}
}
pub(crate) fn from_s3_config(storage: StorageProviderConfig) -> ObjectStorageResult<Option<Self>> {
let config: S3ConfigFile = serde_json::from_value(storage.config)
.map_err(|err| ObjectStorageError::Config(format!("invalid aws-s3 blob storage config: {err}")))?;
let region = config
.region
.ok_or_else(|| ObjectStorageError::Config("aws-s3 blob storage config requires region".to_string()))?;
let endpoint = config.endpoint.or_else(|| Some(resolve_s3_endpoint(&region)));
let credentials = config.credentials.unwrap_or_default();
Ok(Some(Self {
provider: storage.provider,
bucket: storage.bucket,
endpoint,
region: Some(region),
access_key_id: credentials.access_key_id,
secret_access_key: credentials.secret_access_key,
session_token: credentials.session_token,
force_path_style: config.force_path_style.unwrap_or(false),
request_timeout_ms: config.request_timeout_ms,
min_part_size: config.min_part_size,
presign_expires_in_seconds: config.presign.as_ref().and_then(|v| v.expires_in_seconds),
presign_sign_content_type_for_put: config.presign.as_ref().and_then(|v| v.sign_content_type_for_put),
use_presigned_url: config.use_presigned_url.map(|v| v.enabled).unwrap_or(false),
proxy_upload: false,
}))
}
pub(crate) fn from_r2_config(storage: StorageProviderConfig) -> ObjectStorageResult<Option<Self>> {
let config: R2ConfigFile = serde_json::from_value(storage.config)
.map_err(|err| ObjectStorageError::Config(format!("invalid cloudflare-r2 blob storage config: {err}")))?;
let account = match config.jurisdiction {
Some(jurisdiction) => format!("{}.{}", config.account_id, jurisdiction),
None => config.account_id,
};
let credentials = config.credentials.unwrap_or_default();
let (use_presigned_url, proxy_upload) = config
.use_presigned_url
.map(|value| {
(
value.enabled,
value.enabled
&& value.url_prefix.as_ref().is_some_and(|prefix| !prefix.is_empty())
&& value.sign_key.as_ref().is_some_and(|key| !key.is_empty()),
)
})
.unwrap_or((false, false));
Ok(Some(Self {
provider: storage.provider,
bucket: storage.bucket,
endpoint: Some(format!("https://{account}.r2.cloudflarestorage.com")),
region: Some(config.region.unwrap_or_else(|| "auto".to_string())),
access_key_id: credentials.access_key_id,
secret_access_key: credentials.secret_access_key,
session_token: credentials.session_token,
force_path_style: true,
request_timeout_ms: config.request_timeout_ms,
min_part_size: config.min_part_size,
presign_expires_in_seconds: config.presign.as_ref().and_then(|v| v.expires_in_seconds),
presign_sign_content_type_for_put: config.presign.as_ref().and_then(|v| v.sign_content_type_for_put),
use_presigned_url,
proxy_upload,
}))
}
pub(crate) fn build_client(&self) -> ObjectStorageResult<ObjectStorageClient> {
let region = self
.region
.clone()
.ok_or_else(|| ObjectStorageError::Config("object storage region is required".to_string()))?;
let access_key_id = self
.access_key_id
.clone()
.ok_or_else(|| ObjectStorageError::Config("object storage accessKeyId is required".to_string()))?;
let secret_access_key = self
.secret_access_key
.clone()
.ok_or_else(|| ObjectStorageError::Config("object storage secretAccessKey is required".to_string()))?;
let endpoint = self.endpoint.clone().unwrap_or_else(|| resolve_s3_endpoint(&region));
let endpoint = Url::parse(&endpoint)
.map_err(|err| ObjectStorageError::Config(format!("object storage endpoint is invalid: {err}")))?;
let bucket = Bucket::new(
endpoint,
if self.force_path_style {
UrlStyle::Path
} else {
UrlStyle::VirtualHost
},
self.bucket.clone(),
region,
)
.map_err(|err| ObjectStorageError::Config(format!("object storage bucket url is invalid: {err}")))?;
let credentials = match self.session_token.as_ref().filter(|token| !token.is_empty()) {
Some(session_token) => Credentials::new_with_token(access_key_id, secret_access_key, session_token.clone()),
None => Credentials::new(access_key_id, secret_access_key),
};
ObjectStorageClient::new(
bucket,
credentials,
self.request_timeout_ms,
self.presign_expires_in_seconds.unwrap_or(60),
self.presign_sign_content_type_for_put.unwrap_or(true),
)
}
}
fn resolve_s3_endpoint(region: &str) -> String {
if region == "us-east-1" {
"https://s3.amazonaws.com".to_string()
} else {
format!("https://s3.{region}.amazonaws.com")
}
}
@@ -0,0 +1,56 @@
use reqwest::StatusCode;
#[derive(Debug, thiserror::Error)]
pub(crate) enum ObjectStorageError {
#[error("ObjectStorage config error: {0}")]
Config(String),
#[error("{context}: {source}")]
Operation {
context: String,
#[source]
source: Box<ObjectStorageError>,
},
#[error("ObjectStorage http client build failed: {0}")]
HttpClientBuild(#[source] reqwest::Error),
#[error("ObjectStorage http request failed: {0}")]
HttpRequest(#[source] reqwest::Error),
#[error("ObjectStorage invalid http header: {0}")]
InvalidHeader(String),
#[error("ObjectStorage response body exceeds {limit} bytes")]
BodyTooLarge { limit: usize },
#[error("{context}: status={status} body={body}")]
HttpStatus {
context: String,
status: StatusCode,
body: String,
},
#[error("{context}: invalid utf8 response: {source}")]
InvalidUtf8 {
context: String,
#[source]
source: std::string::FromUtf8Error,
},
#[error("{context}: invalid xml response: {source}")]
InvalidXml {
context: String,
#[source]
source: instant_xml::Error,
},
#[error("ObjectStorage invalid input: {0}")]
InvalidInput(String),
}
impl ObjectStorageError {
pub(crate) fn is_not_found(&self) -> bool {
match self {
Self::Operation { source, .. } => source.is_not_found(),
Self::HttpStatus { status, body, .. } => {
*status == StatusCode::NOT_FOUND
&& (body.contains("NoSuchKey") || body.contains("NoSuchUpload") || body.contains("NotFound"))
}
_ => false,
}
}
}
pub(crate) type ObjectStorageResult<T> = std::result::Result<T, ObjectStorageError>;
@@ -0,0 +1,9 @@
pub(crate) mod client;
pub(crate) mod config;
pub(crate) mod error;
#[cfg(test)]
mod tests;
pub(crate) mod types;
pub(crate) use config::ObjectStorageConfig;
pub(crate) use types::StorageProviderConfig;
@@ -0,0 +1,367 @@
use reqwest::StatusCode;
use super::{
config::ObjectStorageConfig,
error::ObjectStorageError,
types::{MultipartUploadPart, ObjectPutMetadata, StorageProviderConfig, completed_multipart_parts, trim_etag},
};
fn storage_config(provider: &str, config: serde_json::Value) -> StorageProviderConfig {
StorageProviderConfig {
provider: provider.to_string(),
bucket: "test-bucket".to_string(),
config,
}
}
#[test]
fn resolves_r2_config_from_config_json_shape() {
let storage = StorageProviderConfig {
provider: "cloudflare-r2".to_string(),
bucket: "workspace-blobs".to_string(),
config: serde_json::json!({
"accountId": "account",
"jurisdiction": "eu",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
},
"usePresignedURL": {
"enabled": true
}
}),
};
let config = ObjectStorageConfig::from_r2_config(storage).unwrap().unwrap();
assert_eq!(config.provider, "cloudflare-r2");
assert_eq!(config.bucket, "workspace-blobs");
assert_eq!(
config.endpoint.as_deref(),
Some("https://account.eu.r2.cloudflarestorage.com")
);
assert_eq!(config.region.as_deref(), Some("auto"));
assert!(config.force_path_style);
assert!(config.use_presigned_url);
assert!(!config.proxy_upload);
assert_eq!(config.access_key_id.as_deref(), Some("key"));
}
#[test]
fn resolves_r2_endpoint_cases_from_config_json_shape() {
for (case, config, expected_endpoint) in [
(
"default account endpoint",
serde_json::json!({
"accountId": "account",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
}
}),
Some("https://account.r2.cloudflarestorage.com"),
),
(
"explicit null jurisdiction",
serde_json::json!({
"accountId": "account",
"jurisdiction": null,
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
}
}),
Some("https://account.r2.cloudflarestorage.com"),
),
(
"eu jurisdiction",
serde_json::json!({
"accountId": "account",
"jurisdiction": "eu",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
}
}),
Some("https://account.eu.r2.cloudflarestorage.com"),
),
] {
let config = ObjectStorageConfig::from_r2_config(storage_config("cloudflare-r2", config))
.unwrap()
.unwrap();
assert_eq!(config.endpoint.as_deref(), expected_endpoint, "{case}");
assert!(config.force_path_style, "{case}");
}
assert!(
ObjectStorageConfig::from_r2_config(storage_config(
"cloudflare-r2",
serde_json::json!({
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
}
})
))
.is_err()
);
}
#[test]
fn object_storage_not_found_requires_object_error_code() {
let bucket_or_route_missing = ObjectStorageError::HttpStatus {
context: "head failed".to_string(),
status: StatusCode::NOT_FOUND,
body: String::new(),
};
let object_missing = ObjectStorageError::HttpStatus {
context: "get failed".to_string(),
status: StatusCode::NOT_FOUND,
body: "<Error><Code>NoSuchKey</Code></Error>".to_string(),
};
let upload_missing = ObjectStorageError::HttpStatus {
context: "abort failed".to_string(),
status: StatusCode::NOT_FOUND,
body: "<Error><Code>NoSuchUpload</Code></Error>".to_string(),
};
assert!(!bucket_or_route_missing.is_not_found());
assert!(object_missing.is_not_found());
assert!(upload_missing.is_not_found());
}
#[test]
fn resolves_r2_proxy_upload_capability_from_config_json_shape() {
let storage = StorageProviderConfig {
provider: "cloudflare-r2".to_string(),
bucket: "workspace-blobs".to_string(),
config: serde_json::json!({
"accountId": "account",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
},
"usePresignedURL": {
"enabled": true,
"urlPrefix": "https://cdn.example.com",
"signKey": "secret"
}
}),
};
let config = ObjectStorageConfig::from_r2_config(storage).unwrap().unwrap();
assert!(config.use_presigned_url);
assert!(config.proxy_upload);
}
#[test]
fn resolves_s3_config_from_config_json_shape() {
let storage = StorageProviderConfig {
provider: "aws-s3".to_string(),
bucket: "workspace-blobs".to_string(),
config: serde_json::json!({
"region": "us-west-2",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret",
"sessionToken": "session"
},
"forcePathStyle": true,
"requestTimeoutMs": 1000,
"minPartSize": 1024,
"presign": {
"expiresInSeconds": 60,
"signContentTypeForPut": false
}
}),
};
let config = ObjectStorageConfig::from_s3_config(storage).unwrap().unwrap();
assert_eq!(config.provider, "aws-s3");
assert_eq!(config.endpoint.as_deref(), Some("https://s3.us-west-2.amazonaws.com"));
assert_eq!(config.session_token.as_deref(), Some("session"));
assert!(config.force_path_style);
assert_eq!(config.request_timeout_ms, Some(1000));
assert_eq!(config.min_part_size, Some(1024));
assert_eq!(config.presign_expires_in_seconds, Some(60));
assert_eq!(config.presign_sign_content_type_for_put, Some(false));
}
#[test]
fn resolves_s3_default_endpoint_cases_from_config_json_shape() {
for (region, expected_endpoint) in [
("us-east-1", "https://s3.amazonaws.com"),
("us-west-2", "https://s3.us-west-2.amazonaws.com"),
] {
let config = ObjectStorageConfig::from_s3_config(storage_config(
"aws-s3",
serde_json::json!({
"region": region,
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
}
}),
))
.unwrap()
.unwrap();
assert_eq!(config.endpoint.as_deref(), Some(expected_endpoint), "{region}");
}
}
#[tokio::test]
async fn object_storage_presign_put_returns_sigv4_url_and_headers() {
let storage = StorageProviderConfig {
provider: "aws-s3".to_string(),
bucket: "test-bucket".to_string(),
config: serde_json::json!({
"region": "us-east-1",
"endpoint": "https://s3.us-east-1.amazonaws.com",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
},
"presign": {
"expiresInSeconds": 60
}
}),
};
let config = ObjectStorageConfig::from_s3_config(storage).unwrap().unwrap();
let Ok(Ok(client)) = std::panic::catch_unwind(|| config.build_client()) else {
eprintln!("skipping object storage presign test: S3 client cannot be built in this environment");
return;
};
let result = client
.presign_put(
"key",
ObjectPutMetadata {
content_type: Some("text/plain".to_string()),
..Default::default()
},
)
.await
.unwrap();
assert!(result.url.contains("X-Amz-Algorithm=AWS4-HMAC-SHA256"));
assert!(result.url.contains("X-Amz-SignedHeaders="));
assert_eq!(
result.headers.get("Content-Type").map(String::as_str),
Some("text/plain")
);
assert!(result.expires_at_ms > 0);
}
#[tokio::test]
async fn object_storage_presign_put_respects_content_length_and_signed_content_type_flag() {
let config = ObjectStorageConfig::from_s3_config(storage_config(
"aws-s3",
serde_json::json!({
"region": "us-east-1",
"endpoint": "https://s3.us-east-1.amazonaws.com",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
},
"presign": {
"expiresInSeconds": 60,
"signContentTypeForPut": false
}
}),
))
.unwrap()
.unwrap();
let client = config.build_client().unwrap();
let result = client
.presign_put(
"key",
ObjectPutMetadata {
content_type: Some("text/plain".to_string()),
content_length: Some(42),
..Default::default()
},
)
.await
.unwrap();
assert_eq!(
result.headers.get("Content-Type").map(String::as_str),
Some("text/plain")
);
assert_eq!(result.headers.get("Content-Length").map(String::as_str), Some("42"));
assert!(!result.url.contains("content-type"));
assert!(result.url.contains("content-length"));
}
#[tokio::test]
async fn object_storage_presign_get_returns_sigv4_url_without_headers() {
let storage = StorageProviderConfig {
provider: "cloudflare-r2".to_string(),
bucket: "test-bucket".to_string(),
config: serde_json::json!({
"accountId": "account",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
},
"presign": {
"expiresInSeconds": 60
}
}),
};
let config = ObjectStorageConfig::from_r2_config(storage).unwrap().unwrap();
let client = config.build_client().unwrap();
let result = client.presign_get("workspace/key").await.unwrap();
assert!(result.url.contains("X-Amz-Algorithm=AWS4-HMAC-SHA256"));
assert!(result.url.contains("X-Amz-SignedHeaders=host"));
assert!(result.url.contains("/test-bucket/workspace/key?"));
assert!(result.headers.is_empty());
assert!(result.expires_at_ms > 0);
}
#[tokio::test]
async fn object_storage_presign_upload_part_returns_sigv4_url() {
let config = ObjectStorageConfig::from_s3_config(storage_config(
"aws-s3",
serde_json::json!({
"region": "us-east-1",
"endpoint": "https://s3.us-east-1.amazonaws.com",
"credentials": {
"accessKeyId": "key",
"secretAccessKey": "secret"
},
"presign": {
"expiresInSeconds": 60
}
}),
))
.unwrap()
.unwrap();
let client = config.build_client().unwrap();
let result = client.presign_upload_part("key", "upload-1", 3).await.unwrap();
assert!(result.url.contains("X-Amz-Algorithm=AWS4-HMAC-SHA256"));
assert!(result.url.contains("partNumber=3"));
assert!(result.url.contains("uploadId=upload-1"));
assert!(result.headers.is_empty());
assert!(result.expires_at_ms > 0);
}
#[test]
fn object_storage_orders_completed_multipart_parts_and_trims_etags() {
let parts = completed_multipart_parts(vec![
MultipartUploadPart {
part_number: 2,
etag: trim_etag("\"b\""),
},
MultipartUploadPart {
part_number: 1,
etag: trim_etag("a"),
},
]);
assert_eq!(parts[0].part_number, 1);
assert_eq!(parts[0].etag, "a");
assert_eq!(parts[1].part_number, 2);
assert_eq!(parts[1].etag, "b");
}
@@ -0,0 +1,182 @@
use std::collections::HashMap;
use serde::Deserialize;
use super::super::{
RuntimeError, RuntimeMultipartUploadInit, RuntimeMultipartUploadPart, RuntimeObjectGetResult, RuntimeObjectListEntry,
RuntimeObjectMetadata, RuntimeObjectStoragePutOptions, RuntimePresignedObjectRequest, RuntimeResult,
};
#[derive(Clone, Debug, Default)]
pub(crate) struct ObjectPutMetadata {
pub(crate) content_type: Option<String>,
pub(crate) content_length: Option<i64>,
pub(crate) checksum_crc32: Option<String>,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct ObjectMetadata {
pub(crate) content_type: String,
pub(crate) content_length: i64,
pub(crate) last_modified_ms: i64,
pub(crate) checksum_crc32: Option<String>,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct ObjectListEntry {
pub(crate) key: String,
pub(crate) content_length: i64,
pub(crate) last_modified_ms: i64,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct ObjectListPage {
pub(crate) entries: Vec<ObjectListEntry>,
pub(crate) next_continuation_token: Option<String>,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct ObjectGetResult {
pub(crate) body: Vec<u8>,
pub(crate) metadata: ObjectMetadata,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct PresignedObjectRequest {
pub(crate) url: String,
pub(crate) headers: HashMap<String, String>,
pub(crate) expires_at_ms: i64,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct MultipartUploadInitResult {
pub(crate) upload_id: String,
pub(crate) expires_at_ms: i64,
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct MultipartUploadPart {
pub(crate) part_number: i32,
pub(crate) etag: String,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct StorageProviderConfig {
pub(crate) provider: String,
pub(crate) bucket: String,
#[serde(default)]
pub(crate) config: serde_json::Value,
}
pub(crate) fn trim_etag(etag: &str) -> String {
etag.trim_matches('"').to_string()
}
pub(crate) fn completed_multipart_parts(mut parts: Vec<MultipartUploadPart>) -> Vec<MultipartUploadPart> {
parts.sort_by_key(|part| part.part_number);
parts
}
impl From<RuntimeObjectStoragePutOptions> for ObjectPutMetadata {
fn from(options: RuntimeObjectStoragePutOptions) -> Self {
Self {
content_type: options.content_type,
content_length: options.content_length,
checksum_crc32: options.checksum_crc32,
}
}
}
impl ObjectPutMetadata {
pub(crate) fn complete_for_body(mut self, body: &[u8]) -> Self {
self.content_length.get_or_insert(body.len() as i64);
self
.checksum_crc32
.get_or_insert_with(|| format!("{:x}", crc32fast::hash(body)));
self
.content_type
.get_or_insert_with(|| crate::file_type::get_mime(body));
self
}
pub(crate) fn into_object_metadata(self, last_modified_ms: i64) -> ObjectMetadata {
ObjectMetadata {
content_type: self
.content_type
.unwrap_or_else(|| "application/octet-stream".to_string()),
content_length: self.content_length.unwrap_or(0),
last_modified_ms,
checksum_crc32: self.checksum_crc32,
}
}
}
impl From<ObjectMetadata> for RuntimeObjectMetadata {
fn from(metadata: ObjectMetadata) -> Self {
Self {
content_type: metadata.content_type,
content_length: metadata.content_length,
last_modified_ms: metadata.last_modified_ms,
checksum_crc32: metadata.checksum_crc32,
}
}
}
impl From<ObjectListEntry> for RuntimeObjectListEntry {
fn from(entry: ObjectListEntry) -> Self {
Self {
key: entry.key,
content_length: entry.content_length,
last_modified_ms: entry.last_modified_ms,
}
}
}
impl TryFrom<PresignedObjectRequest> for RuntimePresignedObjectRequest {
type Error = RuntimeError;
fn try_from(request: PresignedObjectRequest) -> RuntimeResult<Self> {
Ok(Self {
url: request.url,
headers_json: serde_json::to_string(&request.headers)
.map_err(|err| RuntimeError::json("ObjectStorage headers serialization failed", err))?,
expires_at_ms: request.expires_at_ms,
})
}
}
impl From<ObjectGetResult> for RuntimeObjectGetResult {
fn from(result: ObjectGetResult) -> Self {
Self {
body: result.body.into(),
metadata: result.metadata.into(),
}
}
}
impl From<MultipartUploadInitResult> for RuntimeMultipartUploadInit {
fn from(init: MultipartUploadInitResult) -> Self {
Self {
upload_id: init.upload_id,
expires_at_ms: init.expires_at_ms,
}
}
}
impl From<RuntimeMultipartUploadPart> for MultipartUploadPart {
fn from(part: RuntimeMultipartUploadPart) -> Self {
Self {
part_number: part.part_number,
etag: part.etag,
}
}
}
impl From<MultipartUploadPart> for RuntimeMultipartUploadPart {
fn from(part: MultipartUploadPart) -> Self {
Self {
part_number: part.part_number,
etag: part.etag,
}
}
}
@@ -0,0 +1,202 @@
use napi::bindgen_prelude::Buffer;
#[napi_derive::napi(object)]
pub struct RuntimeVerificationTokenRecord {
pub token_type: i32,
pub token: String,
pub credential: Option<String>,
pub expires_at_ms: i64,
}
#[napi_derive::napi(object)]
pub struct BackendRuntimeHealth {
pub started: bool,
pub database_connected: bool,
}
#[napi_derive::napi(object)]
pub struct CoordinationLeaseGrant {
pub key: String,
pub owner: String,
#[napi(ts_type = "bigint | number")]
pub fencing_token: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeMagicLinkOtpConsumeResult {
pub ok: bool,
pub token: Option<String>,
pub reason: Option<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeWorkspaceInviteLinkRecord {
pub workspace_id: String,
pub invite_id: String,
pub inviter_user_id: String,
pub expires_at_ms: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeByokLocalLeaseRecord {
pub lease_id: String,
pub payload: serde_json::Value,
pub expires_at_ms: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeDocHistoryInput {
pub workspace_id: String,
pub doc_id: String,
pub blob: Buffer,
pub timestamp_ms: i64,
pub editor_id: Option<String>,
pub force: bool,
pub history_min_interval_ms: i64,
pub history_max_age_ms: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeObjectStoragePutOptions {
pub content_type: Option<String>,
pub content_length: Option<i64>,
pub checksum_crc32: Option<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeObjectMetadata {
pub content_type: String,
pub content_length: i64,
pub last_modified_ms: i64,
pub checksum_crc32: Option<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeObjectListEntry {
pub key: String,
pub content_length: i64,
pub last_modified_ms: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeObjectGetResult {
pub body: Buffer,
pub metadata: RuntimeObjectMetadata,
}
#[napi_derive::napi(object)]
pub struct RuntimePresignedObjectRequest {
pub url: String,
pub headers_json: String,
pub expires_at_ms: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeMultipartUploadInit {
pub upload_id: String,
pub expires_at_ms: i64,
}
#[napi_derive::napi(object)]
pub struct RuntimeMultipartUploadPart {
pub part_number: i32,
pub etag: String,
}
#[napi_derive::napi(object)]
pub struct RuntimeBlobCleanupResult {
pub scanned: i64,
pub deleted: i64,
pub aborted_multipart: i64,
pub workspace_ids: Vec<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeBlobCompleteResult {
pub ok: bool,
pub reason: Option<String>,
pub content_type: Option<String>,
pub content_length: Option<i64>,
pub last_modified_ms: Option<i64>,
}
#[napi_derive::napi(object)]
pub struct RuntimeBlobMetadataBackfillResult {
pub scanned_objects: i64,
pub headed_objects: i64,
pub upserted_metadata: i64,
pub skipped_existing: i64,
pub skipped_workspace_missing: i64,
pub failed: i64,
pub next_cursor: Option<String>,
pub workspace_ids: Vec<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeDocBlobRefsResult {
pub scanned_docs: i64,
pub parsed_docs: i64,
pub refs_written: i64,
pub refs_deleted: i64,
pub failed_docs: i64,
pub next_cursor: Option<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeBlobCleanupPlanResult {
pub run_id: Option<String>,
pub scanned_blobs: i64,
pub candidates_marked: i64,
pub protected_by_doc_refs: i64,
pub protected_by_metadata: i64,
pub protected_by_other_refs: i64,
pub next_cursor: Option<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeBlobCleanupExecuteResult {
pub scanned_candidates: i64,
pub deleted_objects: i64,
pub deleted_metadata: i64,
pub skipped_still_referenced: i64,
pub failed: i64,
pub workspace_ids: Vec<String>,
}
#[napi_derive::napi(object)]
pub struct RuntimeDocCompactionResult {
pub lease_acquired: bool,
pub merged: bool,
pub workspace_id: String,
pub doc_id: String,
pub updates_merged: i64,
pub history_created: bool,
}
#[napi_derive::napi(object)]
pub struct RuntimeWorkspaceStatsRefreshResult {
pub processed: i64,
pub backlog: i64,
pub skipped: bool,
}
#[napi_derive::napi(object)]
pub struct RuntimeWorkspaceStatsRecalibrationResult {
pub processed: i64,
pub last_sid: i64,
pub skipped: bool,
}
#[napi_derive::napi(object)]
pub struct RuntimeWorkspaceStatsSnapshotResult {
pub snapshotted: i64,
pub skipped: bool,
}
#[napi_derive::napi(object)]
pub struct RuntimeWorkspaceStatsDailyRecalibrationResult {
pub processed: i64,
pub last_sid: i64,
pub snapshotted: i64,
pub skipped: bool,
}