feat(security): harden gateway boundaries and usage policies

Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change.

Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
elky
2026-09-04 03:45:52 +08:00
parent ddcbeb3ae9
commit 579f2c7cc1
1019 changed files with 190437 additions and 26080 deletions
@@ -202,6 +202,22 @@ impl DataBackends {
}
}
pub async fn compare_and_set_system_config_string_value(
&self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
match self.sql_backend() {
Some(backend) => {
backend
.compare_and_set_system_config_string_value(key, expected, replacement)
.await
}
None => Ok(false),
}
}
pub async fn list_system_config_entries(
&self,
) -> Result<Vec<StoredSystemConfigEntry>, DataLayerError> {
@@ -307,6 +323,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Sqlite(sqlite) => {
warm_pool(sqlite.pool(), sqlite.config().pool.min_connections).await
}
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -321,6 +339,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.run_table_maintenance(table_names).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.run_table_maintenance(table_names).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -341,6 +361,8 @@ impl<'a> SqlBackendRef<'a> {
crate::lifecycle::migrate::run_sqlite_migrations(sqlite.pool()).await?;
Ok(true)
}
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -361,6 +383,8 @@ impl<'a> SqlBackendRef<'a> {
crate::lifecycle::backfill::run_sqlite_backfills(sqlite.pool()).await?;
Ok(true)
}
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -380,6 +404,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Sqlite(sqlite) => Ok(Some(
crate::lifecycle::migrate::pending_sqlite_migrations(sqlite.pool()).await?,
)),
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -400,6 +426,8 @@ impl<'a> SqlBackendRef<'a> {
crate::lifecycle::migrate::prepare_sqlite_database_for_startup(sqlite.pool())
.await?,
)),
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -419,6 +447,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Sqlite(sqlite) => Ok(Some(
crate::lifecycle::backfill::pending_sqlite_backfills(sqlite.pool()).await?,
)),
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -445,6 +475,8 @@ impl<'a> SqlBackendRef<'a> {
sqlite.pool().num_idle(),
sqlite.config().pool.max_connections,
),
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -459,6 +491,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.aggregate_wallet_daily_usage(input).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.aggregate_wallet_daily_usage(input).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -473,6 +507,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.aggregate_stats_hourly(input).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.aggregate_stats_hourly(input).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -487,6 +523,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.aggregate_stats_daily(input).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.aggregate_stats_daily(input).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -501,6 +539,38 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.find_system_config_value(key).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.find_system_config_value(key).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
async fn compare_and_set_system_config_string_value(
self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => {
postgres
.compare_and_set_system_config_string_value(key, expected, replacement)
.await
}
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => {
mysql
.compare_and_set_system_config_string_value(key, expected, replacement)
.await
}
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => {
sqlite
.compare_and_set_system_config_string_value(key, expected, replacement)
.await
}
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -514,6 +584,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.list_system_config_entries().await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.list_system_config_entries().await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -542,6 +614,8 @@ impl<'a> SqlBackendRef<'a> {
.upsert_system_config_entry(key, value, description)
.await
}
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -553,6 +627,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.delete_system_config_value(key).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.delete_system_config_value(key).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -564,6 +640,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.read_admin_system_stats().await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.read_admin_system_stats().await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -578,6 +656,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.purge_admin_system_data(target).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.purge_admin_system_data(target).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -591,6 +671,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.export_admin_system_usage_aggregates().await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.export_admin_system_usage_aggregates().await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -635,6 +717,8 @@ impl<'a> SqlBackendRef<'a> {
)
.await
}
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
@@ -649,6 +733,8 @@ impl<'a> SqlBackendRef<'a> {
Self::Mysql(mysql) => mysql.purge_admin_request_bodies_batch(batch_size).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.purge_admin_request_bodies_batch(batch_size).await,
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Self::Disabled(_) => unreachable!("a SQL backend cannot exist without a driver"),
}
}
}
@@ -32,9 +32,9 @@ pub use mysql::MysqlBackend;
pub use postgres::PostgresBackend;
pub use read::DataReadRepositories;
pub use referrals::{
ReferralAdminStats, ReferralDataState, ReferralMutationStatus, ReferralRelationshipListQuery,
ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery,
ReferralRewardRecord, ReferralUserDashboard,
ReferralAdminStats, ReferralDataState, ReferralMutationStatus, ReferralReconciliationSummary,
ReferralRelationshipListQuery, ReferralRelationshipRecord, ReferralRewardConfig,
ReferralRewardListQuery, ReferralRewardRecord, ReferralUserDashboard,
};
#[cfg(feature = "sqlite")]
pub use sqlite::SqliteBackend;
@@ -52,6 +52,11 @@ enum SqlBackendRef<'a> {
Mysql(&'a MysqlBackend),
#[cfg(feature = "sqlite")]
Sqlite(&'a SqliteBackend),
// Keep the reference lifetime represented when this crate is built without
// any SQL driver features. The no-driver build still exposes the
// maintenance facade, but has no concrete backend variant to carry `'a`.
#[cfg(not(any(feature = "postgres", feature = "mysql", feature = "sqlite")))]
Disabled(std::marker::PhantomData<&'a ()>),
}
#[derive(Debug, Clone, Default)]
File diff suppressed because it is too large Load Diff
@@ -302,7 +302,7 @@ mod tests {
.await
.expect("sqlite migrations should run");
let value = serde_json::json!({"enabled": true});
let value = serde_json::json!("enabled");
let stored = backend
.upsert_system_config_entry("feature.local", &value, Some("local flag"))
.await
@@ -313,7 +313,40 @@ mod tests {
.find_system_config_value("feature.local")
.await
.expect("system config should read"),
Some(value)
Some(value.clone())
);
let replacement = serde_json::json!("disabled");
assert!(!backend
.compare_and_set_system_config_string_value("feature.local", "stale", "disabled")
.await
.expect("stale system config compare-and-set should complete"));
assert!(backend
.compare_and_set_system_config_string_value("feature.local", "enabled", "disabled")
.await
.expect("matching system config compare-and-set should complete"));
assert_eq!(
backend
.find_system_config_value("feature.local")
.await
.expect("updated system config should read"),
Some(replacement.clone())
);
sqlx::query("UPDATE system_configs SET value = ? WHERE key = ?")
.bind(r#""\u5bc6\u94a5""#)
.bind("feature.local")
.execute(backend.pool())
.await
.expect("legacy escaped JSON string should persist");
assert!(backend
.compare_and_set_system_config_string_value("feature.local", "密钥", "encrypted-value",)
.await
.expect("escaped JSON string compare-and-set should complete"));
assert_eq!(
backend
.find_system_config_value("feature.local")
.await
.expect("escaped JSON string replacement should read"),
Some(serde_json::json!("encrypted-value"))
);
assert_eq!(
backend
@@ -543,6 +576,22 @@ VALUES ('target-key-1', 'target-user-1', 'hash-target-key', 'target key', 1, 1)
let api_key_id_map =
BTreeMap::from([("source-key-1".to_string(), "target-key-1".to_string())]);
let validation_summary = backend
.import_admin_system_usage_aggregates(
&snapshot,
&user_id_map,
&api_key_id_map,
AdminSystemUsageAggregateImportMode::ValidateError,
)
.await
.expect("usage aggregates should validate");
assert_eq!(validation_summary.stats_daily.created, 1);
assert_eq!(validation_summary.stats_user_daily.created, 1);
assert_eq!(validation_summary.stats_daily_api_key.created, 1);
assert_eq!(sqlite_count(backend.pool(), "stats_daily").await, 0);
assert_eq!(sqlite_count(backend.pool(), "stats_user_daily").await, 0);
assert_eq!(sqlite_count(backend.pool(), "stats_daily_api_key").await, 0);
let summary = backend
.import_admin_system_usage_aggregates(
&snapshot,
@@ -170,9 +170,10 @@ fn should_skip_imported_aggregate(
match mode {
AdminSystemUsageAggregateImportMode::Skip => Ok(true),
AdminSystemUsageAggregateImportMode::Overwrite => Ok(false),
AdminSystemUsageAggregateImportMode::Error => Err(DataLayerError::InvalidInput(format!(
"{table} aggregate already exists for date_unix_secs={date_unix_secs}"
))),
AdminSystemUsageAggregateImportMode::Error
| AdminSystemUsageAggregateImportMode::ValidateError => Err(DataLayerError::InvalidInput(
format!("{table} aggregate already exists for date_unix_secs={date_unix_secs}"),
)),
}
}
@@ -392,7 +392,11 @@ ON DUPLICATE KEY UPDATE
add_aggregate_import_count(&mut summary.stats_daily_api_key, existing.is_some());
}
tx.commit().await.map_sql_err()?;
if mode == AdminSystemUsageAggregateImportMode::ValidateError {
tx.rollback().await.map_sql_err()?;
} else {
tx.commit().await.map_sql_err()?;
}
Ok(summary)
}
@@ -447,6 +451,34 @@ LIMIT 1
.transpose()
}
pub async fn compare_and_set_system_config_string_value(
&self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let now = current_unix_secs();
let replacement =
serialize_json_value(&serde_json::Value::String(replacement.to_string()))?;
let result = sqlx::query(
r#"
UPDATE system_configs
SET value = ?, updated_at = ?
WHERE `key` = ?
AND JSON_TYPE(value) = 'STRING'
AND BINARY JSON_UNQUOTE(value) = BINARY ?
"#,
)
.bind(replacement)
.bind(now as i64)
.bind(key)
.bind(expected)
.execute(self.pool())
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn upsert_system_config_value(
&self,
key: &str,
@@ -501,6 +533,7 @@ ORDER BY `key` ASC
INSERT INTO system_configs (id, `key`, value, description, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
`key` = VALUES(`key`),
value = VALUES(value),
description = COALESCE(VALUES(description), description),
updated_at = VALUES(updated_at)
@@ -18,6 +18,15 @@ WHERE key = $1
LIMIT 1
"#;
const COMPARE_AND_SET_SYSTEM_CONFIG_STRING_VALUE_SQL: &str = r#"
UPDATE system_configs
SET value = TO_JSON($3::text),
updated_at = NOW()
WHERE key = $1
AND JSON_TYPEOF(value) = 'string'
AND value #>> '{}' = $2
"#;
const UPSERT_SYSTEM_CONFIG_VALUE_SQL: &str = r#"
INSERT INTO system_configs (id, key, value, description, created_at, updated_at)
VALUES ($1, $2, $3, $4, NOW(), NOW())
@@ -437,7 +446,11 @@ SET api_key_name = EXCLUDED.api_key_name,
add_aggregate_import_count(&mut summary.stats_daily_api_key, existing.is_some());
}
tx.commit().await.map_postgres_err()?;
if mode == AdminSystemUsageAggregateImportMode::ValidateError {
tx.rollback().await.map_postgres_err()?;
} else {
tx.commit().await.map_postgres_err()?;
}
Ok(summary)
}
@@ -1157,6 +1170,22 @@ impl PostgresBackend {
.map_postgres_err()
}
pub async fn compare_and_set_system_config_string_value(
&self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(COMPARE_AND_SET_SYSTEM_CONFIG_STRING_VALUE_SQL)
.bind(key)
.bind(expected)
.bind(replacement)
.execute(self.pool())
.await
.map_postgres_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn upsert_system_config_value(
&self,
key: &str,
@@ -396,7 +396,11 @@ SET api_key_name = excluded.api_key_name,
add_aggregate_import_count(&mut summary.stats_daily_api_key, existing.is_some());
}
tx.commit().await.map_sql_err()?;
if mode == AdminSystemUsageAggregateImportMode::ValidateError {
tx.rollback().await.map_sql_err()?;
} else {
tx.commit().await.map_sql_err()?;
}
Ok(summary)
}
@@ -1074,6 +1078,35 @@ LIMIT 1
.transpose()
}
pub async fn compare_and_set_system_config_string_value(
&self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let now = current_unix_secs();
let replacement =
serialize_json_value(&serde_json::Value::String(replacement.to_string()))?;
let result = sqlx::query(
r#"
UPDATE system_configs
SET value = ?, updated_at = ?
WHERE key = ?
AND json_valid(value)
AND json_type(value) = 'text'
AND CAST(json_extract(value, '$') AS TEXT) = ? COLLATE BINARY
"#,
)
.bind(replacement)
.bind(now as i64)
.bind(key)
.bind(expected)
.execute(self.pool())
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn upsert_system_config_value(
&self,
key: &str,
@@ -7,7 +7,10 @@ use tracing::info;
// Generated by build.rs from schema/bootstrap/postgres.
pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str =
include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql"));
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260903000000;
// Keep data migrations after the privacy/security frontier executable on a
// fresh database. The bootstrap SQL is schema-only; stamping later data
// migrations would skip required cleanup/anonymization work.
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260821130000;
const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#"
SELECT COUNT(*)::BIGINT
@@ -3,12 +3,18 @@ use std::collections::{BTreeMap, BTreeSet};
#[cfg(all(feature = "postgres", feature = "sqlite"))]
use futures_util::TryStreamExt;
use serde_json::Value;
use sha2::{Digest, Sha256};
#[cfg(all(feature = "postgres", feature = "sqlite"))]
use sqlx::Acquire;
use sqlx::Row;
#[cfg(any(feature = "mysql", feature = "sqlite"))]
use sqlx::{Column, TypeInfo, ValueRef};
use aether_data_contracts::repository::candidates::{
sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data,
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
};
use crate::error::SqlResultExt;
use crate::{DataLayerError, DatabaseDriver, SqlDatabaseConfig};
@@ -46,6 +52,17 @@ use postgres::normalize_postgres_import_payload;
pub const EXPORT_FORMAT_VERSION: u32 = 2;
const MIN_SUPPORTED_EXPORT_FORMAT_VERSION: u32 = 1;
// JSONL imports are ultimately materialized as a `DataImportPlan`, so an
// attacker-controlled document can otherwise consume memory in both the input
// string and the parsed row/payload vectors. Keep these bounds deliberately
// separate from HTTP request limits: database exports may contain large body
// blobs, while still needing a finite parser budget. The total budget is kept
// below the gateway's 256 MiB request-body ceiling because parsing duplicates
// portions of the input in serde values and the import plan.
pub const MAX_JSONL_INPUT_BYTES: usize = 256 * 1024 * 1024;
pub const MAX_JSONL_LINE_BYTES: usize = 16 * 1024 * 1024;
pub const MAX_JSONL_RECORDS: usize = 1_000_000;
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
@@ -223,6 +240,14 @@ const AUXILIARY_TABLES: &[AuxiliaryTable] = &[
name: "usage_counter_deltas",
primary_key: &["id"],
},
AuxiliaryTable {
name: "usage_cost_reservations",
primary_key: &["reservation_token"],
},
AuxiliaryTable {
name: "usage_request_admissions",
primary_key: &["event_token"],
},
AuxiliaryTable {
name: "background_task_runs",
primary_key: &["id"],
@@ -447,6 +472,10 @@ impl DataImportPlan {
.map(Vec::as_slice)
.unwrap_or(&[])
}
fn imports_domain(&self, domain: ExportDomain) -> bool {
self.manifest.domains.contains(&domain)
}
}
#[derive(Debug, Clone, PartialEq)]
@@ -455,6 +484,83 @@ pub struct ExportRow {
pub payload: Value,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct IdentityImportScope {
user_ids: Vec<String>,
oauth_link_ids: Vec<String>,
oauth_provider_types: Vec<String>,
finalizes_oauth_links: bool,
validates_oauth_login_methods: bool,
}
impl IdentityImportScope {
fn from_plan(plan: &DataImportPlan) -> Result<Self, DataLayerError> {
let scope = Self {
user_ids: imported_payload_ids(plan, ExportDomain::Users, "id")?,
oauth_link_ids: imported_payload_ids(plan, ExportDomain::UserOAuthLinks, "id")?,
oauth_provider_types: imported_payload_ids(
plan,
ExportDomain::OAuthProviders,
"provider_type",
)?,
finalizes_oauth_links: plan.imports_domain(ExportDomain::UserOAuthLinks),
validates_oauth_login_methods: plan.imports_domain(ExportDomain::UserOAuthLinks)
|| plan.imports_domain(ExportDomain::OAuthProviders),
};
if let Some(provider_type) = scope.oauth_provider_types.iter().find(|provider_type| {
provider_type.is_empty()
|| provider_type.as_str() != provider_type.trim().to_ascii_lowercase()
}) {
return Err(DataLayerError::InvalidInput(format!(
"OAuth provider import has non-canonical provider_type '{provider_type}'"
)));
}
Ok(scope)
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct IdentityImportState {
affected_user_ids: BTreeSet<String>,
}
fn imported_payload_ids(
plan: &DataImportPlan,
domain: ExportDomain,
payload_field: &str,
) -> Result<Vec<String>, DataLayerError> {
plan.rows(domain)
.iter()
.map(|row| {
let payload_id = row
.payload
.as_object()
.and_then(|payload| payload.get(payload_field))
.and_then(Value::as_str)
.map(str::trim)
.filter(|id| !id.is_empty())
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"{} export row '{}' must contain a non-empty string {}",
domain.as_str(),
row.id,
payload_field
))
})?;
if payload_id != row.id {
return Err(DataLayerError::InvalidInput(format!(
"{} export row id '{}' does not match payload {} '{}'",
domain.as_str(),
row.id,
payload_field,
payload_id
)));
}
Ok(payload_id.to_string())
})
.collect()
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct DataCopyOptions {
pub omit_request_body_details: bool,
@@ -508,6 +614,139 @@ type PostgresImportColumns = BTreeMap<String, PostgresImportColumn>;
#[cfg(any(feature = "mysql", feature = "sqlite"))]
type ImportColumnNames = BTreeSet<String>;
const IMPORTED_CREDENTIAL_REVOKE_REASON: &str = "imported_credentials_revoked";
fn imported_credential_tombstone() -> String {
format!("{:x}", Sha256::digest(uuid::Uuid::new_v4().as_bytes()))
}
fn set_supported_import_value(
object: &mut serde_json::Map<String, Value>,
target_has_column: &impl Fn(&str) -> bool,
column: &str,
value: Value,
) {
if target_has_column(column) {
object.insert(column.to_string(), value);
}
}
fn deactivate_imported_credentials(
table_name: &str,
object: &mut serde_json::Map<String, Value>,
target_has_column: impl Fn(&str) -> bool,
) {
let table_name = table_name
.rsplit('.')
.next()
.unwrap_or(table_name)
.trim_matches(|ch| matches!(ch, '"' | '`'));
match table_name {
"users" => {
if object
.get("password_hash")
.is_some_and(|value| !value.is_null())
{
set_supported_import_value(
object,
&target_has_column,
"password_hash",
Value::String(format!(
"$aether-import-revoked${}",
imported_credential_tombstone()
)),
);
}
}
"api_keys" => {
if object.contains_key("key_hash") {
set_supported_import_value(
object,
&target_has_column,
"key_hash",
Value::String(imported_credential_tombstone()),
);
}
set_supported_import_value(object, &target_has_column, "key_encrypted", Value::Null);
set_supported_import_value(
object,
&target_has_column,
"status",
Value::String("disabled".to_string()),
);
set_supported_import_value(object, &target_has_column, "is_active", Value::Bool(false));
set_supported_import_value(object, &target_has_column, "is_locked", Value::Bool(true));
}
"management_tokens" => {
if object.contains_key("token_hash") {
set_supported_import_value(
object,
&target_has_column,
"token_hash",
Value::String(imported_credential_tombstone()),
);
}
set_supported_import_value(object, &target_has_column, "is_active", Value::Bool(false));
}
"user_sessions" => {
if object.contains_key("refresh_token_hash") {
set_supported_import_value(
object,
&target_has_column,
"refresh_token_hash",
Value::String(imported_credential_tombstone()),
);
}
set_supported_import_value(
object,
&target_has_column,
"prev_refresh_token_hash",
Value::Null,
);
set_supported_import_value(
object,
&target_has_column,
"revoked_at",
Value::from(chrono::Utc::now().timestamp()),
);
set_supported_import_value(
object,
&target_has_column,
"revoke_reason",
Value::String(IMPORTED_CREDENTIAL_REVOKE_REASON.to_string()),
);
}
"proxy_nodes" => {
set_supported_import_value(
object,
&target_has_column,
"tunnel_generation",
Value::String(uuid::Uuid::new_v4().to_string()),
);
set_supported_import_value(
object,
&target_has_column,
"tunnel_connected",
Value::Bool(false),
);
set_supported_import_value(
object,
&target_has_column,
"status",
Value::String("offline".to_string()),
);
set_supported_import_value(
object,
&target_has_column,
"active_connections",
Value::from(0),
);
}
_ => {}
}
}
const USAGE_REQUEST_BODY_DETAIL_COLUMNS: &[&str] = &[
"request_body",
"response_body",
@@ -691,6 +930,27 @@ pub fn encode_jsonl(records: &[DataExportRecord]) -> Result<String, DataLayerErr
for record in records {
let line = serde_json::to_string(record)
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
if line.len() > MAX_JSONL_LINE_BYTES {
return Err(DataLayerError::InvalidInput(format!(
"export JSONL record exceeds the {} byte line limit",
MAX_JSONL_LINE_BYTES
)));
}
let output_len = output
.len()
.checked_add(line.len())
.and_then(|length| length.checked_add(1))
.ok_or_else(|| {
DataLayerError::InvalidInput(
"export JSONL exceeds the input size limit".to_string(),
)
})?;
if output_len > MAX_JSONL_INPUT_BYTES {
return Err(DataLayerError::InvalidInput(format!(
"export JSONL exceeds the {} byte input limit",
MAX_JSONL_INPUT_BYTES
)));
}
output.push_str(&line);
output.push('\n');
}
@@ -698,11 +958,42 @@ pub fn encode_jsonl(records: &[DataExportRecord]) -> Result<String, DataLayerErr
}
pub fn decode_jsonl(input: &str) -> Result<Vec<DataExportRecord>, DataLayerError> {
decode_jsonl_with_limits(
input,
MAX_JSONL_INPUT_BYTES,
MAX_JSONL_LINE_BYTES,
MAX_JSONL_RECORDS,
)
}
fn decode_jsonl_with_limits(
input: &str,
max_input_bytes: usize,
max_line_bytes: usize,
max_records: usize,
) -> Result<Vec<DataExportRecord>, DataLayerError> {
if input.len() > max_input_bytes {
return Err(DataLayerError::InvalidInput(format!(
"export JSONL exceeds the {max_input_bytes} byte input limit"
)));
}
let mut records = Vec::new();
for (line_index, line) in input.lines().enumerate() {
if line.len() > max_line_bytes {
return Err(DataLayerError::InvalidInput(format!(
"export JSONL record on line {} exceeds the {max_line_bytes} byte line limit",
line_index + 1,
)));
}
if line.trim().is_empty() {
continue;
}
if records.len() >= max_records {
return Err(DataLayerError::InvalidInput(format!(
"export JSONL exceeds the {max_records} record limit"
)));
}
let record = serde_json::from_str::<DataExportRecord>(line).map_err(|err| {
DataLayerError::InvalidInput(format!(
"invalid export JSONL record on line {}: {err}",
@@ -745,6 +1036,12 @@ pub fn build_import_plan(input: &str) -> Result<DataImportPlan, DataLayerError>
}
pub fn validate_export_records(records: &[DataExportRecord]) -> Result<(), DataLayerError> {
if records.len() > MAX_JSONL_RECORDS {
return Err(DataLayerError::InvalidInput(format!(
"export JSONL exceeds the {} record limit",
MAX_JSONL_RECORDS
)));
}
let Some(DataExportRecord::Manifest { manifest }) = records.first() else {
return Err(DataLayerError::InvalidInput(
"export JSONL must start with a manifest record".to_string(),
@@ -1154,13 +1451,14 @@ async fn copy_postgres_sqlite_table(
let mut imported = 0usize;
while let Some(row) = rows.try_next().await.map_sql_err()? {
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
let object = payload.as_object().ok_or_else(|| {
let mut payload = row.try_get::<Value, _>("payload").map_sql_err()?;
let object = payload.as_object_mut().ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"postgres copy row for table '{}' did not produce a JSON object",
table.table_name
))
})?;
prepare_postgres_sqlite_copy_payload(table, object);
let mut query = sqlx::query(&target_sql);
for column in &table.columns {
let value = object.get(&column.sqlite.name).ok_or_else(|| {
@@ -1178,6 +1476,21 @@ async fn copy_postgres_sqlite_table(
Ok(imported)
}
#[cfg(all(feature = "postgres", feature = "sqlite"))]
fn prepare_postgres_sqlite_copy_payload(
table: &SchemaCopyTable,
object: &mut serde_json::Map<String, Value>,
) {
deactivate_imported_credentials(&table.table_name, object, |column_name| {
table
.columns
.iter()
.any(|column| column.sqlite.name == column_name)
});
sanitize_request_candidate_auxiliary_payload(&table.table_name, object);
sanitize_payment_security_payload(&table.table_name, object);
}
#[cfg(all(feature = "postgres", feature = "sqlite"))]
fn postgres_schema_copy_select_sql(table: &SchemaCopyTable) -> Result<String, DataLayerError> {
let table_sql = format!(
@@ -1743,10 +2056,85 @@ fn payload_with_table(payload: Value, table_name: &str) -> Result<Value, DataLay
DataLayerError::UnexpectedValue("export row payload must be a JSON object".to_string())
})?;
normalize_billing_payload(table_name, &mut object)?;
sanitize_request_candidate_auxiliary_payload(table_name, &mut object);
sanitize_payment_security_payload(table_name, &mut object);
object.insert("__table".to_string(), Value::String(table_name.to_string()));
Ok(Value::Object(object))
}
fn sanitize_request_candidate_auxiliary_payload(
table_name: &str,
object: &mut serde_json::Map<String, Value>,
) {
if table_name != "request_candidates" {
return;
}
object.insert("error_message".to_string(), Value::Null);
sanitize_request_candidate_auxiliary_string(
object,
"skip_reason",
sanitize_request_candidate_skip_reason,
);
sanitize_request_candidate_auxiliary_string(
object,
"error_type",
sanitize_request_candidate_error_type,
);
sanitize_request_candidate_auxiliary_json(
object,
"extra_data",
sanitize_request_candidate_extra_data,
);
sanitize_request_candidate_auxiliary_json(
object,
"required_capabilities",
sanitize_request_candidate_required_capabilities,
);
}
fn sanitize_payment_security_payload(
table_name: &str,
object: &mut serde_json::Map<String, Value>,
) {
match table_name {
"payment_orders" => {
object.insert("gateway_response".to_string(), Value::Null);
}
"payment_callbacks" => {
object.insert("payload".to_string(), Value::Null);
}
_ => {}
}
}
fn sanitize_request_candidate_auxiliary_string(
object: &mut serde_json::Map<String, Value>,
field: &str,
sanitize: fn(Option<String>) -> Option<String>,
) {
let value = object
.remove(field)
.and_then(|value| value.as_str().map(ToOwned::to_owned));
object.insert(
field.to_string(),
sanitize(value).map_or(Value::Null, Value::String),
);
}
fn sanitize_request_candidate_auxiliary_json(
object: &mut serde_json::Map<String, Value>,
field: &str,
sanitize: fn(Option<Value>) -> Option<Value>,
) {
let value = object.remove(field).and_then(|value| match value {
Value::Null => None,
Value::String(raw) => serde_json::from_str::<Value>(&raw).ok(),
value => Some(value),
});
object.insert(field.to_string(), sanitize(value).unwrap_or(Value::Null));
}
fn normalize_billing_payload(
table_name: &str,
object: &mut serde_json::Map<String, Value>,
@@ -1800,5 +2188,197 @@ fn domain_payload_table(
))
})?,
};
sanitize_request_candidate_auxiliary_payload(&table_name, &mut object);
sanitize_payment_security_payload(&table_name, &mut object);
Ok((table_name, Value::Object(object)))
}
#[cfg(test)]
mod payment_export_security_tests {
use serde_json::json;
use super::{domain_payload_table, payload_with_table, ExportRow};
#[test]
fn wallet_exports_and_imports_drop_payment_capabilities_and_raw_callbacks() {
let order = payload_with_table(
json!({
"id": "order-1",
"gateway_response": {
"client_secret": "pi_1_secret_replayable",
"_stripe_client_secret_encrypted": "ciphertext",
"customer": {"email": "[email protected]"},
"payment_url": "https://pay.example/checkout?token=secret",
},
}),
"payment_orders",
)
.expect("payment order export should sanitize");
assert!(order["gateway_response"].is_null());
let callback = ExportRow {
id: "payment_callbacks:callback-1".to_string(),
payload: json!({
"__table": "payment_callbacks",
"id": "callback-1",
"payload": {
"client_secret": "pi_1_secret_replayable",
"customer_email": "[email protected]",
},
}),
};
let (table, callback) = domain_payload_table(&callback, "wallet", Some("wallets"))
.expect("payment callback import should sanitize");
assert_eq!(table, "payment_callbacks");
assert!(callback["payload"].is_null());
let encoded = format!("{order}{callback}");
for forbidden in [
"client_secret",
"replayable",
"ciphertext",
"customer",
"[email protected]",
"token=secret",
] {
assert!(!encoded.contains(forbidden), "exported {forbidden}");
}
}
}
#[cfg(test)]
mod request_candidate_export_security_tests {
use serde_json::json;
#[cfg(all(feature = "postgres", feature = "sqlite"))]
use serde_json::Value;
use super::{domain_payload_table, payload_with_table, ExportRow};
#[cfg(all(feature = "postgres", feature = "sqlite"))]
use super::{
prepare_postgres_sqlite_copy_payload, PostgresImportColumn, SchemaCopyColumn,
SchemaCopyTable, SqliteCopyColumn,
};
#[test]
fn request_candidate_auxiliary_export_and_import_drop_sensitive_diagnostics() {
let raw = json!({
"id": "candidate-1",
"error_message": "Bearer export-secret",
"skip_reason": "secret skip reason",
"error_type": "secret error type",
"extra_data": "{\"upstream_url\":\"https://user:[email protected]/private/export-secret?token=secret\",\"unknown\":\"secret\",\"header_rules\":[{\"id\":\"secret-rule\",\"action\":\"set\",\"name\":\"authorization\",\"value\":\"secret\"}]}",
"required_capabilities": "{\"cache_1h\":\"true\",\"tenant_secret\":\"secret\"}"
});
let exported = payload_with_table(raw, "request_candidates")
.expect("candidate export payload should sanitize");
assert!(exported["error_message"].is_null());
assert_eq!(exported["skip_reason"], "unclassified_skip");
assert_eq!(exported["error_type"], "unclassified_error");
assert_eq!(
exported["extra_data"]["upstream_url"],
"https://example.com/"
);
assert_eq!(exported["extra_data"]["header_rules"]["count"], 1);
assert_eq!(exported["required_capabilities"]["cache_1h"], true);
let encoded = exported.to_string();
for sensitive in [
"export-secret",
"user:pass",
"secret-rule",
"authorization",
"tenant_secret",
] {
assert!(!encoded.contains(sensitive));
}
let imported_row = ExportRow {
id: "request_candidates:[\"candidate-1\"]".to_string(),
payload: json!({
"__table": "request_candidates",
"id": "candidate-1",
"error_message": "Bearer import-secret",
"extra_data": {"free_text": "import-secret"},
"required_capabilities": {"vision": 1, "secret": "import-secret"}
}),
};
let (table, imported) = domain_payload_table(&imported_row, "auxiliary", None)
.expect("candidate import payload should sanitize");
assert_eq!(table, "request_candidates");
assert!(imported["error_message"].is_null());
assert!(imported["extra_data"].is_null());
assert_eq!(imported["required_capabilities"], json!({"vision": true}));
assert!(!imported.to_string().contains("import-secret"));
}
#[cfg(all(feature = "postgres", feature = "sqlite"))]
#[test]
fn postgres_to_sqlite_fast_copy_sanitizes_request_candidate_diagnostics() {
let table = SchemaCopyTable {
table_name: "request_candidates".to_string(),
columns: [
"error_message",
"skip_reason",
"error_type",
"extra_data",
"required_capabilities",
]
.into_iter()
.map(|name| SchemaCopyColumn {
sqlite: SqliteCopyColumn {
name: name.to_string(),
declared_type: "TEXT".to_string(),
not_null: false,
has_default: false,
primary_key_position: 0,
},
postgres: PostgresImportColumn {
data_type: "text".to_string(),
udt_name: "text".to_string(),
is_nullable: true,
has_default: false,
},
})
.collect(),
};
let mut payload = json!({
"error_message": "Bearer fast-copy-secret",
"skip_reason": "fast-copy-secret",
"error_type": "fast-copy-secret",
"extra_data": {
"upstream_url": "https://user:[email protected]/private?token=fast-copy-secret",
"image_progress": {
"phase": "upstream_streaming",
"message": "fast-copy-secret"
}
},
"required_capabilities": {
"vision": 1,
"tenant_secret": "fast-copy-secret"
}
})
.as_object()
.cloned()
.expect("copy payload should be an object");
prepare_postgres_sqlite_copy_payload(&table, &mut payload);
assert!(payload["error_message"].is_null());
assert_eq!(payload["skip_reason"], "unclassified_skip");
assert_eq!(payload["error_type"], "unclassified_error");
assert_eq!(
payload["extra_data"]["upstream_url"],
"https://example.com/"
);
assert_eq!(
payload["extra_data"]["image_progress"],
json!({"phase": "upstream_streaming"})
);
assert_eq!(payload["required_capabilities"], json!({"vision": true}));
assert!(!Value::Object(payload)
.to_string()
.contains("fast-copy-secret"));
}
}
@@ -72,7 +72,9 @@ pub async fn import_mysql_plan(
pool: &crate::driver::mysql::MysqlPool,
plan: &DataImportPlan,
) -> Result<usize, DataLayerError> {
let identity_scope = IdentityImportScope::from_plan(plan)?;
let mut tx = pool.begin().await.map_sql_err()?;
let identity_state = capture_mysql_identity_import_state(&mut tx, &identity_scope).await?;
let mut imported = 0usize;
let mut column_cache = BTreeMap::<String, MysqlImportColumns>::new();
for domain in &plan.manifest.domains {
@@ -105,10 +107,210 @@ pub async fn import_mysql_plan(
imported = imported.saturating_add(1);
}
}
enforce_mysql_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?;
tx.commit().await.map_sql_err()?;
Ok(imported)
}
async fn capture_mysql_identity_import_state(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
scope: &IdentityImportScope,
) -> Result<IdentityImportState, DataLayerError> {
let mut affected_user_ids = if scope.finalizes_oauth_links {
scope.user_ids.iter().cloned().collect::<BTreeSet<_>>()
} else {
BTreeSet::new()
};
for link_id in &scope.oauth_link_ids {
if let Some(user_id) =
sqlx::query_scalar::<_, String>("SELECT user_id FROM user_oauth_links WHERE id = ?")
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
affected_user_ids.insert(user_id);
}
}
for provider_type in &scope.oauth_provider_types {
let user_ids = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM user_oauth_links WHERE provider_type = ?",
)
.bind(provider_type)
.fetch_all(&mut **tx)
.await
.map_sql_err()?;
affected_user_ids.extend(user_ids);
}
Ok(IdentityImportState { affected_user_ids })
}
async fn enforce_mysql_identity_import_invariants(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
scope: &IdentityImportScope,
mut state: IdentityImportState,
) -> Result<(), DataLayerError> {
for user_id in &scope.user_ids {
let auth_source =
sqlx::query_scalar::<_, String>("SELECT auth_source FROM users WHERE id = ? LIMIT 1")
.bind(&user_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"imported users row '{user_id}' did not produce a user record"
))
})?;
if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") {
return Err(DataLayerError::InvalidInput(format!(
"imported user '{user_id}' has unsupported auth_source '{auth_source}'"
)));
}
if auth_source == "oauth" {
sqlx::query("UPDATE users SET email_verified = 0 WHERE id = ?")
.bind(&user_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
}
}
for link_id in &scope.oauth_link_ids {
let user_id = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM user_oauth_links WHERE id = ? LIMIT 1",
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"imported OAuth link row '{link_id}' did not produce a link record"
))
})?;
state.affected_user_ids.insert(user_id);
}
for link_id in &scope.oauth_link_ids {
if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>(
r#"
SELECT imported.provider_type, imported.provider_user_id
FROM user_oauth_links imported
JOIN user_oauth_links duplicate
ON duplicate.provider_type = imported.provider_type
AND duplicate.provider_user_id = imported.provider_user_id
AND duplicate.id <> imported.id
WHERE imported.id = ?
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once"
)));
}
}
for link_id in &scope.oauth_link_ids {
if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>(
r#"
SELECT imported.user_id, imported.provider_type
FROM user_oauth_links imported
JOIN user_oauth_links duplicate
ON duplicate.user_id = imported.user_id
AND duplicate.provider_type = imported.provider_type
AND duplicate.id <> imported.id
WHERE imported.id = ?
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import links user '{user_id}' to provider '{provider_type}' more than once"
)));
}
}
for link_id in &scope.oauth_link_ids {
if let Some(invalid_id) = sqlx::query_scalar::<_, String>(
r#"
SELECT links.id
FROM user_oauth_links links
LEFT JOIN users ON users.id = links.user_id
LEFT JOIN oauth_providers providers ON providers.provider_type = links.provider_type
WHERE links.id = ?
AND (
users.id IS NULL
OR providers.provider_type IS NULL
OR BINARY links.provider_type <> BINARY LOWER(TRIM(links.provider_type))
OR links.provider_type = ''
OR BINARY links.provider_user_id <> BINARY TRIM(links.provider_user_id)
OR links.provider_user_id = ''
OR BINARY providers.provider_type <> BINARY LOWER(TRIM(providers.provider_type))
)
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import produced invalid or orphaned link '{invalid_id}'"
)));
}
}
if !scope.validates_oauth_login_methods {
return Ok(());
}
for user_id in state.affected_user_ids {
if sqlx::query_scalar::<_, String>(
r#"
SELECT users.id
FROM users
WHERE users.id = ?
AND users.auth_source = 'oauth'
AND users.is_active = 1
AND users.is_deleted = 0
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links links
JOIN oauth_providers providers ON providers.provider_type = links.provider_type
WHERE links.user_id = users.id
AND providers.is_enabled = 1
AND BINARY links.provider_type = BINARY LOWER(TRIM(links.provider_type))
AND BINARY links.provider_user_id = BINARY TRIM(links.provider_user_id)
AND links.provider_user_id <> ''
)
LIMIT 1
"#,
)
.bind(&user_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.is_some()
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import would leave active user '{user_id}' without an enabled identity binding"
)));
}
}
Ok(())
}
fn mysql_domain_table(
domain: ExportDomain,
) -> Result<(&'static str, &'static str), DataLayerError> {
@@ -262,7 +464,11 @@ async fn import_mysql_row(
row: &ExportRow,
target_columns: &MysqlImportColumns,
) -> Result<(), DataLayerError> {
let object = filter_import_payload("mysql", table_name, domain, row, &target_columns.names)?;
let mut object =
filter_import_payload("mysql", table_name, domain, row, &target_columns.names)?;
deactivate_imported_credentials(table_name, &mut object, |column_name| {
target_columns.names.contains(column_name)
});
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
for primary_key in &target_columns.primary_key {
@@ -67,7 +67,9 @@ pub async fn import_postgres_plan(
pool: &crate::driver::postgres::PostgresPool,
plan: &DataImportPlan,
) -> Result<usize, DataLayerError> {
let identity_scope = IdentityImportScope::from_plan(plan)?;
let mut tx = pool.begin().await.map_sql_err()?;
let identity_state = capture_postgres_identity_import_state(&mut tx, &identity_scope).await?;
let mut imported = 0usize;
let mut column_cache = BTreeMap::<String, PostgresImportColumns>::new();
for domain in &plan.manifest.domains {
@@ -113,6 +115,7 @@ pub async fn import_postgres_plan(
imported = imported.saturating_add(1);
}
}
enforce_postgres_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?;
if !plan.rows(ExportDomain::Auxiliary).is_empty() {
reset_postgres_auxiliary_sequences(&mut tx).await?;
}
@@ -120,6 +123,207 @@ pub async fn import_postgres_plan(
Ok(imported)
}
async fn capture_postgres_identity_import_state(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
scope: &IdentityImportScope,
) -> Result<IdentityImportState, DataLayerError> {
let mut affected_user_ids = if scope.finalizes_oauth_links {
scope.user_ids.iter().cloned().collect::<BTreeSet<_>>()
} else {
BTreeSet::new()
};
for link_id in &scope.oauth_link_ids {
if let Some(user_id) = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM public.user_oauth_links WHERE id = $1",
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
affected_user_ids.insert(user_id);
}
}
for provider_type in &scope.oauth_provider_types {
let user_ids = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM public.user_oauth_links WHERE provider_type = $1",
)
.bind(provider_type)
.fetch_all(&mut **tx)
.await
.map_sql_err()?;
affected_user_ids.extend(user_ids);
}
Ok(IdentityImportState { affected_user_ids })
}
async fn enforce_postgres_identity_import_invariants(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
scope: &IdentityImportScope,
mut state: IdentityImportState,
) -> Result<(), DataLayerError> {
for user_id in &scope.user_ids {
let auth_source = sqlx::query_scalar::<_, String>(
"SELECT auth_source::text FROM public.users WHERE id = $1 LIMIT 1",
)
.bind(&user_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"imported users row '{user_id}' did not produce a user record"
))
})?;
if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") {
return Err(DataLayerError::InvalidInput(format!(
"imported user '{user_id}' has unsupported auth_source '{auth_source}'"
)));
}
if auth_source == "oauth" {
sqlx::query("UPDATE public.users SET email_verified = FALSE WHERE id = $1")
.bind(&user_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
}
}
for link_id in &scope.oauth_link_ids {
let user_id = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM public.user_oauth_links WHERE id = $1 LIMIT 1",
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"imported OAuth link row '{link_id}' did not produce a link record"
))
})?;
state.affected_user_ids.insert(user_id);
}
for link_id in &scope.oauth_link_ids {
if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>(
r#"
SELECT imported.provider_type, imported.provider_user_id
FROM public.user_oauth_links imported
JOIN public.user_oauth_links duplicate
ON duplicate.provider_type = imported.provider_type
AND duplicate.provider_user_id = imported.provider_user_id
AND duplicate.id <> imported.id
WHERE imported.id = $1
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once"
)));
}
}
for link_id in &scope.oauth_link_ids {
if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>(
r#"
SELECT imported.user_id, imported.provider_type
FROM public.user_oauth_links imported
JOIN public.user_oauth_links duplicate
ON duplicate.user_id = imported.user_id
AND duplicate.provider_type = imported.provider_type
AND duplicate.id <> imported.id
WHERE imported.id = $1
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import links user '{user_id}' to provider '{provider_type}' more than once"
)));
}
}
for link_id in &scope.oauth_link_ids {
if let Some(invalid_id) = sqlx::query_scalar::<_, String>(
r#"
SELECT links.id
FROM public.user_oauth_links links
LEFT JOIN public.users users ON users.id = links.user_id
LEFT JOIN public.oauth_providers providers ON providers.provider_type = links.provider_type
WHERE links.id = $1
AND (
users.id IS NULL
OR providers.provider_type IS NULL
OR links.provider_type <> LOWER(TRIM(links.provider_type))
OR links.provider_type = ''
OR links.provider_user_id <> TRIM(links.provider_user_id)
OR links.provider_user_id = ''
OR providers.provider_type <> LOWER(TRIM(providers.provider_type))
)
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import produced invalid or orphaned link '{invalid_id}'"
)));
}
}
if !scope.validates_oauth_login_methods {
return Ok(());
}
for user_id in state.affected_user_ids {
if sqlx::query_scalar::<_, String>(
r#"
SELECT users.id
FROM public.users users
WHERE users.id = $1
AND users.auth_source = 'oauth'::public.authsource
AND users.is_active IS TRUE
AND users.is_deleted IS FALSE
AND NOT EXISTS (
SELECT 1
FROM public.user_oauth_links links
JOIN public.oauth_providers providers ON providers.provider_type = links.provider_type
WHERE links.user_id = users.id
AND providers.is_enabled IS TRUE
AND links.provider_type = LOWER(TRIM(links.provider_type))
AND links.provider_user_id = TRIM(links.provider_user_id)
AND links.provider_user_id <> ''
)
LIMIT 1
"#,
)
.bind(&user_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.is_some()
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import would leave active user '{user_id}' without an enabled identity binding"
)));
}
}
Ok(())
}
async fn reset_postgres_auxiliary_sequences(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
) -> Result<(), DataLayerError> {
@@ -359,7 +563,13 @@ async fn export_postgres_wallet_records(
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
for row in rows {
let id = row.try_get::<String, _>("export_id").map_sql_err()?;
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
let mut payload = row.try_get::<Value, _>("payload").map_sql_err()?;
let object = payload.as_object_mut().ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"wallet export row in table '{export_table}' must be an object"
))
})?;
sanitize_payment_security_payload(export_table, object);
records.push(DataExportRecord::row(
ExportDomain::Wallets,
format!("{export_table}:{id}"),
@@ -378,7 +588,10 @@ async fn import_postgres_row(
row: &ExportRow,
target_columns: &PostgresImportColumns,
) -> Result<(), DataLayerError> {
let object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?;
let mut object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?;
deactivate_imported_credentials(table_name, &mut object, |column_name| {
target_columns.contains_key(column_name)
});
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
let column_sql = columns
@@ -66,7 +66,9 @@ pub async fn import_sqlite_plan(
pool: &crate::driver::sqlite::SqlitePool,
plan: &DataImportPlan,
) -> Result<usize, DataLayerError> {
let identity_scope = IdentityImportScope::from_plan(plan)?;
let mut tx = pool.begin().await.map_sql_err()?;
let identity_state = capture_sqlite_identity_import_state(&mut tx, &identity_scope).await?;
let mut imported = 0usize;
let mut column_cache = BTreeMap::<String, SqliteImportColumns>::new();
for domain in &plan.manifest.domains {
@@ -99,10 +101,210 @@ pub async fn import_sqlite_plan(
imported = imported.saturating_add(1);
}
}
enforce_sqlite_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?;
tx.commit().await.map_sql_err()?;
Ok(imported)
}
async fn capture_sqlite_identity_import_state(
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
scope: &IdentityImportScope,
) -> Result<IdentityImportState, DataLayerError> {
let mut affected_user_ids = if scope.finalizes_oauth_links {
scope.user_ids.iter().cloned().collect::<BTreeSet<_>>()
} else {
BTreeSet::new()
};
for link_id in &scope.oauth_link_ids {
if let Some(user_id) =
sqlx::query_scalar::<_, String>("SELECT user_id FROM user_oauth_links WHERE id = ?")
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
affected_user_ids.insert(user_id);
}
}
for provider_type in &scope.oauth_provider_types {
let user_ids = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM user_oauth_links WHERE provider_type = ?",
)
.bind(provider_type)
.fetch_all(&mut **tx)
.await
.map_sql_err()?;
affected_user_ids.extend(user_ids);
}
Ok(IdentityImportState { affected_user_ids })
}
async fn enforce_sqlite_identity_import_invariants(
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
scope: &IdentityImportScope,
mut state: IdentityImportState,
) -> Result<(), DataLayerError> {
for user_id in &scope.user_ids {
let auth_source =
sqlx::query_scalar::<_, String>("SELECT auth_source FROM users WHERE id = ? LIMIT 1")
.bind(&user_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"imported users row '{user_id}' did not produce a user record"
))
})?;
if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") {
return Err(DataLayerError::InvalidInput(format!(
"imported user '{user_id}' has unsupported auth_source '{auth_source}'"
)));
}
if auth_source == "oauth" {
sqlx::query("UPDATE users SET email_verified = 0 WHERE id = ?")
.bind(&user_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
}
}
for link_id in &scope.oauth_link_ids {
let user_id = sqlx::query_scalar::<_, String>(
"SELECT user_id FROM user_oauth_links WHERE id = ? LIMIT 1",
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"imported OAuth link row '{link_id}' did not produce a link record"
))
})?;
state.affected_user_ids.insert(user_id);
}
for link_id in &scope.oauth_link_ids {
if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>(
r#"
SELECT imported.provider_type, imported.provider_user_id
FROM user_oauth_links imported
JOIN user_oauth_links duplicate
ON duplicate.provider_type = imported.provider_type
AND duplicate.provider_user_id = imported.provider_user_id
AND duplicate.id <> imported.id
WHERE imported.id = ?
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once"
)));
}
}
for link_id in &scope.oauth_link_ids {
if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>(
r#"
SELECT imported.user_id, imported.provider_type
FROM user_oauth_links imported
JOIN user_oauth_links duplicate
ON duplicate.user_id = imported.user_id
AND duplicate.provider_type = imported.provider_type
AND duplicate.id <> imported.id
WHERE imported.id = ?
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import links user '{user_id}' to provider '{provider_type}' more than once"
)));
}
}
for link_id in &scope.oauth_link_ids {
if let Some(invalid_id) = sqlx::query_scalar::<_, String>(
r#"
SELECT links.id
FROM user_oauth_links links
LEFT JOIN users ON users.id = links.user_id
LEFT JOIN oauth_providers providers ON providers.provider_type = links.provider_type
WHERE links.id = ?
AND (
users.id IS NULL
OR providers.provider_type IS NULL
OR links.provider_type <> LOWER(TRIM(links.provider_type))
OR links.provider_type = ''
OR links.provider_user_id <> TRIM(links.provider_user_id)
OR links.provider_user_id = ''
OR providers.provider_type <> LOWER(TRIM(providers.provider_type))
)
LIMIT 1
"#,
)
.bind(link_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import produced invalid or orphaned link '{invalid_id}'"
)));
}
}
if !scope.validates_oauth_login_methods {
return Ok(());
}
for user_id in state.affected_user_ids {
if sqlx::query_scalar::<_, String>(
r#"
SELECT users.id
FROM users
WHERE users.id = ?
AND users.auth_source = 'oauth'
AND users.is_active = 1
AND users.is_deleted = 0
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links links
JOIN oauth_providers providers ON providers.provider_type = links.provider_type
WHERE links.user_id = users.id
AND providers.is_enabled = 1
AND links.provider_type = LOWER(TRIM(links.provider_type))
AND links.provider_user_id = TRIM(links.provider_user_id)
AND links.provider_user_id <> ''
)
LIMIT 1
"#,
)
.bind(&user_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.is_some()
{
return Err(DataLayerError::InvalidInput(format!(
"OAuth import would leave active user '{user_id}' without an enabled identity binding"
)));
}
}
Ok(())
}
fn sqlite_domain_table(
domain: ExportDomain,
) -> Result<(&'static str, &'static str), DataLayerError> {
@@ -262,7 +464,11 @@ async fn import_sqlite_row(
row: &ExportRow,
target_columns: &SqliteImportColumns,
) -> Result<(), DataLayerError> {
let object = filter_import_payload("sqlite", table_name, domain, row, &target_columns.names)?;
let mut object =
filter_import_payload("sqlite", table_name, domain, row, &target_columns.names)?;
deactivate_imported_credentials(table_name, &mut object, |column_name| {
target_columns.names.contains(column_name)
});
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
let column_sql = columns
@@ -1,16 +1,17 @@
use std::collections::{BTreeMap, BTreeSet};
use serde_json::json;
use serde_json::{json, Value};
use super::{
build_import_plan, decode_jsonl, encode_jsonl, export_mysql_core_jsonl, export_mysql_jsonl,
export_postgres_core_jsonl, export_sqlite_core_jsonl, filter_import_payload,
import_mysql_jsonl, import_postgres_jsonl, import_sqlite_jsonl, mysql_core_export_domains,
normalize_imported_binary, normalize_imported_integer_timestamp,
normalize_postgres_import_payload, postgres_bytea_json_value, postgres_core_export_domains,
sqlite_core_export_domains, sqlite_schema_copy_insert_sql, DataExportManifest,
DataExportRecord, DataImportPlan, ExportDomain, ExportRow, PostgresImportColumn,
SchemaCopyColumn, SchemaCopyTable, SqliteCopyColumn, AUXILIARY_TABLES,
build_import_plan, deactivate_imported_credentials, decode_jsonl, decode_jsonl_with_limits,
encode_jsonl, export_mysql_core_jsonl, export_mysql_jsonl, export_postgres_core_jsonl,
export_sqlite_core_jsonl, filter_import_payload, import_mysql_jsonl, import_postgres_jsonl,
import_sqlite_jsonl, mysql_core_export_domains, normalize_imported_binary,
normalize_imported_integer_timestamp, normalize_postgres_import_payload,
postgres_bytea_json_value, postgres_core_export_domains, sqlite_core_export_domains,
sqlite_schema_copy_insert_sql, DataExportManifest, DataExportRecord, DataImportPlan,
ExportDomain, ExportRow, PostgresImportColumn, SchemaCopyColumn, SchemaCopyTable,
SqliteCopyColumn, AUXILIARY_TABLES,
};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::lifecycle::migrate::{
@@ -181,6 +182,25 @@ not-json"#,
assert!(err.to_string().contains("line 2"));
}
#[test]
fn jsonl_rejects_input_and_record_limits_before_materializing_rows() {
let oversized = "x".repeat(11);
let err = decode_jsonl_with_limits(&oversized, 10, 100, 10)
.expect_err("input byte limit should be enforced");
assert!(err.to_string().contains("10 byte input limit"));
let manifest = r#"{"record_type":"manifest","manifest":{"format_version":1,"created_at_unix_secs":1,"source_driver":null,"domains":[]}}"#;
let err = decode_jsonl_with_limits(manifest, usize::MAX, 10, 10)
.expect_err("line byte limit should be enforced");
assert!(err.to_string().contains("byte line limit"));
let row = r#"{"record_type":"manifest","manifest":{"format_version":1,"created_at_unix_secs":1,"source_driver":null,"domains":[]}}"#;
let input = format!("{row}\n{row}\n");
let err = decode_jsonl_with_limits(&input, usize::MAX, usize::MAX, 1)
.expect_err("record limit should be enforced");
assert!(err.to_string().contains("1 record limit"));
}
#[test]
fn jsonl_rejects_duplicate_domain_ids() {
let records = vec![
@@ -409,6 +429,231 @@ fn mysql_and_sqlite_import_payloads_ignore_unknown_null_columns() {
);
}
#[test]
fn imported_identity_credentials_are_replaced_with_disabled_tombstones() {
let columns = BTreeSet::from([
"password_hash".to_string(),
"key_hash".to_string(),
"key_encrypted".to_string(),
"status".to_string(),
"is_active".to_string(),
"is_locked".to_string(),
"token_hash".to_string(),
"refresh_token_hash".to_string(),
"prev_refresh_token_hash".to_string(),
"revoked_at".to_string(),
"revoke_reason".to_string(),
]);
let mut user =
serde_json::Map::from_iter([("password_hash".to_string(), json!("$2b$12$backup-hash"))]);
deactivate_imported_credentials("users", &mut user, |column| columns.contains(column));
assert_ne!(user["password_hash"], json!("$2b$12$backup-hash"));
assert!(user["password_hash"]
.as_str()
.is_some_and(|value| value.starts_with("$aether-import-revoked$")));
let mut api_key = serde_json::Map::from_iter([
("key_hash".to_string(), json!("backup-key-hash")),
("key_encrypted".to_string(), json!("backup-ciphertext")),
("is_active".to_string(), json!(true)),
("is_locked".to_string(), json!(false)),
("status".to_string(), json!("active")),
]);
deactivate_imported_credentials("api_keys", &mut api_key, |column| columns.contains(column));
assert_ne!(api_key["key_hash"], json!("backup-key-hash"));
assert_eq!(api_key["key_encrypted"], Value::Null);
assert_eq!(api_key["is_active"], json!(false));
assert_eq!(api_key["is_locked"], json!(true));
assert_eq!(api_key["status"], json!("disabled"));
let mut token = serde_json::Map::from_iter([
("token_hash".to_string(), json!("backup-token-hash")),
("is_active".to_string(), json!(true)),
]);
deactivate_imported_credentials("management_tokens", &mut token, |column| {
columns.contains(column)
});
assert_ne!(token["token_hash"], json!("backup-token-hash"));
assert_eq!(token["is_active"], json!(false));
let mut session = serde_json::Map::from_iter([
(
"refresh_token_hash".to_string(),
json!("backup-refresh-hash"),
),
(
"prev_refresh_token_hash".to_string(),
json!("backup-previous-hash"),
),
("revoked_at".to_string(), Value::Null),
("revoke_reason".to_string(), Value::Null),
]);
deactivate_imported_credentials("user_sessions", &mut session, |column| {
columns.contains(column)
});
assert_ne!(session["refresh_token_hash"], json!("backup-refresh-hash"));
assert_eq!(session["prev_refresh_token_hash"], Value::Null);
assert!(session["revoked_at"].as_i64().is_some());
assert_eq!(
session["revoke_reason"],
json!("imported_credentials_revoked")
);
}
#[test]
fn imported_proxy_nodes_receive_a_new_offline_tunnel_generation() {
let columns = BTreeSet::from([
"tunnel_generation".to_string(),
"tunnel_connected".to_string(),
"status".to_string(),
"active_connections".to_string(),
]);
let mut node = serde_json::Map::from_iter([
(
"tunnel_generation".to_string(),
json!("backup-tunnel-generation"),
),
("tunnel_connected".to_string(), json!(true)),
("status".to_string(), json!("online")),
("active_connections".to_string(), json!(42)),
(
"proxy_metadata".to_string(),
json!({"tunnel_security": {"encryption_key": "preserved-psk"}}),
),
]);
deactivate_imported_credentials("public.proxy_nodes", &mut node, |column| {
columns.contains(column)
});
let generation = node["tunnel_generation"]
.as_str()
.expect("imported node generation should be a string");
assert_ne!(generation, "backup-tunnel-generation");
assert!(uuid::Uuid::parse_str(generation).is_ok());
assert_eq!(node["tunnel_connected"], json!(false));
assert_eq!(node["status"], json!("offline"));
assert_eq!(node["active_connections"], json!(0));
assert_eq!(
node["proxy_metadata"]["tunnel_security"]["encryption_key"],
json!("preserved-psk")
);
}
#[tokio::test]
async fn sqlite_import_rotates_proxy_node_generations_and_clears_online_state() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO proxy_nodes (
id, tunnel_generation, name, ip, port, status, active_connections,
tunnel_mode, tunnel_connected, proxy_metadata, created_at, updated_at
) VALUES (
'import-existing-node', 'target-live-generation', 'existing node', '127.0.0.1',
8080, 'online', 9, 1, 1, '{"target":"metadata"}', 1, 1
)
"#,
)
.execute(&pool)
.await
.expect("existing proxy node should insert");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::ProxyNodes],
)),
DataExportRecord::row(
ExportDomain::ProxyNodes,
"import-existing-node",
json!({
"id": "import-existing-node",
"tunnel_generation": "backup-stale-generation",
"name": "restored existing node",
"ip": "127.0.0.1",
"port": 8080,
"status": "online",
"active_connections": 42,
"tunnel_mode": true,
"tunnel_connected": true,
"proxy_metadata": {
"tunnel_security": {"encryption_key": "preserved-psk"}
},
"created_at": 1,
"updated_at": 2
}),
),
DataExportRecord::row(
ExportDomain::ProxyNodes,
"import-legacy-node",
json!({
"id": "import-legacy-node",
"name": "legacy backup node",
"ip": "127.0.0.2",
"port": 8081,
"status": "online",
"active_connections": 7,
"tunnel_mode": true,
"tunnel_connected": true,
"created_at": 1,
"updated_at": 2
}),
),
])
.expect("proxy node import fixture should encode");
assert_eq!(
import_sqlite_jsonl(&pool, &encoded)
.await
.expect("proxy nodes should import"),
2
);
let restored = sqlx::query_as::<_, (String, String, bool, i32, Option<String>)>(
r#"
SELECT tunnel_generation, status, tunnel_connected, active_connections, proxy_metadata
FROM proxy_nodes
WHERE id = 'import-existing-node'
"#,
)
.fetch_one(&pool)
.await
.expect("restored proxy node should load");
assert_ne!(restored.0, "target-live-generation");
assert_ne!(restored.0, "backup-stale-generation");
assert!(uuid::Uuid::parse_str(&restored.0).is_ok());
assert_eq!(restored.1, "offline");
assert!(!restored.2);
assert_eq!(restored.3, 0);
assert_eq!(
restored
.4
.as_deref()
.and_then(|value| serde_json::from_str::<Value>(value).ok())
.and_then(|value| value["tunnel_security"]["encryption_key"]
.as_str()
.map(str::to_string)),
Some("preserved-psk".to_string())
);
let legacy_generation: String = sqlx::query_scalar(
"SELECT tunnel_generation FROM proxy_nodes WHERE id = 'import-legacy-node'",
)
.fetch_one(&pool)
.await
.expect("legacy imported proxy node should load");
assert!(uuid::Uuid::parse_str(&legacy_generation).is_ok());
}
#[test]
fn postgres_to_sqlite_copy_uses_primary_key_upsert_instead_of_replace() {
let table = SchemaCopyTable {
@@ -639,6 +884,548 @@ async fn sqlite_import_rolls_back_rows_after_late_failure() {
assert_eq!(count, 0);
}
#[tokio::test]
async fn sqlite_users_import_fails_closed_for_oauth_email_verification() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::Users],
)),
DataExportRecord::row(
ExportDomain::Users,
"oauth-user",
json!({
"id": "oauth-user",
"email": "[email protected]",
"email_verified": true,
"username": "oauth-user",
"role": "user",
"auth_source": "oauth",
"created_at": 1,
"updated_at": 1
}),
),
DataExportRecord::row(
ExportDomain::Users,
"local-user",
json!({
"id": "local-user",
"email": "[email protected]",
"email_verified": true,
"username": "local-user",
"role": "user",
"auth_source": "local",
"created_at": 1,
"updated_at": 1
}),
),
])
.expect("users fixture should encode");
assert_eq!(
import_sqlite_jsonl(&pool, &encoded)
.await
.expect("users-only staged restore should succeed without OAuth links"),
2
);
let verification =
sqlx::query_as::<_, (String, i64)>("SELECT id, email_verified FROM users ORDER BY id ASC")
.fetch_all(&pool)
.await
.expect("verification state should load");
assert_eq!(
verification,
vec![("local-user".to_string(), 1), ("oauth-user".to_string(), 0)]
);
}
#[tokio::test]
async fn sqlite_users_and_providers_can_restore_before_oauth_links() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::Users, ExportDomain::OAuthProviders],
)),
DataExportRecord::row(
ExportDomain::Users,
"oauth-user",
json!({
"id": "oauth-user",
"email": "[email protected]",
"email_verified": true,
"username": "oauth-user",
"role": "user",
"auth_source": "oauth",
"created_at": 1,
"updated_at": 1
}),
),
DataExportRecord::row(
ExportDomain::OAuthProviders,
"linuxdo",
json!({
"provider_type": "linuxdo",
"display_name": "Linux.do",
"client_id": "client",
"redirect_uri": "https://gateway.example.test/oauth/callback",
"frontend_callback_url": "https://app.example.test/auth/callback",
"is_enabled": true,
"created_at": 1,
"updated_at": 1
}),
),
])
.expect("staged identity fixture should encode");
assert_eq!(
import_sqlite_jsonl(&pool, &encoded)
.await
.expect("users and Providers should restore before links"),
2
);
let email_verified: i64 =
sqlx::query_scalar("SELECT email_verified FROM users WHERE id = 'oauth-user'")
.fetch_one(&pool)
.await
.expect("staged OAuth user should load");
assert_eq!(email_verified, 0);
}
#[tokio::test]
async fn sqlite_oauth_link_import_rolls_back_without_enabled_login_binding() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::raw_sql(
r#"
INSERT INTO users (
id, email, email_verified, username, role, auth_source,
is_active, is_deleted, created_at, updated_at
) VALUES (
'oauth-user', 'oauth@example.test', 0, 'oauth-user', 'user', 'oauth',
1, 0, 1, 1
);
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES (
'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback',
'https://app.example.test/auth/callback', 0, 1, 1
);
"#,
)
.execute(&pool)
.await
.expect("OAuth fixtures should seed");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::UserOAuthLinks],
)),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"link-disabled",
json!({
"id": "link-disabled",
"user_id": "oauth-user",
"provider_type": "linuxdo",
"provider_user_id": "subject-1",
"linked_at": 1
}),
),
])
.expect("OAuth link fixture should encode");
let err = import_sqlite_jsonl(&pool, &encoded)
.await
.expect_err("disabled-only OAuth binding should fail");
assert!(err
.to_string()
.contains("without an enabled identity binding"));
let link_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links")
.fetch_one(&pool)
.await
.expect("rolled-back OAuth link count should load");
assert_eq!(link_count, 0);
}
#[tokio::test]
async fn sqlite_oauth_provider_import_rolls_back_if_it_removes_last_enabled_binding() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::raw_sql(
r#"
INSERT INTO users (
id, email, email_verified, username, role, auth_source,
is_active, is_deleted, created_at, updated_at
) VALUES (
'oauth-user', 'oauth@example.test', 0, 'oauth-user', 'user', 'oauth',
1, 0, 1, 1
);
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES (
'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback',
'https://app.example.test/auth/callback', 1, 1, 1
);
INSERT INTO user_oauth_links (
id, user_id, provider_type, provider_user_id, linked_at
) VALUES (
'existing-link', 'oauth-user', 'linuxdo', 'subject-1', 1
);
"#,
)
.execute(&pool)
.await
.expect("OAuth fixtures should seed");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::OAuthProviders],
)),
DataExportRecord::row(
ExportDomain::OAuthProviders,
"linuxdo",
json!({
"provider_type": "linuxdo",
"display_name": "Linux.do disabled",
"client_id": "client",
"redirect_uri": "https://gateway.example.test/oauth/callback",
"frontend_callback_url": "https://app.example.test/auth/callback",
"is_enabled": false,
"created_at": 1,
"updated_at": 2
}),
),
])
.expect("disabled Provider fixture should encode");
let err = import_sqlite_jsonl(&pool, &encoded)
.await
.expect_err("disabling the last OAuth login method should fail");
assert!(err
.to_string()
.contains("without an enabled identity binding"));
let provider = sqlx::query_as::<_, (String, i64)>(
"SELECT display_name, is_enabled FROM oauth_providers WHERE provider_type = 'linuxdo'",
)
.fetch_one(&pool)
.await
.expect("rolled-back Provider should load");
assert_eq!(provider, ("Linux.do".to_string(), 1));
}
#[tokio::test]
async fn sqlite_oauth_link_reassignment_rolls_back_if_old_owner_loses_last_binding() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::raw_sql(
r#"
INSERT INTO users (id, username, role, auth_source, created_at, updated_at) VALUES
('oauth-owner', 'oauth-owner', 'user', 'oauth', 1, 1),
('local-target', 'local-target', 'user', 'local', 1, 1);
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES (
'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback',
'https://app.example.test/auth/callback', 1, 1, 1
);
INSERT INTO user_oauth_links (
id, user_id, provider_type, provider_user_id, linked_at
) VALUES (
'reassigned-link', 'oauth-owner', 'linuxdo', 'subject-1', 1
);
"#,
)
.execute(&pool)
.await
.expect("OAuth reassignment fixtures should seed");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::UserOAuthLinks],
)),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"reassigned-link",
json!({
"id": "reassigned-link",
"user_id": "local-target",
"provider_type": "linuxdo",
"provider_user_id": "subject-1",
"linked_at": 2
}),
),
])
.expect("OAuth reassignment fixture should encode");
let err = import_sqlite_jsonl(&pool, &encoded)
.await
.expect_err("taking the old owner's last OAuth binding should fail");
assert!(err
.to_string()
.contains("without an enabled identity binding"));
let owner: String =
sqlx::query_scalar("SELECT user_id FROM user_oauth_links WHERE id = 'reassigned-link'")
.fetch_one(&pool)
.await
.expect("rolled-back OAuth link should load");
assert_eq!(owner, "oauth-owner");
}
#[tokio::test]
async fn sqlite_oauth_link_import_ignores_unrelated_legacy_identity_damage() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::raw_sql(
r#"
INSERT INTO users (id, username, role, auth_source, created_at, updated_at) VALUES
('broken-oauth', 'broken-oauth', 'user', 'oauth', 1, 1),
('local-user', 'local-user', 'user', 'local', 1, 1);
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES (
'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback',
'https://app.example.test/auth/callback', 1, 1, 1
);
INSERT INTO user_oauth_links (
id, user_id, provider_type, provider_user_id, linked_at
) VALUES (
'legacy-orphan', 'missing-user', 'missing-provider', 'legacy-subject', 1
);
"#,
)
.execute(&pool)
.await
.expect("legacy damaged identity fixtures should seed");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::UserOAuthLinks],
)),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"valid-link",
json!({
"id": "valid-link",
"user_id": "local-user",
"provider_type": "linuxdo",
"provider_user_id": "valid-subject",
"linked_at": 1
}),
),
])
.expect("valid OAuth link fixture should encode");
assert_eq!(
import_sqlite_jsonl(&pool, &encoded)
.await
.expect("unrelated legacy damage must not block a valid scoped import"),
1
);
let valid_link_count: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links WHERE id = 'valid-link'")
.fetch_one(&pool)
.await
.expect("valid OAuth link count should load");
assert_eq!(valid_link_count, 1);
}
#[tokio::test]
async fn sqlite_oauth_link_import_rejects_duplicate_identity_in_legacy_schema() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::raw_sql(
r#"
DROP INDEX uq_user_oauth_links_provider_user;
INSERT INTO users (id, username, role, auth_source, created_at, updated_at) VALUES
('local-a', 'local-a', 'user', 'local', 1, 1),
('local-b', 'local-b', 'user', 'local', 1, 1);
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES (
'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback',
'https://app.example.test/auth/callback', 1, 1, 1
);
"#,
)
.execute(&pool)
.await
.expect("legacy schema fixture should seed");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::UserOAuthLinks],
)),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"link-a",
json!({
"id": "link-a",
"user_id": "local-a",
"provider_type": "linuxdo",
"provider_user_id": "same-subject",
"linked_at": 1
}),
),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"link-b",
json!({
"id": "link-b",
"user_id": "local-b",
"provider_type": "linuxdo",
"provider_user_id": "same-subject",
"linked_at": 1
}),
),
])
.expect("duplicate identity fixture should encode");
let err = import_sqlite_jsonl(&pool, &encoded)
.await
.expect_err("duplicate provider identity should fail");
assert!(err.to_string().contains("more than once"));
let link_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links")
.fetch_one(&pool)
.await
.expect("rolled-back OAuth link count should load");
assert_eq!(link_count, 0);
}
#[tokio::test]
async fn sqlite_oauth_link_import_rejects_duplicate_user_provider_in_legacy_schema() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::raw_sql(
r#"
DROP INDEX uq_user_oauth_links_user_provider;
INSERT INTO users (id, username, role, auth_source, created_at, updated_at)
VALUES ('local-user', 'local-user', 'user', 'local', 1, 1);
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES (
'linuxdo', 'Linux.do', 'client', 'https://gateway.example.test/oauth/callback',
'https://app.example.test/auth/callback', 1, 1, 1
);
"#,
)
.execute(&pool)
.await
.expect("legacy schema fixture should seed");
let encoded = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_700_000_000,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::UserOAuthLinks],
)),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"link-a",
json!({
"id": "link-a",
"user_id": "local-user",
"provider_type": "linuxdo",
"provider_user_id": "subject-a",
"linked_at": 1
}),
),
DataExportRecord::row(
ExportDomain::UserOAuthLinks,
"link-b",
json!({
"id": "link-b",
"user_id": "local-user",
"provider_type": "linuxdo",
"provider_user_id": "subject-b",
"linked_at": 1
}),
),
])
.expect("duplicate user-provider fixture should encode");
let err = import_sqlite_jsonl(&pool, &encoded)
.await
.expect_err("duplicate user provider should fail");
assert!(err.to_string().contains("more than once"));
let link_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM user_oauth_links")
.fetch_one(&pool)
.await
.expect("rolled-back OAuth link count should load");
assert_eq!(link_count, 0);
}
#[tokio::test]
async fn sqlite_core_export_reads_migrated_database_rows() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
@@ -694,6 +1481,24 @@ VALUES (
'request-1', 'candidate-1', 2, 'provider-1',
'endpoint-1', 'provider-key-1', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z'
);
INSERT INTO usage_cost_reservations (
request_id, subject_id, reservation_token, admitted_at,
reserved_cost_units, state, reservation_expires_at, retain_until,
created_at, updated_at
)
VALUES (
'request-1', 'user-1', 'reservation-1', 1,
500, 'reserved', 2, 3,
1, 1
);
INSERT INTO usage_request_admissions (
request_id, subject_id, event_token, admitted_at,
retain_until, state, created_at
)
VALUES (
'request-1', 'user-1', 'admission-1', 1,
3, 'active', 1
);
"#,
)
.execute(&pool)
@@ -757,6 +1562,19 @@ VALUES (
.any(|row| row.payload["__table"] == "usage_routing_snapshots"
&& row.payload["candidate_id"] == "candidate-1"
&& row.payload["selected_provider_id"] == "provider-1"));
assert!(import_plan
.rows(ExportDomain::Auxiliary)
.iter()
.any(|row| row.payload["__table"] == "usage_cost_reservations"
&& row.payload["reservation_token"] == "reservation-1"
&& row.payload["reserved_cost_units"] == 500
&& row.payload["state"] == "reserved"));
assert!(import_plan
.rows(ExportDomain::Auxiliary)
.iter()
.any(|row| row.payload["__table"] == "usage_request_admissions"
&& row.payload["event_token"] == "admission-1"
&& row.payload["state"] == "active"));
let target_pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
@@ -769,14 +1587,19 @@ VALUES (
let imported = import_sqlite_jsonl(&target_pool, &encoded)
.await
.expect("sqlite import should load exported rows");
assert_eq!(imported, 20);
assert_eq!(imported, 22);
let imported_api_key =
sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'")
let imported_api_key = sqlx::query_as::<_, (String, Option<String>, bool, bool, String)>(
"SELECT key_hash, key_encrypted, is_active, is_locked, status FROM api_keys WHERE id = 'api-key-1'",
)
.fetch_one(&target_pool)
.await
.expect("imported api key should load");
assert_eq!(imported_api_key.0, "ciphertext-1");
assert_ne!(imported_api_key.0, "hash-1");
assert_eq!(imported_api_key.1, None);
assert!(!imported_api_key.2);
assert!(imported_api_key.3);
assert_eq!(imported_api_key.4, "disabled");
let imported_usage = sqlx::query_as::<_, (String, i64, String)>(
"SELECT request_id, created_at_unix_ms, typeof(created_at_unix_ms) FROM \"usage\" WHERE request_id = 'request-1'",
@@ -844,6 +1667,36 @@ WHERE request_id = 'request-1'
("candidate-1".to_string(), 2, "provider-1".to_string())
);
let imported_reservation = sqlx::query_as::<_, (String, i64, String)>(
r#"
SELECT subject_id, reserved_cost_units, state
FROM usage_cost_reservations
WHERE reservation_token = 'reservation-1'
"#,
)
.fetch_one(&target_pool)
.await
.expect("imported usage cost reservation should load");
assert_eq!(
imported_reservation,
("user-1".to_string(), 500, "reserved".to_string())
);
let imported_admission = sqlx::query_as::<_, (String, String, Option<i64>)>(
r#"
SELECT subject_id, state, released_at
FROM usage_request_admissions
WHERE event_token = 'admission-1'
"#,
)
.fetch_one(&target_pool)
.await
.expect("imported usage request admission should load");
assert_eq!(
imported_admission,
("user-1".to_string(), "active".to_string(), None)
);
if let Some(database_url) = std::env::var("AETHER_TEST_POSTGRES_URL")
.ok()
.filter(|value| !value.trim().is_empty())
@@ -869,15 +1722,19 @@ WHERE request_id = 'request-1'
let imported = import_postgres_jsonl(&postgres_pool, &encoded)
.await
.expect("postgres import should load exported rows");
assert_eq!(imported, 20);
assert_eq!(imported, 22);
let imported_api_key = sqlx::query_as::<_, (String,)>(
"SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'",
let imported_api_key = sqlx::query_as::<_, (String, Option<String>, bool, bool, String)>(
"SELECT key_hash, key_encrypted, is_active, is_locked, status FROM api_keys WHERE id = 'api-key-1'",
)
.fetch_one(&postgres_pool)
.await
.expect("imported postgres api key should load");
assert_eq!(imported_api_key.0, "ciphertext-1");
assert_ne!(imported_api_key.0, "hash-1");
assert_eq!(imported_api_key.1, None);
assert!(!imported_api_key.2);
assert!(imported_api_key.3);
assert_eq!(imported_api_key.4, "disabled");
}
}
@@ -1317,13 +2174,18 @@ async fn mysql_core_export_reads_migrated_database_rows_when_url_is_set() {
.expect("mysql import should be idempotent");
assert!(imported >= 6);
let imported_api_key =
sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = ?")
.bind(&api_key_id)
.fetch_one(&pool)
.await
.expect("imported mysql api key should load");
assert_eq!(imported_api_key.0, "ciphertext-1");
let imported_api_key = sqlx::query_as::<_, (String, Option<String>, bool, bool, String)>(
"SELECT key_hash, key_encrypted, is_active, is_locked, status FROM api_keys WHERE id = ?",
)
.bind(&api_key_id)
.fetch_one(&pool)
.await
.expect("imported mysql api key should load");
assert_ne!(imported_api_key.0, "hash-1");
assert_eq!(imported_api_key.1, None);
assert!(!imported_api_key.2);
assert!(imported_api_key.3);
assert_eq!(imported_api_key.4, "disabled");
}
fn unique_suffix() -> String {
@@ -7,7 +7,7 @@ use std::time::{Duration, Instant};
use sqlx::{
migrate::{AppliedMigration, Migrate},
query, query_scalar, Connection, PgConnection, PgPool, SqlitePool,
query, query_scalar, Connection, PgConnection, PgPool, Row, SqlitePool,
};
use aether_data_contracts::repository::{
@@ -23,7 +23,8 @@ use super::{
prepare_database_for_startup,
};
use crate::lifecycle::bootstrap::postgres::{
snapshot_migrations as empty_database_snapshot_migrations, EMPTY_DATABASE_SNAPSHOT_SQL,
snapshot_migrations as empty_database_snapshot_migrations,
EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION, EMPTY_DATABASE_SNAPSHOT_SQL,
};
#[derive(Debug)]
@@ -411,8 +412,12 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
20260720000000,
20260727000000,
20260731000000,
20260814000000,
20260815000000,
20260816000000,
20260821000000,
20260903000000,
20260821120000,
20260821130000,
]
);
}
@@ -490,6 +495,9 @@ fn portable_driver_migrations_create_the_postgres_table_set() {
.iter()
.filter(|migration| migration.migration_type.is_up_migration())
.flat_map(|migration| create_table_names(migration.sql.as_ref()))
// SQLite rebuilds tables to add foreign keys. These staging tables are
// renamed to the canonical table names before the migration finishes.
.filter(|table| !table.ends_with("_with_user_fk"))
.collect::<BTreeSet<_>>();
assert_eq!(mysql_tables, postgres_tables, "MySQL table set drifted");
@@ -1023,6 +1031,487 @@ fn worker_boot_cleanup_migration_is_enabled_for_every_driver() {
}
}
#[test]
fn request_candidate_sensitive_diagnostic_purge_is_enabled_for_every_driver() {
const VERSION: i64 = 20260822000000;
for (driver, migrator) in [
("postgres", &POSTGRES_MIGRATOR),
("mysql", &super::mysql::MIGRATOR),
("sqlite", &super::sqlite::MIGRATOR),
] {
let migration = migrator
.iter()
.find(|migration| migration.version == VERSION)
.unwrap_or_else(|| {
panic!("{driver} request candidate diagnostic purge should be embedded")
});
let sql = migration.sql.as_ref();
for required in [
"UPDATE request_candidates",
"username = NULL",
"api_key_name = NULL",
"extra_data = NULL",
"required_capabilities = NULL",
"error_message = NULL",
"error_type = NULL",
"skip_reason = NULL",
] {
assert!(
sql.contains(required),
"{driver} request candidate diagnostic purge is missing {required}"
);
}
}
}
#[test]
fn deleted_user_history_anonymization_is_enabled_for_every_driver() {
const VERSION: i64 = 20260827050000;
const HISTORY_TABLES: &[&str] = &[
"request_candidates",
"video_tasks",
"usage",
"stats_user_daily",
"stats_user_summary",
"stats_user_daily_model",
"stats_user_daily_provider",
"stats_user_daily_api_format",
"stats_user_daily_model_provider",
"stats_user_daily_cost_savings",
"stats_user_daily_cost_savings_provider",
"stats_user_daily_cost_savings_model",
"stats_user_daily_cost_savings_model_provider",
];
for (driver, migrator) in [
("postgres", &POSTGRES_MIGRATOR),
("mysql", &super::mysql::MIGRATOR),
("sqlite", &super::sqlite::MIGRATOR),
] {
let migration = migrator
.iter()
.find(|migration| migration.version == VERSION)
.unwrap_or_else(|| panic!("{driver} user-history anonymization should be embedded"));
let sql = migration.sql.as_ref();
for table in HISTORY_TABLES {
assert!(
sql.contains(&format!("UPDATE {table}"))
|| sql.contains(&format!("UPDATE `{table}`"))
|| sql.contains(&format!("UPDATE public.{table}")),
"{driver} user-history anonymization is missing {table}"
);
}
assert_eq!(
sql.matches("SET username = NULL").count(),
HISTORY_TABLES.len(),
"{driver} must erase every username snapshot"
);
assert!(
sql.matches("api_key_name = NULL").count() >= 4,
"{driver} must erase API key names from request, video, usage, and API key aggregates"
);
assert!(
sql.matches("WHERE NOT EXISTS").count() >= HISTORY_TABLES.len() + 1,
"{driver} migration must leave existing users untouched"
);
assert!(
sql.contains("UPDATE stats_daily_api_key")
|| sql.contains("UPDATE public.stats_daily_api_key"),
"{driver} migration must anonymize orphaned API key aggregate names"
);
for fact_table in [
"user_plan_entitlements",
"wallets",
"user_referrals",
"referral_rewards",
] {
assert!(
sql.contains(&format!("UPDATE {fact_table}"))
|| sql.contains(&format!("UPDATE public.{fact_table}")),
"{driver} deleted-user fact migration is missing {fact_table}"
);
}
for required in [
"THEN 'revoked'",
"status = 'disabled'",
"THEN 'voided'",
"source_json = NULL",
"failure_reason = NULL",
"admin_note = NULL",
"payment_callbacks",
"payload = NULL",
"error_message = NULL",
] {
assert!(
sql.contains(required),
"{driver} deleted-user fact migration is missing {required}"
);
}
}
let postgres_migration = POSTGRES_MIGRATOR
.iter()
.find(|migration| migration.version == VERSION)
.expect("postgres user-history anonymization should be embedded");
for constraint in [
"request_candidates_user_id_fkey",
"video_tasks_user_id_fkey",
"usage_user_id_fkey",
"stats_user_daily_user_id_fkey",
"stats_user_summary_user_id_fkey",
"stats_user_daily_model_user_id_fkey",
"stats_user_daily_provider_user_id_fkey",
"stats_user_daily_api_format_user_id_fkey",
"stats_user_daily_model_provider_user_id_fkey",
"stats_user_daily_cost_savings_user_id_fkey",
"stats_user_daily_cost_savings_provider_user_id_fkey",
"stats_user_daily_cost_savings_model_user_id_fkey",
"stats_user_daily_cost_savings_model_provider_user_id_fkey",
"stats_hourly_user_model_user_id_fkey",
"user_model_usage_counts_user_id_fkey",
"request_candidates_api_key_id_fkey",
"video_tasks_api_key_id_fkey",
"usage_api_key_id_fkey",
"stats_daily_api_key_api_key_id_fkey",
"audit_logs_user_id_fkey",
"payment_orders_user_id_fkey",
"refund_requests_user_id_fkey",
"wallet_transactions_operator_id_fkey",
"wallets_user_id_fkey",
"wallets_api_key_id_fkey",
"user_plan_entitlements_user_id_fkey",
"entitlement_usage_ledgers_user_id_fkey",
"user_referrals_inviter_user_id_fkey",
"user_referrals_invitee_user_id_fkey",
"referral_rewards_inviter_user_id_fkey",
"referral_rewards_invitee_user_id_fkey",
] {
assert!(
postgres_migration
.sql
.contains(&format!("DROP CONSTRAINT IF EXISTS {constraint}")),
"postgres migration must decouple {constraint}"
);
}
}
#[test]
fn deleted_user_history_anonymization_remediation_is_enabled_for_every_driver() {
const VERSION: i64 = 20260829000000;
const HISTORY_TABLES: &[&str] = &[
"request_candidates",
"video_tasks",
"usage",
"stats_user_daily",
"stats_user_summary",
"stats_user_daily_model",
"stats_user_daily_provider",
"stats_user_daily_api_format",
"stats_user_daily_model_provider",
"stats_user_daily_cost_savings",
"stats_user_daily_cost_savings_provider",
"stats_user_daily_cost_savings_model",
"stats_user_daily_cost_savings_model_provider",
];
const OWNER_HISTORY_TABLES: &[&str] = &[
"wallets",
"audit_logs",
"wallet_transactions",
"payment_orders",
"payment_callbacks",
"refund_requests",
"redeem_code_batches",
];
for (driver, migrator) in [
("postgres", &POSTGRES_MIGRATOR),
("mysql", &super::mysql::MIGRATOR),
("sqlite", &super::sqlite::MIGRATOR),
] {
let migration = migrator
.iter()
.find(|migration| migration.version == VERSION)
.unwrap_or_else(|| panic!("{driver} history remediation should be embedded"));
let sql = migration.sql.as_ref();
for table in HISTORY_TABLES {
assert!(
sql.contains(&format!("UPDATE {table}"))
|| sql.contains(&format!("UPDATE `{table}`"))
|| sql.contains(&format!("UPDATE public.{table}")),
"{driver} remediation is missing {table}"
);
}
for table in OWNER_HISTORY_TABLES {
assert!(
sql.contains(&format!("UPDATE {table}"))
|| sql.contains(&format!("UPDATE `{table}`"))
|| sql.contains(&format!("UPDATE public.{table}")),
"{driver} remediation is missing owner-linked history table {table}"
);
}
for table in ["request_candidates", "video_tasks", "usage"] {
let has_api_key_projection = sql.split(';').any(|statement| {
let statement = statement.replace('`', "");
(statement.contains(&format!("UPDATE {table}"))
|| statement.contains(&format!("UPDATE public.{table}")))
&& statement.contains("SET api_key_name = NULL")
&& statement.contains("api_key_id IS NOT NULL")
});
assert!(
has_api_key_projection,
"{driver} remediation must clear orphaned api_key_name values in {table}"
);
}
for field in ["requested_by", "approved_by", "processed_by"] {
assert!(
sql.contains(&format!("SET {field} = NULL"))
&& sql.contains(&format!("{field} IS NOT NULL")),
"{driver} remediation must clear only orphaned {field} references"
);
}
let normalized_sql = sql.replace("public.", "").replace('`', "");
assert!(
normalized_sql.split(';').any(|statement| {
statement.contains("UPDATE payment_callbacks")
&& statement.contains("SET payload = NULL")
&& statement.contains("WHERE payload IS NOT NULL")
}),
"{driver} remediation must purge every persisted callback payload"
);
assert!(
normalized_sql.contains("NOT EXISTS (\n SELECT 1 FROM wallets")
|| normalized_sql.contains("NOT EXISTS (\n SELECT 1 FROM wallets"),
"{driver} remediation must clear sensitive fields for missing wallets"
);
}
}
#[tokio::test]
async fn sqlite_history_remediation_fails_closed_for_orphaned_financial_records() {
const VERSION: i64 = 20260829000000;
let pool = SqlitePool::connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
let mut connection = pool.acquire().await.expect("sqlite connection should open");
connection
.ensure_migrations_table()
.await
.expect("migration table should be created");
for migration in super::sqlite::MIGRATOR
.iter()
.filter(|migration| migration.version < VERSION)
{
connection
.apply(migration)
.await
.expect("pre-remediation migration should apply");
}
drop(connection);
query(
r#"
INSERT INTO users (id, username, email, auth_source, created_at, updated_at)
VALUES ('remediation-live-user', 'remediation-live', 'remediation-live@example.com', 'local', 1, 1);
INSERT INTO wallets (
id, user_id, balance, gift_balance, limit_mode, currency, status,
total_recharged, total_consumed, total_refunded, total_adjusted,
created_at, updated_at
) VALUES (
'remediation-live-wallet', 'remediation-live-user', 0, 0, 'finite', 'USD', 'active',
0, 0, 0, 0, 1, 1
);
INSERT INTO payment_orders (
id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd,
refundable_amount_usd, payment_method, gateway_response, status, created_at
) VALUES
('remediation-orphan-order', 'remediation-orphan-order-no', 'missing-wallet', NULL,
10, 0, 0, 'epay', '{"secret":"orphan"}', 'pending', 1),
('remediation-live-order', 'remediation-live-order-no', 'remediation-live-wallet',
'remediation-live-user', 10, 0, 0, 'epay', '{"secret":"live"}', 'pending', 1);
INSERT INTO payment_callbacks (
id, payment_order_id, payment_method, callback_key, order_no,
payload_hash, signature_valid, status, payload, error_message, created_at
) VALUES
('remediation-unmatched-callback', NULL, 'epay', 'remediation-unmatched-key',
'remediation-unmatched-order-no', 'hash-unmatched', 0, 'failed',
'SECRET-UNMATCHED', 'private unmatched error', 1),
('remediation-orphan-callback', 'remediation-orphan-order', 'epay',
'remediation-orphan-key', 'remediation-orphan-order-no', 'hash-orphan', 0,
'failed', 'SECRET-ORPHAN', 'private orphan error', 1),
('remediation-live-callback', 'remediation-live-order', 'epay',
'remediation-live-key', 'remediation-live-order-no', 'hash-live', 0,
'failed', 'SECRET-LIVE', 'retain diagnostic', 1);
INSERT INTO wallet_transactions (
id, wallet_id, category, reason_code, amount, balance_before, balance_after,
recharge_balance_before, recharge_balance_after, gift_balance_before,
gift_balance_after, description, created_at
) VALUES (
'remediation-orphan-wallet-tx', 'missing-wallet', 'adjust', 'manual', 1,
0, 1, 0, 1, 0, 0, 'private orphan transaction note', 1
);
INSERT INTO refund_requests (
id, refund_no, wallet_id, user_id, source_type, refund_mode, amount_usd,
status, reason, payout_reference, payout_proof, failure_reason,
created_at, updated_at
) VALUES (
'remediation-orphan-refund', 'remediation-orphan-refund-no', 'missing-wallet',
NULL, 'payment_order', 'original', 1, 'pending_approval', 'private refund reason',
'private payout reference', 'private payout proof', 'private failure', 1, 1
);
"#,
)
.execute(&pool)
.await
.expect("orphan financial fixtures should insert");
let migration = super::sqlite::MIGRATOR
.iter()
.find(|migration| migration.version == VERSION)
.expect("history remediation migration should be embedded");
sqlx::raw_sql(migration.sql.as_ref())
.execute(&pool)
.await
.expect("history remediation should apply");
let payload_count: i64 =
query_scalar("SELECT COUNT(*) FROM payment_callbacks WHERE payload IS NOT NULL")
.fetch_one(&pool)
.await
.expect("callback payload count should query");
assert_eq!(
payload_count, 0,
"raw callback payloads must be purged globally"
);
let orphan_error_count: i64 = query_scalar(
"SELECT COUNT(*) FROM payment_callbacks WHERE id IN ('remediation-unmatched-callback', 'remediation-orphan-callback') AND error_message IS NOT NULL",
)
.fetch_one(&pool)
.await
.expect("orphan callback errors should query");
assert_eq!(
orphan_error_count, 0,
"orphan callback diagnostics must be purged"
);
let live_error: Option<String> = query_scalar(
"SELECT error_message FROM payment_callbacks WHERE id = 'remediation-live-callback'",
)
.fetch_one(&pool)
.await
.expect("live callback error should query");
assert_eq!(live_error.as_deref(), Some("retain diagnostic"));
let orphan_gateway_response: Option<String> = query_scalar(
"SELECT gateway_response FROM payment_orders WHERE id = 'remediation-orphan-order'",
)
.fetch_one(&pool)
.await
.expect("orphan gateway response should query");
assert_eq!(orphan_gateway_response, None);
let orphan_description: Option<String> = query_scalar(
"SELECT description FROM wallet_transactions WHERE id = 'remediation-orphan-wallet-tx'",
)
.fetch_one(&pool)
.await
.expect("orphan transaction description should query");
assert_eq!(orphan_description, None);
let orphan_refund = query(
"SELECT reason, payout_reference, payout_proof, failure_reason FROM refund_requests WHERE id = 'remediation-orphan-refund'",
)
.fetch_one(&pool)
.await
.expect("orphan refund should query");
for column in [
"reason",
"payout_reference",
"payout_proof",
"failure_reason",
] {
assert_eq!(
orphan_refund
.try_get::<Option<String>, _>(column)
.expect("orphan refund column should decode"),
None,
"orphan refund {column} must be anonymized"
);
}
}
#[test]
fn background_task_sensitive_diagnostic_purge_is_enabled_for_every_driver() {
const VERSION: i64 = 20260822020000;
for (driver, migrator) in [
("postgres", &POSTGRES_MIGRATOR),
("mysql", &super::mysql::MIGRATOR),
("sqlite", &super::sqlite::MIGRATOR),
] {
let migration = migrator
.iter()
.find(|migration| migration.version == VERSION)
.unwrap_or_else(|| {
panic!("{driver} background task diagnostic purge should be embedded")
});
let sql = migration.sql.as_ref();
for required in [
"owner_instance = NULL",
"created_by = CASE",
"progress_message = NULL",
"payload_json = NULL",
"result_json = NULL",
"error_message = CASE",
"payload_json = NULL",
"unclassified_event",
] {
assert!(
sql.contains(required),
"{driver} background task diagnostic purge is missing {required}"
);
}
}
}
#[test]
fn identity_oauth_raw_userinfo_purge_is_enabled_for_every_driver() {
const VERSION: i64 = 20260827030000;
for (driver, migrator) in [
("postgres", &POSTGRES_MIGRATOR),
("mysql", &super::mysql::MIGRATOR),
("sqlite", &super::sqlite::MIGRATOR),
] {
let migration = migrator
.iter()
.find(|migration| migration.version == VERSION)
.unwrap_or_else(|| panic!("{driver} identity OAuth userinfo purge should be embedded"));
let sql = migration.sql.as_ref();
for required in [
"UPDATE",
"user_oauth_links",
"SET extra_data = NULL",
"WHERE extra_data IS NOT NULL",
] {
assert!(
sql.contains(required),
"{driver} identity OAuth userinfo purge is missing {required}"
);
}
}
}
#[test]
fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
let mysql_versions = super::mysql::MIGRATOR
@@ -1066,8 +1555,29 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260725030000,
20260727000000,
20260731000000,
20260814000000,
20260815000000,
20260816000000,
20260817000000,
20260821000000,
20260821120000,
20260821130000,
20260822000000,
20260822010000,
20260822020000,
20260827000000,
20260827010000,
20260827020000,
20260827030000,
20260827040000,
20260827050000,
20260829000000,
20260831000000,
20260831010000,
20260831020000,
20260831030000,
20260903000000,
20260903010000,
]
);
assert_eq!(
@@ -1102,12 +1612,224 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260725040000,
20260727000000,
20260731000000,
20260814000000,
20260815000000,
20260816000000,
20260821000000,
20260821120000,
20260821130000,
20260822000000,
20260822010000,
20260822020000,
20260827000000,
20260827010000,
20260827020000,
20260827030000,
20260827040000,
20260827050000,
20260829000000,
20260831000000,
20260831010000,
20260831020000,
20260831030000,
20260903000000,
20260903010000,
]
);
}
#[tokio::test]
async fn sqlite_legacy_oauth_email_verification_migration_fails_closed() {
const VERSION: i64 = 20260903010000;
let pool = SqlitePool::connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
super::run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
query(
r#"
INSERT INTO users (
id, email, username, auth_source, email_verified, created_at, updated_at
) VALUES
('legacy-oauth', 'oauth@example.com', 'legacy-oauth', 'oauth', 1, 1, 1),
('local-user', 'local@example.com', 'local-user', 'local', 1, 1, 1),
('ldap-user', 'ldap@example.com', 'ldap-user', 'ldap', 1, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("email verification fixtures should insert");
let migration = super::sqlite::MIGRATOR
.iter()
.find(|migration| migration.version == VERSION)
.expect("legacy OAuth verification migration should be embedded");
sqlx::raw_sql(migration.sql.as_ref())
.execute(&pool)
.await
.expect("legacy OAuth verification migration should apply");
for (user_id, expected) in [
("legacy-oauth", 0_i64),
("local-user", 1_i64),
("ldap-user", 1_i64),
] {
let verified: i64 = query_scalar("SELECT email_verified FROM users WHERE id = ?")
.bind(user_id)
.fetch_one(&pool)
.await
.expect("user verification state should load");
assert_eq!(verified, expected, "unexpected state for {user_id}");
}
}
#[tokio::test]
async fn sqlite_gateway_order_uniqueness_migration_rejects_historical_duplicates() {
const VERSION: i64 = 20260821120000;
let pool = SqlitePool::connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
let mut connection = pool.acquire().await.expect("sqlite connection should open");
connection
.ensure_migrations_table()
.await
.expect("migration table should be created");
for migration in super::sqlite::MIGRATOR
.iter()
.filter(|migration| migration.version < VERSION)
{
connection
.apply(migration)
.await
.expect("pre-uniqueness migration should apply");
}
drop(connection);
query(
r#"
INSERT INTO wallets (
id, user_id, balance, gift_balance, limit_mode, currency, status,
total_recharged, total_consumed, total_refunded, total_adjusted,
created_at, updated_at
) VALUES
('duplicate-wallet-a', 'duplicate-user-a', 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, 1, 1),
('duplicate-wallet-b', 'duplicate-user-b', 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, 1, 1);
INSERT INTO payment_orders (
id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd,
refundable_amount_usd, payment_method, gateway_order_id, status, created_at
) VALUES
('duplicate-order-a', 'duplicate-no-a', 'duplicate-wallet-a', 'duplicate-user-a', 1, 0, 0, ' EPAY ', 'duplicate-gateway-id', 'pending', 1),
('duplicate-order-b', 'duplicate-no-b', 'duplicate-wallet-b', 'duplicate-user-b', 1, 0, 0, 'epay', 'duplicate-gateway-id', 'pending', 1);
"#,
)
.execute(&pool)
.await
.expect("historical duplicate fixtures should insert");
let migration = super::sqlite::MIGRATOR
.iter()
.find(|migration| migration.version == VERSION)
.expect("gateway-order uniqueness migration should be embedded");
let error = sqlx::raw_sql(migration.sql.as_ref())
.execute(&pool)
.await
.expect_err("historical financial duplicates must block migration");
assert!(error.to_string().to_ascii_lowercase().contains("unique"));
let order_count: i64 = query_scalar(
"SELECT COUNT(*) FROM payment_orders WHERE gateway_order_id = 'duplicate-gateway-id'",
)
.fetch_one(&pool)
.await
.expect("duplicate financial records should remain intact");
assert_eq!(order_count, 2);
}
#[tokio::test]
async fn sqlite_gateway_order_uniqueness_migration_normalizes_legacy_payment_methods() {
const VERSION: i64 = 20260821120000;
let pool = SqlitePool::connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
let mut connection = pool.acquire().await.expect("sqlite connection should open");
connection
.ensure_migrations_table()
.await
.expect("migration table should be created");
for migration in super::sqlite::MIGRATOR
.iter()
.filter(|migration| migration.version < VERSION)
{
connection
.apply(migration)
.await
.expect("pre-uniqueness migration should apply");
}
drop(connection);
query(
r#"
INSERT INTO wallets (
id, user_id, balance, gift_balance, limit_mode, currency, status,
total_recharged, total_consumed, total_refunded, total_adjusted,
created_at, updated_at
) VALUES ('legacy-method-wallet', 'legacy-method-user', 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, 1, 1);
INSERT INTO payment_orders (
id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd,
refundable_amount_usd, payment_method, gateway_order_id, status, created_at
) VALUES ('legacy-method-order', 'legacy-method-no', 'legacy-method-wallet', 'legacy-method-user', 1, 0, 0, ' EPAY ', 'CaseSensitiveTxn', 'pending', 1);
INSERT INTO payment_callbacks (
id, payment_method, callback_key, signature_valid, status, created_at
) VALUES ('legacy-method-callback', ' EPAY ', 'legacy-method-key', 0, 'received', 1);
"#,
)
.execute(&pool)
.await
.expect("legacy payment methods should insert");
let migration = super::sqlite::MIGRATOR
.iter()
.find(|migration| migration.version == VERSION)
.expect("gateway-order uniqueness migration should be embedded");
sqlx::raw_sql(migration.sql.as_ref())
.execute(&pool)
.await
.expect("non-conflicting legacy payment methods should normalize");
let order_method: String =
query_scalar("SELECT payment_method FROM payment_orders WHERE id = 'legacy-method-order'")
.fetch_one(&pool)
.await
.expect("normalized order should load");
let callback_method: String = query_scalar(
"SELECT payment_method FROM payment_callbacks WHERE id = 'legacy-method-callback'",
)
.fetch_one(&pool)
.await
.expect("normalized callback should load");
assert_eq!(order_method, "epay");
assert_eq!(callback_method, "epay");
query(
r#"
INSERT INTO payment_orders (
id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd,
refundable_amount_usd, payment_method, gateway_order_id, status, created_at
) VALUES ('case-sensitive-order', 'case-sensitive-no', 'legacy-method-wallet', 'legacy-method-user', 1, 0, 0, 'epay', 'casesensitivetxn', 'pending', 1);
"#,
)
.execute(&pool)
.await
.expect("case-distinct opaque gateway identifiers should remain distinct");
}
#[tokio::test]
async fn sqlite_imported_timestamp_migration_normalizes_text_storage() {
let pool = SqlitePool::connect("sqlite::memory:")
@@ -2210,14 +2932,34 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
20260720000000,
20260727000000,
20260731000000,
20260814000000,
20260815000000,
20260816000000,
20260821000000,
20260821120000,
20260821130000,
20260822000000,
20260822010000,
20260822020000,
20260827000000,
20260827010000,
20260827020000,
20260827030000,
20260827040000,
20260827050000,
20260829000000,
20260831000000,
20260831010000,
20260831030000,
20260901000000,
20260903000000,
20260903010000,
]
);
}
#[test]
fn pending_migrations_from_applied_is_empty_after_empty_database_snapshot_stamp() {
fn pending_migrations_from_applied_only_returns_post_snapshot_migrations() {
let applied = empty_database_snapshot_migrations(&POSTGRES_MIGRATOR)
.expect("empty database snapshot migrations should resolve")
.into_iter()
@@ -2228,11 +2970,15 @@ fn pending_migrations_from_applied_is_empty_after_empty_database_snapshot_stamp(
.collect::<Vec<_>>();
let pending = pending_migrations_from_applied(&applied);
let expected = all_up_migrations()
.into_iter()
.filter(|migration| migration.version > EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION)
.collect::<Vec<_>>();
assert!(
pending.is_empty(),
"empty database snapshot-stamped databases should not require a manual migration before first startup"
);
assert_eq!(pending, expected);
assert!(pending
.iter()
.all(|migration| migration.version > EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION));
}
#[tokio::test]
File diff suppressed because it is too large Load Diff
@@ -4,8 +4,8 @@ pub use aether_data_contracts::repository::auth::{
read_resolved_auth_api_key_snapshot, read_resolved_auth_api_key_snapshot_by_key_hash,
read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyExportSummary,
AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository, AuthRepository,
CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, ResolvedAuthApiKeySnapshot,
ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery,
CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord,
ResolvedAuthApiKeySnapshot, ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery,
StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord,
UpdateUserApiKeyBasicRecord,
};
@@ -3,8 +3,8 @@ use std::sync::RwLock;
use async_trait::async_trait;
use super::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
};
use crate::DataLayerError;
@@ -49,15 +49,73 @@ impl AuthModuleReadRepository for InMemoryAuthModuleReadRepository {
#[async_trait]
impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository {
async fn upsert_ldap_config(
async fn compare_and_swap_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
self.ldap_config
expected: Option<&StoredLdapModuleConfig>,
replacement: &StoredLdapModuleConfig,
bind_password_update: &LdapBindPasswordUpdate,
) -> Result<CompareAndSwapLdapConfigResult, DataLayerError> {
let mut config = self
.ldap_config
.write()
.expect("auth module ldap repository lock")
.replace(config.clone());
Ok(Some(config.clone()))
.expect("auth module ldap repository lock");
if config.as_ref() != expected {
return Ok(CompareAndSwapLdapConfigResult::Conflict);
}
let bind_password_encrypted = match bind_password_update {
LdapBindPasswordUpdate::Preserve => expected
.ok_or_else(|| {
DataLayerError::InvalidConfiguration(
"LDAP bind password cannot be preserved while creating the singleton"
.to_string(),
)
})?
.bind_password_encrypted
.clone(),
LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()),
LdapBindPasswordUpdate::Clear => None,
};
let persisted = StoredLdapModuleConfig {
bind_password_encrypted,
..replacement.clone()
};
*config = Some(persisted.clone());
Ok(CompareAndSwapLdapConfigResult::Applied(persisted))
}
async fn delete_ldap_config_if_matches(
&self,
expected: &StoredLdapModuleConfig,
) -> Result<bool, DataLayerError> {
let mut config = self
.ldap_config
.write()
.expect("auth module ldap repository lock");
if config.as_ref() != Some(expected) {
return Ok(false);
}
config.take();
Ok(true)
}
async fn compare_and_swap_ldap_bind_password(
&self,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let mut config = self
.ldap_config
.write()
.expect("auth module ldap repository lock");
let Some(config) = config.as_mut() else {
return Ok(false);
};
if config.bind_password_encrypted.as_deref() != Some(expected) {
return Ok(false);
}
config.bind_password_encrypted = Some(replacement.to_string());
Ok(true)
}
}
@@ -65,9 +123,27 @@ impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository {
mod tests {
use super::InMemoryAuthModuleReadRepository;
use crate::repository::auth_modules::{
AuthModuleReadRepository, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
};
fn ldap_config() -> StoredLdapModuleConfig {
StoredLdapModuleConfig {
server_url: "ldaps://ldap.example.com".to_string(),
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
bind_password_encrypted: Some("encrypted-password".to_string()),
base_dn: "dc=example,dc=com".to_string(),
user_search_filter: Some("(uid={username})".to_string()),
username_attr: Some("uid".to_string()),
email_attr: Some("mail".to_string()),
display_name_attr: Some("displayName".to_string()),
is_enabled: true,
is_exclusive: false,
use_starttls: true,
connect_timeout: Some(10),
}
}
#[tokio::test]
async fn reads_seeded_auth_module_configs() {
let repository = InMemoryAuthModuleReadRepository::seed(
@@ -79,20 +155,7 @@ mod tests {
"https://example.com/callback".to_string(),
)
.expect("oauth provider should build")],
Some(StoredLdapModuleConfig {
server_url: "ldaps://ldap.example.com".to_string(),
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
bind_password_encrypted: Some("encrypted-password".to_string()),
base_dn: "dc=example,dc=com".to_string(),
user_search_filter: Some("(uid={username})".to_string()),
username_attr: Some("uid".to_string()),
email_attr: Some("mail".to_string()),
display_name_attr: Some("displayName".to_string()),
is_enabled: true,
is_exclusive: false,
use_starttls: true,
connect_timeout: Some(10),
}),
Some(ldap_config()),
);
let oauth = repository
@@ -111,4 +174,187 @@ mod tests {
"ldaps://ldap.example.com"
);
}
#[tokio::test]
async fn ldap_compensation_delete_requires_an_exact_match() {
let expected = ldap_config();
let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(expected.clone()));
let mismatched = StoredLdapModuleConfig {
is_enabled: false,
..expected.clone()
};
assert!(!repository
.delete_ldap_config_if_matches(&mismatched)
.await
.expect("mismatched delete should execute"));
assert!(repository
.delete_ldap_config_if_matches(&expected)
.await
.expect("matching delete should execute"));
assert!(repository
.get_ldap_config()
.await
.expect("LDAP config should remain readable")
.is_none());
}
#[tokio::test]
async fn ldap_compare_and_swap_separates_preserve_set_and_clear() {
let original = ldap_config();
let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(original.clone()));
let replacement = StoredLdapModuleConfig {
server_url: "ldap://updated.example.com".to_string(),
bind_password_encrypted: Some("stale-ciphertext-must-be-ignored".to_string()),
..original.clone()
};
let preserved = repository
.compare_and_swap_ldap_config(
Some(&original),
&replacement,
&LdapBindPasswordUpdate::Preserve,
)
.await
.expect("preserve CAS should execute");
let CompareAndSwapLdapConfigResult::Applied(preserved) = preserved else {
panic!("fresh snapshot should apply");
};
assert_eq!(
preserved.bind_password_encrypted.as_deref(),
Some("encrypted-password")
);
let set = repository
.compare_and_swap_ldap_config(
Some(&preserved),
&preserved,
&LdapBindPasswordUpdate::Set("rotated-ciphertext".to_string()),
)
.await
.expect("set CAS should execute");
let CompareAndSwapLdapConfigResult::Applied(set) = set else {
panic!("fresh snapshot should apply");
};
assert_eq!(
set.bind_password_encrypted.as_deref(),
Some("rotated-ciphertext")
);
let cleared = repository
.compare_and_swap_ldap_config(Some(&set), &set, &LdapBindPasswordUpdate::Clear)
.await
.expect("clear CAS should execute");
let CompareAndSwapLdapConfigResult::Applied(cleared) = cleared else {
panic!("fresh snapshot should apply");
};
assert!(cleared.bind_password_encrypted.is_none());
}
#[tokio::test]
async fn ldap_compare_and_swap_rejects_stale_password_and_config_snapshots() {
let original = ldap_config();
let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(original.clone()));
assert!(repository
.compare_and_swap_ldap_bind_password("encrypted-password", "rotated-ciphertext")
.await
.expect("password rotation should execute"));
let stale_password_result = repository
.compare_and_swap_ldap_config(
Some(&original),
&StoredLdapModuleConfig {
base_dn: "dc=updated,dc=example".to_string(),
..original.clone()
},
&LdapBindPasswordUpdate::Preserve,
)
.await
.expect("stale password CAS should execute");
assert_eq!(
stale_password_result,
CompareAndSwapLdapConfigResult::Conflict
);
assert_eq!(
repository
.get_ldap_config()
.await
.expect("LDAP config should load")
.and_then(|config| config.bind_password_encrypted)
.as_deref(),
Some("rotated-ciphertext")
);
let current = repository
.get_ldap_config()
.await
.expect("LDAP config should load")
.expect("LDAP config should exist");
let changed = StoredLdapModuleConfig {
is_enabled: false,
..current.clone()
};
let applied = repository
.compare_and_swap_ldap_config(
Some(&current),
&changed,
&LdapBindPasswordUpdate::Preserve,
)
.await
.expect("fresh config CAS should execute");
assert!(matches!(
applied,
CompareAndSwapLdapConfigResult::Applied(_)
));
let stale_config_result = repository
.compare_and_swap_ldap_config(
Some(&current),
&current,
&LdapBindPasswordUpdate::Preserve,
)
.await
.expect("stale config CAS should execute");
assert_eq!(
stale_config_result,
CompareAndSwapLdapConfigResult::Conflict
);
}
#[tokio::test]
async fn ldap_compare_and_swap_allows_only_one_initial_create() {
let repository = InMemoryAuthModuleReadRepository::default();
let replacement = StoredLdapModuleConfig {
bind_password_encrypted: None,
..ldap_config()
};
let first = repository
.compare_and_swap_ldap_config(
None,
&replacement,
&LdapBindPasswordUpdate::Set("first-ciphertext".to_string()),
)
.await
.expect("first create should execute");
assert!(matches!(first, CompareAndSwapLdapConfigResult::Applied(_)));
let second = repository
.compare_and_swap_ldap_config(
None,
&replacement,
&LdapBindPasswordUpdate::Set("second-ciphertext".to_string()),
)
.await
.expect("second create should execute");
assert_eq!(second, CompareAndSwapLdapConfigResult::Conflict);
assert_eq!(
repository
.get_ldap_config()
.await
.expect("LDAP config should load")
.and_then(|config| config.bind_password_encrypted)
.as_deref(),
Some("first-ciphertext")
);
}
}
@@ -1,8 +1,8 @@
mod memory;
pub use aether_data_contracts::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository};
@@ -53,7 +53,8 @@ impl InMemoryBackgroundTaskRepository {
I: IntoIterator<Item = StoredBackgroundTaskRun>,
{
let mut index = InMemoryBackgroundTaskIndex::default();
for run in runs {
for mut run in runs {
run.sanitize_persisted_data();
index.runs.insert(run.id.clone(), run);
}
Self {
@@ -164,8 +165,9 @@ impl BackgroundTaskReadRepository for InMemoryBackgroundTaskRepository {
impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
mut run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.sanitize_for_persistence();
run.validate()?;
let stored = run.into_stored();
self.index
@@ -192,8 +194,9 @@ impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository {
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
mut event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.sanitize_for_persistence();
event.validate()?;
let stored = event.into_stored();
let mut guard = self.index.write().expect("background task repository lock");
@@ -5,8 +5,9 @@ use async_trait::async_trait;
use super::{
AdminBillingMutationOutcome, BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, StoredBillingModelContext,
UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord,
};
use crate::DataLayerError;
@@ -215,6 +216,76 @@ impl BillingReadRepository for InMemoryBillingReadRepository {
.cloned())
}
async fn compare_and_swap_payment_gateway_secret(
&self,
update: &PaymentGatewaySecretCasUpdate,
) -> Result<bool, DataLayerError> {
let provider = update.provider.trim().to_ascii_lowercase();
let mut configs = self
.gateway_configs_by_provider
.write()
.expect("billing repository lock");
let Some(record) = configs.get_mut(&provider) else {
return Ok(false);
};
if record.merchant_key_encrypted.as_deref()
!= Some(update.expected_merchant_key_encrypted.as_str())
{
return Ok(false);
}
record.merchant_key_encrypted = Some(update.merchant_key_encrypted.clone());
Ok(true)
}
async fn compare_and_swap_payment_gateway_config(
&self,
mutation: &PaymentGatewayConfigCasWriteInput,
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, DataLayerError> {
let input = &mutation.input;
let provider = input.provider.trim().to_ascii_lowercase();
let now = current_unix_secs();
let mut configs = self
.gateway_configs_by_provider
.write()
.expect("billing repository lock");
let existing = configs.get(&provider);
if mutation.expected_existing {
let Some(existing) = existing else {
return Ok(AdminBillingMutationOutcome::NotFound);
};
if existing.merchant_key_encrypted != mutation.expected_merchant_key_encrypted {
return Ok(AdminBillingMutationOutcome::NotFound);
}
} else if existing.is_some() {
return Ok(AdminBillingMutationOutcome::NotFound);
}
let created_at = existing
.map(|value| value.created_at_unix_secs)
.unwrap_or(now);
let merchant_key_encrypted = if input.preserve_existing_secret {
existing.and_then(|value| value.merchant_key_encrypted.clone())
} else {
input.merchant_key_encrypted.clone()
};
let record = PaymentGatewayConfigRecord {
provider: provider.clone(),
enabled: input.enabled,
endpoint_url: input.endpoint_url.clone(),
callback_base_url: input.callback_base_url.clone(),
merchant_id: input.merchant_id.clone(),
merchant_key_encrypted,
pay_currency: input.pay_currency.clone(),
usd_exchange_rate: input.usd_exchange_rate,
min_recharge_usd: input.min_recharge_usd,
channels_json: input.channels_json.clone(),
created_at_unix_secs: created_at,
updated_at_unix_secs: now,
};
configs.insert(provider, record.clone());
Ok(AdminBillingMutationOutcome::Applied(record))
}
async fn upsert_payment_gateway_config(
&self,
input: &PaymentGatewayConfigWriteInput,
@@ -495,7 +566,10 @@ mod tests {
use serde_json::json;
use super::InMemoryBillingReadRepository;
use crate::repository::billing::{BillingReadRepository, StoredBillingModelContext};
use crate::repository::billing::{
AdminBillingMutationOutcome, BillingReadRepository, PaymentGatewayConfigCasWriteInput,
PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, StoredBillingModelContext,
};
fn sample_context() -> StoredBillingModelContext {
StoredBillingModelContext::new(
@@ -603,4 +677,105 @@ mod tests {
Some("gpt-5-upstream")
);
}
fn gateway_input(secret: Option<&str>) -> PaymentGatewayConfigWriteInput {
PaymentGatewayConfigWriteInput {
provider: "stripe".to_string(),
enabled: true,
endpoint_url: "https://api.stripe.com".to_string(),
callback_base_url: Some("https://example.com".to_string()),
merchant_id: "merchant".to_string(),
merchant_key_encrypted: secret.map(ToOwned::to_owned),
preserve_existing_secret: false,
pay_currency: "USD".to_string(),
usd_exchange_rate: 1.0,
min_recharge_usd: 1.0,
channels_json: json!({"channels": []}),
}
}
#[tokio::test]
async fn payment_gateway_cas_prevents_create_overwrite_and_uses_exact_secret_fence() {
let repository = InMemoryBillingReadRepository::default();
let create = PaymentGatewayConfigCasWriteInput {
input: gateway_input(Some("ciphertext-a")),
expected_existing: false,
expected_merchant_key_encrypted: None,
};
assert!(matches!(
repository
.compare_and_swap_payment_gateway_config(&create)
.await
.expect("create should succeed"),
AdminBillingMutationOutcome::Applied(_)
));
let mut competing_create = create.clone();
competing_create.input.merchant_id = "overwritten".to_string();
assert_eq!(
repository
.compare_and_swap_payment_gateway_config(&competing_create)
.await
.expect("conflicting create should be handled"),
AdminBillingMutationOutcome::NotFound
);
let mut stale_update = create.clone();
stale_update.expected_existing = true;
stale_update.expected_merchant_key_encrypted = Some("ciphertext-stale".to_string());
stale_update.input.merchant_id = "stale-update".to_string();
assert_eq!(
repository
.compare_and_swap_payment_gateway_config(&stale_update)
.await
.expect("stale update should be handled"),
AdminBillingMutationOutcome::NotFound
);
let stored = repository
.find_payment_gateway_config("stripe")
.await
.expect("lookup should succeed")
.expect("config should exist");
assert_eq!(stored.merchant_id, "merchant");
}
#[tokio::test]
async fn payment_gateway_secret_cas_changes_no_other_fields() {
let repository = InMemoryBillingReadRepository::default();
repository
.upsert_payment_gateway_config(&gateway_input(Some("legacy-ciphertext")))
.await
.expect("seed should succeed");
let before = repository
.find_payment_gateway_config("stripe")
.await
.expect("lookup should succeed")
.expect("config should exist");
assert!(!repository
.compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate {
provider: "stripe".to_string(),
expected_merchant_key_encrypted: "wrong-ciphertext".to_string(),
merchant_key_encrypted: "v2-ciphertext".to_string(),
})
.await
.expect("stale secret CAS should be handled"));
assert!(repository
.compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate {
provider: "stripe".to_string(),
expected_merchant_key_encrypted: "legacy-ciphertext".to_string(),
merchant_key_encrypted: "v2-ciphertext".to_string(),
})
.await
.expect("secret CAS should succeed"));
let mut expected = before.clone();
expected.merchant_key_encrypted = Some("v2-ciphertext".to_string());
let after = repository
.find_payment_gateway_config("stripe")
.await
.expect("lookup should succeed")
.expect("config should exist");
assert_eq!(after, expected);
}
}
@@ -1,14 +1,18 @@
use std::collections::{BTreeMap, BTreeSet};
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
request_candidate_lifecycle_would_regress, PublicHealthStatusCount, PublicHealthTimelineBucket,
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
StoredRequestCandidate, UpsertRequestCandidateRecord,
};
use crate::DataLayerError;
use async_trait::async_trait;
fn sanitize_stored_candidate(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
candidate.sanitize_sensitive_diagnostics();
candidate
}
fn merge_extra_data(
existing: Option<serde_json::Value>,
@@ -38,7 +42,7 @@ impl InMemoryRequestCandidateRepository {
I: IntoIterator<Item = StoredRequestCandidate>,
{
let mut by_id = BTreeMap::new();
for item in items {
for item in items.into_iter().map(sanitize_stored_candidate) {
by_id.insert(item.id.clone(), item);
}
Self {
@@ -60,6 +64,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
.values()
.filter(|row| row.request_id == request_id)
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.candidate_index
@@ -84,6 +89,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
.expect("request candidate repository lock")
.values()
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
@@ -106,6 +112,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
.values()
.filter(|row| row.provider_id.as_deref() == Some(provider_id))
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
@@ -141,6 +148,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
)
})
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
@@ -294,8 +302,9 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
mut candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.sanitize_for_persistence();
candidate.validate()?;
let mut by_id = self
@@ -309,7 +318,8 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
&& row.candidate_index == candidate.candidate_index
&& row.retry_index == candidate.retry_index
})
.cloned();
.cloned()
.map(sanitize_stored_candidate);
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|row| {
request_candidate_lifecycle_would_regress(row.status, candidate.status)
@@ -336,29 +346,36 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
.map(|row| row.id.clone())
.unwrap_or_else(|| candidate.id.clone()),
request_id: candidate.request_id.clone(),
user_id: candidate
.user_id
.or_else(|| existing.as_ref().and_then(|row| row.user_id.clone())),
api_key_id: candidate
.api_key_id
.or_else(|| existing.as_ref().and_then(|row| row.api_key_id.clone())),
username: candidate
.username
.or_else(|| existing.as_ref().and_then(|row| row.username.clone())),
api_key_name: candidate
.api_key_name
.or_else(|| existing.as_ref().and_then(|row| row.api_key_name.clone())),
user_id: existing
.as_ref()
.and_then(|row| row.user_id.clone())
.or(candidate.user_id),
api_key_id: existing
.as_ref()
.and_then(|row| row.api_key_id.clone())
.or(candidate.api_key_id),
username: existing
.as_ref()
.and_then(|row| row.username.clone())
.or(candidate.username),
api_key_name: existing
.as_ref()
.and_then(|row| row.api_key_name.clone())
.or(candidate.api_key_name),
candidate_index: candidate.candidate_index,
retry_index: candidate.retry_index,
provider_id: candidate
.provider_id
.or_else(|| existing.as_ref().and_then(|row| row.provider_id.clone())),
endpoint_id: candidate
.endpoint_id
.or_else(|| existing.as_ref().and_then(|row| row.endpoint_id.clone())),
key_id: candidate
.key_id
.or_else(|| existing.as_ref().and_then(|row| row.key_id.clone())),
provider_id: existing
.as_ref()
.and_then(|row| row.provider_id.clone())
.or(candidate.provider_id),
endpoint_id: existing
.as_ref()
.and_then(|row| row.endpoint_id.clone())
.or(candidate.endpoint_id),
key_id: existing
.as_ref()
.and_then(|row| row.key_id.clone())
.or(candidate.key_id),
status: merged_status,
skip_reason: candidate
.skip_reason
@@ -380,13 +397,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
.error_type
.or_else(|| existing.as_ref().and_then(|row| row.error_type.clone()))
},
error_message: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.error_message.clone())
} else {
candidate
.error_message
.or_else(|| existing.as_ref().and_then(|row| row.error_message.clone()))
},
error_message: None,
latency_ms: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.latency_ms)
} else {
@@ -407,9 +418,10 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
.and_then(|row| row.required_capabilities.clone())
}),
created_at_unix_ms,
started_at_unix_ms: candidate
.started_at_unix_ms
.or_else(|| existing.as_ref().and_then(|row| row.started_at_unix_ms)),
started_at_unix_ms: existing
.as_ref()
.and_then(|row| row.started_at_unix_ms)
.or(candidate.started_at_unix_ms),
finished_at_unix_ms: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.finished_at_unix_ms)
} else {
@@ -418,6 +430,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
.or_else(|| existing.as_ref().and_then(|row| row.finished_at_unix_ms))
},
};
let stored = sanitize_stored_candidate(stored);
by_id.insert(stored.id.clone(), stored.clone());
Ok(stored)
@@ -531,6 +544,138 @@ mod tests {
assert_eq!(rows[1].id, "cand-1");
}
#[tokio::test]
async fn seed_and_reads_sanitize_candidates_that_bypass_contract_constructors() {
let raw_candidate = StoredRequestCandidate {
id: "cand-raw".to_string(),
request_id: "req-raw".to_string(),
user_id: None,
api_key_id: None,
username: None,
api_key_name: None,
candidate_index: 0,
retry_index: 0,
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: None,
status: RequestCandidateStatus::Failed,
skip_reason: Some("secret=/private/path".to_string()),
is_cached: false,
status_code: Some(500),
error_type: Some("token=secret".to_string()),
error_message: Some("Bearer secret-token".to_string()),
latency_ms: Some(10),
concurrent_requests: Some(1),
extra_data: Some(json!({
"gateway_execution_runtime": true,
"request_headers": {"authorization": "Bearer secret-token"},
"request_body": {"password": "secret"}
})),
required_capabilities: Some(json!({
"streaming": "true",
"internal_capability": "secret"
})),
created_at_unix_ms: 100,
started_at_unix_ms: Some(100),
finished_at_unix_ms: Some(110),
};
let repository = InMemoryRequestCandidateRepository::seed(vec![raw_candidate.clone()]);
{
let stored = repository
.by_id
.read()
.expect("request candidate repository lock");
let candidate = stored
.get("cand-raw")
.expect("seeded candidate should exist");
assert!(candidate.error_message.is_none());
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
assert_eq!(
candidate.extra_data,
Some(json!({"gateway_execution_runtime": true}))
);
assert_eq!(
candidate.required_capabilities,
Some(json!({"streaming": true}))
);
}
let mut bypassed_candidate = raw_candidate;
bypassed_candidate.id = "cand-bypassed".to_string();
bypassed_candidate.request_id = "req-bypassed".to_string();
repository
.by_id
.write()
.expect("request candidate repository lock")
.insert(bypassed_candidate.id.clone(), bypassed_candidate);
let rows = repository
.list_recent(10)
.await
.expect("list recent should succeed");
let candidate = rows
.iter()
.find(|candidate| candidate.id == "cand-bypassed")
.expect("bypassed candidate should be returned");
assert!(candidate.error_message.is_none());
assert_eq!(
candidate.extra_data,
Some(json!({"gateway_execution_runtime": true}))
);
assert_eq!(
candidate.required_capabilities,
Some(json!({"streaming": true}))
);
let merged = repository
.upsert(UpsertRequestCandidateRecord {
id: "cand-merged".to_string(),
request_id: "req-bypassed".to_string(),
user_id: None,
api_key_id: None,
username: None,
api_key_name: None,
candidate_index: 0,
retry_index: 0,
provider_id: None,
endpoint_id: None,
key_id: None,
status: RequestCandidateStatus::Success,
skip_reason: None,
is_cached: None,
status_code: Some(200),
error_type: None,
error_message: Some("Bearer new-secret".to_string()),
latency_ms: Some(12),
concurrent_requests: None,
extra_data: Some(json!({
"stream_completed": true,
"request_body": {"password": "new-secret"}
})),
required_capabilities: Some(json!({
"vision": 1,
"internal_capability": "new-secret"
})),
created_at_unix_ms: Some(100),
started_at_unix_ms: Some(100),
finished_at_unix_ms: Some(112),
})
.await
.expect("candidate merge should succeed");
assert_eq!(merged.id, "cand-bypassed");
assert!(merged.error_message.is_none());
assert_eq!(
merged.extra_data,
Some(json!({
"gateway_execution_runtime": true,
"stream_completed": true
}))
);
assert_eq!(merged.required_capabilities, Some(json!({"vision": true})));
}
#[tokio::test]
async fn aggregates_finalized_health_data_by_endpoint_ids() {
let repository = InMemoryRequestCandidateRepository::seed(vec![
@@ -594,7 +739,7 @@ mod tests {
"execution_strategy": "local_cross_format",
"provider_name": "primary",
})),
required_capabilities: None,
required_capabilities: Some(json!({"streaming": true})),
created_at_unix_ms: Some(100),
started_at_unix_ms: None,
finished_at_unix_ms: None,
@@ -608,15 +753,15 @@ mod tests {
.upsert(UpsertRequestCandidateRecord {
id: "cand-1-replacement".to_string(),
request_id: "req-1".to_string(),
user_id: None,
api_key_id: None,
username: None,
api_key_name: None,
user_id: Some("attacker-user".to_string()),
api_key_id: Some("attacker-api-key".to_string()),
username: Some("mallory".to_string()),
api_key_name: Some("attacker-key".to_string()),
candidate_index: 0,
retry_index: 0,
provider_id: None,
endpoint_id: None,
key_id: None,
provider_id: Some("attacker-provider".to_string()),
endpoint_id: Some("attacker-endpoint".to_string()),
key_id: Some("attacker-provider-key".to_string()),
status: RequestCandidateStatus::Success,
skip_reason: None,
is_cached: None,
@@ -629,7 +774,7 @@ mod tests {
"provider_api_format": "openai:responses",
"provider_name": "updated",
})),
required_capabilities: None,
required_capabilities: Some(json!({"vision": true})),
created_at_unix_ms: None,
started_at_unix_ms: Some(101),
finished_at_unix_ms: Some(102),
@@ -638,6 +783,14 @@ mod tests {
.expect("update should succeed");
assert_eq!(updated.id, "cand-1");
assert_eq!(updated.status, RequestCandidateStatus::Success);
assert_eq!(updated.user_id.as_deref(), Some("user-1"));
assert_eq!(updated.api_key_id.as_deref(), Some("api-key-1"));
assert!(updated.username.is_none());
assert!(updated.api_key_name.is_none());
assert_eq!(updated.provider_id.as_deref(), Some("provider-1"));
assert_eq!(updated.endpoint_id.as_deref(), Some("endpoint-1"));
assert_eq!(updated.key_id.as_deref(), Some("key-1"));
assert_eq!(updated.required_capabilities, Some(json!({"vision": true})));
assert_eq!(updated.status_code, Some(200));
assert_eq!(updated.latency_ms, Some(25));
assert_eq!(
@@ -659,13 +812,13 @@ mod tests {
.extra_data
.as_ref()
.and_then(|value| value.get("provider_name")),
Some(&json!("updated"))
None
);
assert_eq!(updated.started_at_unix_ms, Some(101));
}
#[tokio::test]
async fn upsert_keeps_terminal_candidate_state_when_streaming_arrives_late() {
async fn upsert_keeps_first_terminal_candidate_fact_when_another_terminal_arrives_late() {
let existing = StoredRequestCandidate::new(
"cand-1".to_string(),
"req-1".to_string(),
@@ -686,7 +839,7 @@ mod tests {
Some("retryable upstream failure".to_string()),
Some(45),
Some(1),
Some(json!({"terminal": true})),
Some(json!({"stream_completed": true})),
None,
100,
Some(101),
@@ -708,7 +861,7 @@ mod tests {
provider_id: None,
endpoint_id: None,
key_id: None,
status: RequestCandidateStatus::Streaming,
status: RequestCandidateStatus::Success,
skip_reason: None,
is_cached: None,
status_code: Some(200),
@@ -716,7 +869,7 @@ mod tests {
error_message: None,
latency_ms: Some(9_999),
concurrent_requests: Some(2),
extra_data: Some(json!({"late": true})),
extra_data: Some(json!({"gateway_execution_runtime": true})),
required_capabilities: None,
created_at_unix_ms: None,
started_at_unix_ms: Some(102),
@@ -729,16 +882,16 @@ mod tests {
assert_eq!(updated.status, RequestCandidateStatus::Failed);
assert_eq!(updated.status_code, Some(503));
assert_eq!(updated.error_type.as_deref(), Some("upstream_error"));
assert_eq!(
updated.error_message.as_deref(),
Some("retryable upstream failure")
);
assert!(updated.error_message.is_none());
assert_eq!(updated.latency_ms, Some(45));
assert_eq!(updated.concurrent_requests, Some(2));
assert_eq!(updated.finished_at_unix_ms, Some(145));
assert_eq!(
updated.extra_data,
Some(json!({"terminal": true, "late": true}))
Some(json!({
"gateway_execution_runtime": true,
"stream_completed": true
}))
);
}
@@ -41,6 +41,40 @@ impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository {
Ok(guard.get(file_name).cloned())
}
async fn find_active_by_file_name_for_user(
&self,
file_name: &str,
user_id: &str,
now_unix_secs: u64,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let guard = self.by_file.read().expect("gemini mapping repository lock");
Ok(guard
.get(file_name)
.filter(|mapping| {
mapping.user_id.as_deref() == Some(user_id)
&& mapping.expires_at_unix_secs > now_unix_secs
})
.cloned())
}
async fn find_active_by_file_name_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
now_unix_secs: u64,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let guard = self.by_file.read().expect("gemini mapping repository lock");
Ok(guard
.get(file_name)
.filter(|mapping| {
mapping.key_id == key_id
&& mapping.user_id.as_deref() == Some(user_id)
&& mapping.expires_at_unix_secs > now_unix_secs
})
.cloned())
}
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
@@ -52,6 +86,12 @@ impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository {
.map(|value| value.to_ascii_lowercase());
let mut items = guard
.values()
.filter(|item| {
query
.user_id
.as_deref()
.is_none_or(|user_id| item.user_id.as_deref() == Some(user_id))
})
.filter(|item| query.include_expired || item.expires_at_unix_secs > query.now_unix_secs)
.filter(|item| {
search.as_deref().is_none_or(|needle| {
@@ -146,6 +186,39 @@ impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository {
Ok(mapping)
}
async fn upsert_if_owner_matches(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
record.validate()?;
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
let (id, created_at_unix_ms) = match guard.get(&record.file_name) {
Some(existing)
if existing.key_id == record.key_id && existing.user_id == record.user_id =>
{
(existing.id.clone(), existing.created_at_unix_ms)
}
Some(_) => return Ok(None),
None => (record.id.clone(), current_unix_secs()),
};
let mapping = StoredGeminiFileMapping {
id,
file_name: record.file_name.clone(),
key_id: record.key_id.clone(),
user_id: record.user_id.clone(),
display_name: record.display_name.clone(),
mime_type: record.mime_type.clone(),
source_hash: record.source_hash.clone(),
created_at_unix_ms,
expires_at_unix_secs: record.expires_at_unix_secs,
};
guard.insert(record.file_name, mapping.clone());
Ok(Some(mapping))
}
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
let mut guard = self
.by_file
@@ -154,6 +227,44 @@ impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository {
Ok(guard.remove(file_name).is_some())
}
async fn delete_by_file_name_for_user(
&self,
file_name: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
if guard
.get(file_name)
.and_then(|item| item.user_id.as_deref())
!= Some(user_id)
{
return Ok(false);
}
Ok(guard.remove(file_name).is_some())
}
async fn delete_by_file_name_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
let owner_matches = guard
.get(file_name)
.is_some_and(|item| item.key_id == key_id && item.user_id.as_deref() == Some(user_id));
if !owner_matches {
return Ok(false);
}
Ok(guard.remove(file_name).is_some())
}
async fn delete_by_id(
&self,
mapping_id: &str,
@@ -224,6 +335,35 @@ mod tests {
Ok(())
}
#[tokio::test]
async fn owner_scoped_reads_bind_user_provider_key_and_expiry() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::default();
repo.upsert(sample_record("id-owner", "files/owned"))
.await?;
assert!(repo
.find_active_by_file_name_for_user("files/owned", "user-1", 100)
.await?
.is_some());
assert!(repo
.find_active_by_file_name_for_user("files/owned", "user-2", 100)
.await?
.is_none());
assert!(repo
.find_active_by_file_name_for_owner("files/owned", "key-1", "user-1", 100)
.await?
.is_some());
assert!(repo
.find_active_by_file_name_for_owner("files/owned", "key-2", "user-1", 100)
.await?
.is_none());
assert!(repo
.find_active_by_file_name_for_user("files/owned", "user-1", 4_102_444_800)
.await?
.is_none());
Ok(())
}
#[tokio::test]
async fn delete_removes_entry() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::default();
@@ -247,6 +387,34 @@ mod tests {
Ok(())
}
#[tokio::test]
async fn owner_checked_upsert_cannot_reassign_existing_mapping() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::default();
let first = repo.upsert(sample_record("id-1", "files/owned")).await?;
let mut attacker = sample_record("id-2", "files/owned");
attacker.key_id = "key-2".to_string();
attacker.user_id = Some("user-2".to_string());
assert!(repo.upsert_if_owner_matches(attacker).await?.is_none());
let unchanged = repo
.find_by_file_name("files/owned")
.await?
.expect("mapping should remain");
assert_eq!(unchanged.key_id, "key-1");
assert_eq!(unchanged.user_id.as_deref(), Some("user-1"));
let mut refresh = sample_record("id-3", "files/owned");
refresh.display_name = Some("refreshed".to_string());
let refreshed = repo
.upsert_if_owner_matches(refresh)
.await?
.expect("same owner should refresh");
assert_eq!(refreshed.id, first.id);
assert_eq!(refreshed.display_name.as_deref(), Some("refreshed"));
Ok(())
}
#[tokio::test]
async fn list_and_summarize_mappings() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::seed(vec![
@@ -257,6 +425,7 @@ mod tests {
let page = repo
.list_mappings(&GeminiFileMappingListQuery {
user_id: None,
include_expired: false,
search: Some("ga".to_string()),
offset: 0,
@@ -279,6 +448,46 @@ mod tests {
Ok(())
}
#[tokio::test]
async fn owner_filter_and_delete_do_not_cross_user_boundaries() -> Result<(), DataLayerError> {
let mut first = repo_item("id-1", "files/alpha", "image/png", 10, 200);
first.user_id = Some("user-1".to_string());
let mut second = repo_item("id-2", "files/beta", "image/png", 20, 200);
second.user_id = Some("user-2".to_string());
let repo = InMemoryGeminiFileMappingRepository::seed([first, second]);
let page = repo
.list_mappings(&GeminiFileMappingListQuery {
user_id: Some("user-1".to_string()),
include_expired: false,
search: None,
offset: 0,
limit: 10,
now_unix_secs: 100,
})
.await?;
assert_eq!(page.total, 1);
assert_eq!(page.items[0].file_name, "files/alpha");
assert!(
!repo
.delete_by_file_name_for_user("files/alpha", "user-2")
.await?
);
assert!(repo.find_by_file_name("files/alpha").await?.is_some());
assert!(
!repo
.delete_by_file_name_for_owner("files/alpha", "key-2", "user-1")
.await?
);
assert!(repo.find_by_file_name("files/alpha").await?.is_some());
assert!(
repo.delete_by_file_name_for_user("files/alpha", "user-1")
.await?
);
Ok(())
}
#[tokio::test]
async fn delete_by_id_and_cleanup_expired() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::seed(vec![
@@ -6,9 +6,10 @@ use async_trait::async_trait;
use crate::DataLayerError;
use aether_data_contracts::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenWithUser, UpdateManagementTokenRecord,
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
};
#[derive(Debug, Default)]
@@ -49,6 +50,189 @@ impl InMemoryManagementTokenRepository {
fn remove_hash_for_token(hashes: &mut BTreeMap<String, String>, token_id: &str) {
hashes.retain(|_, existing_token_id| existing_token_id != token_id);
}
fn update_management_token_scoped(
&self,
record: &UpdateManagementTokenRecord,
expected_user_id: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
let mut items = self
.items
.write()
.expect("management token repository lock");
let Some(index) = items.iter().position(|item| {
item.token.id == record.token_id
&& expected_user_id
.map(|user_id| item.token.user_id == user_id)
.unwrap_or(true)
}) else {
return Ok(None);
};
if let Some(name) = &record.name {
if items.iter().enumerate().any(|(position, item)| {
position != index
&& item.token.user_id == items[index].token.user_id
&& item.token.name == *name
}) {
return Err(DataLayerError::InvalidInput(format!(
"已存在名为 '{}' 的 Token",
name
)));
}
items[index].token.name = name.clone();
}
if record.clear_description {
items[index].token.description = None;
} else if let Some(description) = &record.description {
items[index].token.description = Some(description.clone());
}
if record.clear_allowed_ips {
items[index].token.allowed_ips = None;
} else if let Some(allowed_ips) = &record.allowed_ips {
items[index].token.allowed_ips = Some(allowed_ips.clone());
}
if let Some(permissions) = &record.permissions {
items[index].token.permissions = Some(permissions.clone());
}
if record.clear_expires_at {
items[index].token.expires_at_unix_secs = None;
} else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs {
items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs);
}
if let Some(is_active) = record.is_active {
items[index].token.is_active = is_active;
}
items[index].token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(items[index].token.clone()))
}
fn delete_management_token_scoped(
&self,
token_id: &str,
expected_user_id: Option<&str>,
) -> bool {
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
let original_len = items.len();
items.retain(|item| {
item.token.id != token_id
|| expected_user_id
.map(|user_id| item.token.user_id != user_id)
.unwrap_or(false)
});
if items.len() != original_len {
Self::remove_hash_for_token(&mut hashes, token_id);
return true;
}
false
}
fn set_management_token_active_scoped(
&self,
token_id: &str,
expected_user_id: Option<&str>,
is_active: bool,
) -> Option<StoredManagementToken> {
let mut items = self
.items
.write()
.expect("management token repository lock");
let item = items.iter_mut().find(|item| {
item.token.id == token_id
&& expected_user_id
.map(|user_id| item.token.user_id == user_id)
.unwrap_or(true)
})?;
item.token.is_active = is_active;
item.token.updated_at_unix_secs = Self::now_unix_secs();
Some(item.token.clone())
}
fn activate_management_token_if_matches_inner(
&self,
mutation: &ActivateManagementTokenIfMatches,
) -> Result<bool, DataLayerError> {
mutation.validate()?;
// The in-memory token store does not own the independently stored user row and cannot
// atomically verify role/status/security_version with this mutation. Pretending that the
// user summary cached beside the token is authoritative would recreate the TOCTOU, so
// one-time install activation is intentionally unavailable on this backend.
Ok(false)
}
fn delete_inactive_management_token_if_matches_inner(
&self,
mutation: &ActivateManagementTokenIfMatches,
) -> Result<bool, DataLayerError> {
mutation.validate()?;
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
if hashes.get(&mutation.token_hash).map(String::as_str)
!= Some(mutation.expected_token.id.as_str())
{
return Ok(false);
}
let Some(index) = items.iter().position(|item| {
mutation.matches_locked_token_snapshot(&item.token, &mutation.token_hash)
}) else {
return Ok(false);
};
items.remove(index);
Self::remove_hash_for_token(&mut hashes, &mutation.expected_token.id);
Ok(true)
}
fn regenerate_management_token_secret_scoped(
&self,
mutation: &RegenerateManagementTokenSecret,
expected_user_id: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
let Some(item) = items.iter_mut().find(|item| {
item.token.id == mutation.token_id
&& expected_user_id
.map(|user_id| item.token.user_id == user_id)
.unwrap_or(true)
}) else {
return Ok(None);
};
Self::remove_hash_for_token(&mut hashes, &mutation.token_id);
hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone());
item.token.token_prefix = mutation.token_prefix.clone();
item.token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(item.token.clone()))
}
}
#[async_trait]
@@ -167,76 +351,27 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
&self,
record: &UpdateManagementTokenRecord,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
self.update_management_token_scoped(record, None)
}
let mut items = self
.items
.write()
.expect("management token repository lock");
let Some(index) = items
.iter()
.position(|item| item.token.id == record.token_id)
else {
return Ok(None);
};
if let Some(name) = &record.name {
if items.iter().enumerate().any(|(position, item)| {
position != index
&& item.token.user_id == items[index].token.user_id
&& item.token.name == *name
}) {
return Err(DataLayerError::InvalidInput(format!(
"已存在名为 '{}' 的 Token",
name
)));
}
items[index].token.name = name.clone();
}
if record.clear_description {
items[index].token.description = None;
} else if let Some(description) = &record.description {
items[index].token.description = Some(description.clone());
}
if record.clear_allowed_ips {
items[index].token.allowed_ips = None;
} else if let Some(allowed_ips) = &record.allowed_ips {
items[index].token.allowed_ips = Some(allowed_ips.clone());
}
if let Some(permissions) = &record.permissions {
items[index].token.permissions = Some(permissions.clone());
}
if record.clear_expires_at {
items[index].token.expires_at_unix_secs = None;
} else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs {
items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs);
}
if let Some(is_active) = record.is_active {
items[index].token.is_active = is_active;
}
items[index].token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(items[index].token.clone()))
async fn update_management_token_for_user(
&self,
record: &UpdateManagementTokenRecord,
user_id: &str,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
self.update_management_token_scoped(record, Some(user_id))
}
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
let original_len = items.len();
items.retain(|item| item.token.id != token_id);
Self::remove_hash_for_token(&mut hashes, token_id);
Ok(items.len() != original_len)
Ok(self.delete_management_token_scoped(token_id, None))
}
async fn delete_management_token_for_user(
&self,
token_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
Ok(self.delete_management_token_scoped(token_id, Some(user_id)))
}
async fn set_management_token_active(
@@ -244,43 +379,45 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
token_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let mut items = self
.items
.write()
.expect("management token repository lock");
let Some(item) = items.iter_mut().find(|item| item.token.id == token_id) else {
return Ok(None);
};
item.token.is_active = is_active;
item.token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(item.token.clone()))
Ok(self.set_management_token_active_scoped(token_id, None, is_active))
}
async fn set_management_token_active_for_user(
&self,
token_id: &str,
user_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
Ok(self.set_management_token_active_scoped(token_id, Some(user_id), is_active))
}
async fn activate_management_token_if_matches(
&self,
mutation: &ActivateManagementTokenIfMatches,
) -> Result<bool, DataLayerError> {
self.activate_management_token_if_matches_inner(mutation)
}
async fn delete_inactive_management_token_if_matches(
&self,
mutation: &ActivateManagementTokenIfMatches,
) -> Result<bool, DataLayerError> {
self.delete_inactive_management_token_if_matches_inner(mutation)
}
async fn regenerate_management_token_secret(
&self,
mutation: &RegenerateManagementTokenSecret,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
self.regenerate_management_token_secret_scoped(mutation, None)
}
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
let Some(item) = items
.iter_mut()
.find(|item| item.token.id == mutation.token_id)
else {
return Ok(None);
};
Self::remove_hash_for_token(&mut hashes, &mutation.token_id);
hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone());
item.token.token_prefix = mutation.token_prefix.clone();
item.token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(item.token.clone()))
async fn regenerate_management_token_secret_for_user(
&self,
mutation: &RegenerateManagementTokenSecret,
user_id: &str,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
self.regenerate_management_token_secret_scoped(mutation, Some(user_id))
}
async fn record_management_token_usage(
@@ -307,10 +444,10 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
mod tests {
use super::InMemoryManagementTokenRepository;
use crate::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
ManagementTokenReadRepository, ManagementTokenWriteRepository,
RegenerateManagementTokenSecret, StoredManagementToken, StoredManagementTokenUserSummary,
StoredManagementTokenWithUser, UpdateManagementTokenRecord,
};
fn sample_token(id: &str, user_id: &str, is_active: bool) -> StoredManagementTokenWithUser {
@@ -456,4 +593,108 @@ mod tests {
.expect("hash lookup should succeed");
assert!(deleted_by_hash.is_none());
}
#[tokio::test]
async fn owner_scoped_mutations_never_cross_user_boundaries() {
let repository = InMemoryManagementTokenRepository::seed_with_hashes(
vec![sample_token("token-1", "user-1", true)],
vec![("hash-1".to_string(), "token-1".to_string())],
);
let update = UpdateManagementTokenRecord {
token_id: "token-1".to_string(),
name: Some("hijacked".to_string()),
description: None,
clear_description: false,
allowed_ips: None,
clear_allowed_ips: false,
permissions: None,
expires_at_unix_secs: None,
clear_expires_at: false,
is_active: None,
};
assert!(repository
.update_management_token_for_user(&update, "user-2")
.await
.expect("scoped update should execute")
.is_none());
assert!(repository
.set_management_token_active_for_user("token-1", "user-2", false)
.await
.expect("scoped toggle should execute")
.is_none());
assert!(repository
.regenerate_management_token_secret_for_user(
&RegenerateManagementTokenSecret {
token_id: "token-1".to_string(),
token_hash: "hash-hijacked".to_string(),
token_prefix: Some("ae_hijacked".to_string()),
},
"user-2",
)
.await
.expect("scoped regeneration should execute")
.is_none());
assert!(!repository
.delete_management_token_for_user("token-1", "user-2")
.await
.expect("scoped delete should execute"));
let unchanged = repository
.get_management_token_with_user_by_hash("hash-1")
.await
.expect("original hash lookup should succeed")
.expect("token should remain");
assert_eq!(unchanged.token.name, "token-1");
assert!(unchanged.token.is_active);
assert!(repository
.get_management_token_with_user_by_hash("hash-hijacked")
.await
.expect("replacement hash lookup should succeed")
.is_none());
}
#[tokio::test]
async fn install_activation_fails_closed_without_atomic_user_state() {
let mut pending = sample_token("token-1", "user-1", false);
pending.token.allowed_ips = Some(serde_json::json!(["127.0.0.1"]));
pending.token.permissions = Some(serde_json::json!(["admin:proxy_nodes:write"]));
pending.token.expires_at_unix_secs = Some(1_800_000_000);
let expected_token = pending.token.clone();
let repository = InMemoryManagementTokenRepository::seed_with_hashes(
[pending],
[("hash-1".to_string(), "token-1".to_string())],
);
let expected = ActivateManagementTokenIfMatches {
expected_token,
token_hash: "hash-1".to_string(),
expected_user_security_version: 4,
now_unix_secs: 1_700_000_000,
};
let mut mismatched = expected.clone();
mismatched.expected_token.permissions =
Some(serde_json::json!(["admin:proxy_nodes:admin"]));
assert!(!repository
.activate_management_token_if_matches(&mismatched)
.await
.expect("mismatched activation should execute"));
assert!(!repository
.activate_management_token_if_matches(&expected)
.await
.expect("memory activation should fail closed"));
assert!(
!repository
.get_management_token_with_user("token-1")
.await
.expect("token lookup should execute")
.expect("token should remain")
.token
.is_active
);
assert!(repository
.delete_inactive_management_token_if_matches(&expected)
.await
.expect("exact inactive snapshot cleanup should execute"));
}
}
@@ -1,10 +1,10 @@
mod memory;
pub use aether_data_contracts::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary,
StoredManagementTokenWithUser, UpdateManagementTokenRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlManagementTokenRepository;
@@ -7,7 +7,7 @@ use async_trait::async_trait;
use crate::DataLayerError;
use aether_data_contracts::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
StoredOAuthProviderConfig, UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord,
};
#[derive(Debug, Default)]
@@ -65,15 +65,30 @@ impl OAuthProviderReadRepository for InMemoryOAuthProviderRepository {
#[async_trait]
impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository {
async fn upsert_oauth_provider_config(
async fn upsert_oauth_provider_config_guarded(
&self,
record: &UpsertOAuthProviderConfigRecord,
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
_ldap_exclusive: bool,
force_disable: bool,
locked_users_snapshot: usize,
) -> Result<UpsertOAuthProviderConfigOutcome, DataLayerError> {
record.validate()?;
let mut items = self.items.write().expect("oauth provider repository lock");
let now = Self::now_unix_secs();
let existing = items.get(&record.provider_type).cloned();
if !force_disable
&& locked_users_snapshot > 0
&& existing
.as_ref()
.is_some_and(|provider| provider.is_enabled && !record.is_enabled)
{
return Ok(
UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation {
affected_count: locked_users_snapshot,
},
);
}
let created_at = existing
.as_ref()
.and_then(|item| item.created_at_unix_ms)
@@ -106,13 +121,34 @@ impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository {
.with_timestamps(created_at, now);
items.insert(record.provider_type.clone(), item.clone());
Ok(item)
Ok(UpsertOAuthProviderConfigOutcome::Upserted(item))
}
async fn delete_oauth_provider_config(
async fn compare_and_swap_oauth_provider_client_secret(
&self,
provider_type: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let mut items = self.items.write().expect("oauth provider repository lock");
let Some(item) = items.get_mut(provider_type) else {
return Ok(false);
};
if item.client_secret_encrypted.as_deref() != Some(expected) {
return Ok(false);
}
item.client_secret_encrypted = Some(replacement.to_string());
Ok(true)
}
async fn delete_oauth_provider_config_if_unlinked(
&self,
provider_type: &str,
has_links_snapshot: bool,
) -> Result<bool, DataLayerError> {
if has_links_snapshot {
return Ok(false);
}
let mut items = self.items.write().expect("oauth provider repository lock");
Ok(items.remove(provider_type).is_some())
}
@@ -123,7 +159,8 @@ mod tests {
use super::InMemoryOAuthProviderRepository;
use crate::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
StoredOAuthProviderConfig, UpsertOAuthProviderConfigOutcome,
UpsertOAuthProviderConfigRecord,
};
fn sample_provider(provider_type: &str) -> StoredOAuthProviderConfig {
@@ -138,19 +175,31 @@ mod tests {
}
fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord {
let is_custom_oidc = provider_type.starts_with("custom_oidc");
let endpoint_host = if is_custom_oidc {
"idp.example".to_string()
} else {
format!("{provider_type}.example.com")
};
UpsertOAuthProviderConfigRecord {
provider_type: provider_type.to_string(),
display_name: format!("{provider_type} display"),
client_id: format!("{provider_type}-client"),
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")),
token_url_override: Some(format!("https://{provider_type}.example.com/token")),
userinfo_url_override: None,
authorization_url_override: Some(format!("https://{endpoint_host}/auth")),
token_url_override: Some(format!("https://{endpoint_host}/token")),
userinfo_url_override: is_custom_oidc
.then(|| format!("https://{endpoint_host}/userinfo")),
scopes: Some(vec!["openid".to_string(), "profile".to_string()]),
redirect_uri: format!("https://{provider_type}.example.com/redirect"),
frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(),
attribute_mapping: Some(serde_json::json!({"email": "email"})),
extra_config: Some(serde_json::json!({"team": true})),
extra_config: is_custom_oidc.then(|| {
serde_json::json!({
"allowed_domains": [endpoint_host],
"team": true,
})
}),
icon_url: None,
is_enabled: true,
}
@@ -171,28 +220,106 @@ mod tests {
assert_eq!(listed[0].provider_type, "github");
assert_eq!(listed[1].provider_type, "linuxdo");
let created = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
..sample_upsert("google")
})
let UpsertOAuthProviderConfigOutcome::Upserted(created) = repository
.upsert_oauth_provider_config_guarded(
&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
..sample_upsert("custom_oidc")
},
false,
false,
0,
)
.await
.expect("create should succeed");
.expect("create should succeed")
else {
panic!("create unexpectedly required confirmation");
};
assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1"));
let updated = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Clear,
..sample_upsert("google")
})
let UpsertOAuthProviderConfigOutcome::Upserted(updated) = repository
.upsert_oauth_provider_config_guarded(
&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Clear,
..sample_upsert("custom_oidc")
},
false,
false,
0,
)
.await
.expect("update should succeed");
.expect("update should succeed")
else {
panic!("update unexpectedly required confirmation");
};
assert!(updated.client_secret_encrypted.is_none());
let deleted = repository
.delete_oauth_provider_config("google")
.delete_oauth_provider_config_if_unlinked("custom_oidc", false)
.await
.expect("delete should succeed");
assert!(deleted);
}
#[tokio::test]
async fn client_secret_cas_preserves_concurrent_non_secret_fields_and_timestamp() {
let repository = InMemoryOAuthProviderRepository::default();
repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Set("legacy-secret".to_string()),
..sample_upsert("custom_oidc")
})
.await
.expect("provider should create");
let concurrent = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
display_name: "concurrent display update".to_string(),
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
..sample_upsert("custom_oidc")
})
.await
.expect("non-secret update should persist");
assert!(repository
.compare_and_swap_oauth_provider_client_secret(
"custom_oidc",
"legacy-secret",
"record-bound-v2",
)
.await
.expect("secret CAS should execute"));
let migrated = repository
.get_oauth_provider_config("custom_oidc")
.await
.expect("provider should read")
.expect("provider should exist");
assert_eq!(migrated.display_name, "concurrent display update");
assert_eq!(
migrated.updated_at_unix_secs,
concurrent.updated_at_unix_secs
);
assert_eq!(
migrated.client_secret_encrypted.as_deref(),
Some("record-bound-v2")
);
assert!(!repository
.compare_and_swap_oauth_provider_client_secret(
"custom_oidc",
"legacy-secret",
"must-not-win",
)
.await
.expect("stale CAS should execute"));
assert_eq!(
repository
.get_oauth_provider_config("custom_oidc")
.await
.expect("provider should read")
.expect("provider should exist")
.client_secret_encrypted
.as_deref(),
Some("record-bound-v2")
);
}
}
@@ -1,8 +1,10 @@
mod memory;
pub use aether_data_contracts::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderRepository,
OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config,
validate_oauth_redirect_uri, EncryptedSecretUpdate, OAuthProviderReadRepository,
OAuthProviderRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlOAuthProviderRepository;
@@ -7,10 +7,12 @@ use serde_json::{json, Map, Value};
use super::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate,
ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -454,6 +456,44 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
Ok(stored.clone())
}
async fn compare_and_swap_provider_config(
&self,
update: &ProviderCatalogProviderConfigCasUpdate,
) -> Result<bool, DataLayerError> {
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(provider) = index.providers.get_mut(&update.provider_id) else {
return Ok(false);
};
if provider.config != update.expected_config {
return Ok(false);
}
provider.config = update.config.clone();
provider.updated_at_unix_secs = Some(current_unix_secs());
Ok(true)
}
async fn compare_and_swap_provider_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(provider) = index.providers.get_mut(&update.record_id) else {
return Ok(false);
};
if provider.proxy != update.expected_proxy {
return Ok(false);
}
provider.proxy = update.proxy.clone();
provider.updated_at_unix_secs = Some(current_unix_secs());
Ok(true)
}
async fn delete_provider(&self, provider_id: &str) -> Result<bool, DataLayerError> {
let mut index = self
.index
@@ -504,6 +544,25 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
Ok(stored.clone())
}
async fn compare_and_swap_endpoint_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(endpoint) = index.endpoints.get_mut(&update.record_id) else {
return Ok(false);
};
if endpoint.proxy != update.expected_proxy {
return Ok(false);
}
endpoint.proxy = update.proxy.clone();
endpoint.updated_at_unix_secs = Some(current_unix_secs());
Ok(true)
}
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, DataLayerError> {
let mut index = self
.index
@@ -548,6 +607,52 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
Ok(stored.clone())
}
async fn compare_and_swap_key_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(key) = index.keys.get_mut(&update.record_id) else {
return Ok(false);
};
if key.proxy != update.expected_proxy {
return Ok(false);
}
key.proxy = update.proxy.clone();
key.updated_at_unix_secs = Some(current_unix_secs());
Ok(true)
}
async fn compare_and_swap_key_credentials(
&self,
update: &ProviderCatalogKeyCredentialsCasUpdate,
) -> Result<bool, DataLayerError> {
if update.key_id.trim().is_empty() || update.expected_provider_id.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
"provider catalog key credential CAS requires key_id and provider_id".to_string(),
));
}
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(key) = index.keys.get_mut(&update.key_id) else {
return Ok(false);
};
if key.provider_id != update.expected_provider_id
|| key.encrypted_api_key != update.expected_encrypted_api_key
|| key.encrypted_auth_config != update.expected_encrypted_auth_config
{
return Ok(false);
}
key.encrypted_api_key = update.encrypted_api_key.clone();
key.encrypted_auth_config = update.encrypted_auth_config.clone();
Ok(true)
}
async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
@@ -882,40 +987,11 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
Ok(true)
}
async fn update_key_oauth_credentials(
&self,
key_id: &str,
encrypted_api_key: &str,
encrypted_auth_config: Option<&str>,
expires_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
if encrypted_api_key.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
"provider catalog oauth api_key is empty".to_string(),
));
}
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(key) = index.keys.get_mut(key_id) else {
return Ok(false);
};
key.encrypted_api_key = Some(encrypted_api_key.to_string());
key.encrypted_auth_config = encrypted_auth_config.map(ToOwned::to_owned);
key.expires_at_unix_secs = expires_at_unix_secs;
key.updated_at_unix_secs = Some(current_unix_secs());
Ok(true)
}
async fn update_key_oauth_runtime_state(
&self,
key_id: &str,
oauth_invalid_at_unix_secs: Option<u64>,
oauth_invalid_reason: Option<&str>,
encrypted_auth_config_update: Option<&str>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
let mut index = self
@@ -928,9 +1004,6 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
key.oauth_invalid_at_unix_secs = oauth_invalid_at_unix_secs;
key.oauth_invalid_reason = oauth_invalid_reason.map(ToOwned::to_owned);
if let Some(encrypted_auth_config) = encrypted_auth_config_update {
key.encrypted_auth_config = Some(encrypted_auth_config.to_string());
}
key.updated_at_unix_secs = Some(updated_at_unix_secs.unwrap_or_else(current_unix_secs));
Ok(true)
}
@@ -1576,15 +1649,15 @@ mod tests {
}
#[tokio::test]
async fn updates_oauth_credentials_for_existing_key() {
async fn unfenced_oauth_runtime_state_update_preserves_credentials() {
let repository = InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-1")],
vec![sample_endpoint("endpoint-1", "provider-1")],
vec![sample_key("key-1", "provider-1")
.with_transport_fields(
None,
"ciphertext-placeholder".to_string(),
Some("ciphertext-auth-1".to_string()),
"ciphertext-api".to_string(),
Some("ciphertext-auth".to_string()),
None,
None,
None,
@@ -1596,29 +1669,26 @@ mod tests {
);
assert!(repository
.update_key_oauth_credentials(
"key-1",
"ciphertext-updated-token",
Some("ciphertext-auth-2"),
Some(4_102_444_800),
)
.update_key_oauth_runtime_state("key-1", Some(123), Some("refresh failed"), Some(456),)
.await
.expect("update should succeed"));
.expect("runtime state should update"));
let stored = repository
.list_keys_by_ids(&["key-1".to_string()])
.await
.expect("keys should read");
assert_eq!(stored.len(), 1);
.expect("key should read")
.pop()
.expect("key should exist");
assert_eq!(stored.encrypted_api_key.as_deref(), Some("ciphertext-api"));
assert_eq!(
stored[0].encrypted_api_key.as_deref(),
Some("ciphertext-updated-token")
stored.encrypted_auth_config.as_deref(),
Some("ciphertext-auth")
);
assert_eq!(stored.oauth_invalid_at_unix_secs, Some(123));
assert_eq!(
stored[0].encrypted_auth_config.as_deref(),
Some("ciphertext-auth-2")
stored.oauth_invalid_reason.as_deref(),
Some("refresh failed")
);
assert_eq!(stored[0].expires_at_unix_secs, Some(4_102_444_800));
}
#[tokio::test]
@@ -3,11 +3,12 @@ mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate,
ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogUpstreamMetadataNamespaceExpectation,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
@@ -1,3 +1,5 @@
use sha2::{Digest, Sha256};
const KIRO_DEVICE_AUTH_SESSION_PREFIX: &str = "device_auth_session:";
const PROVIDER_OAUTH_BATCH_TASK_PREFIX: &str = "provider_oauth_batch_task:";
const PROVIDER_OAUTH_STATE_PREFIX: &str = "provider_oauth_state:";
@@ -8,7 +10,11 @@ pub const PROVIDER_OAUTH_STATE_TTL_SECS: u64 = 600;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminProviderOAuthDeviceSession {
pub session_id: String,
pub provider_id: String,
pub initiated_by_user_id: String,
pub initiated_by_session_id: Option<String>,
pub initiated_by_management_token_id: Option<String>,
pub region: String,
pub client_id: String,
pub client_secret: String,
@@ -34,28 +40,54 @@ pub struct StoredAdminProviderOAuthDeviceSession {
pub error_msg: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminProviderOAuthState {
pub nonce: String,
pub key_id: String,
pub provider_id: String,
pub provider_type: String,
pub pkce_verifier: Option<String>,
#[serde(default)]
pub expected_encrypted_auth_config: Option<String>,
pub initiated_by_user_id: String,
#[serde(default)]
pub initiated_by_session_id: Option<String>,
#[serde(default)]
pub initiated_by_management_token_id: Option<String>,
pub created_at: u64,
}
pub fn provider_oauth_device_session_storage_key(session_id: &str) -> String {
format!("{KIRO_DEVICE_AUTH_SESSION_PREFIX}{session_id}")
}
pub fn provider_oauth_device_session_secret_purpose(session_id: &str) -> String {
let storage_key = provider_oauth_device_session_storage_key(session_id);
format!(
"provider-oauth-device-session:sha256:{:x}",
Sha256::digest(storage_key.as_bytes())
)
}
pub fn provider_oauth_state_storage_key(nonce: &str) -> String {
format!("{PROVIDER_OAUTH_STATE_PREFIX}{nonce}")
format!(
"{PROVIDER_OAUTH_STATE_PREFIX}sha256:{:x}",
Sha256::digest(nonce.as_bytes())
)
}
pub fn provider_oauth_batch_task_storage_key(task_id: &str) -> String {
format!("{PROVIDER_OAUTH_BATCH_TASK_PREFIX}{task_id}")
}
pub fn provider_oauth_batch_task_secret_purpose(task_id: &str) -> String {
let storage_key = provider_oauth_batch_task_storage_key(task_id);
format!(
"provider-oauth-batch-task:sha256:{:x}",
Sha256::digest(storage_key.as_bytes())
)
}
pub fn build_provider_oauth_batch_task_status_payload(
provider_id: &str,
state: &serde_json::Map<String, serde_json::Value>,
@@ -142,7 +174,8 @@ pub fn build_provider_oauth_batch_task_status_payload(
#[cfg(test)]
mod tests {
use super::{
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key,
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_secret_purpose,
provider_oauth_batch_task_storage_key, provider_oauth_device_session_secret_purpose,
provider_oauth_device_session_storage_key, provider_oauth_state_storage_key,
KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS,
PROVIDER_OAUTH_STATE_TTL_SECS,
@@ -155,14 +188,23 @@ mod tests {
provider_oauth_device_session_storage_key("session-123"),
"device_auth_session:session-123"
);
assert_eq!(
provider_oauth_state_storage_key("nonce-123"),
"provider_oauth_state:nonce-123"
);
let first_purpose = provider_oauth_device_session_secret_purpose("session-123");
let second_purpose = provider_oauth_device_session_secret_purpose("session-456");
assert!(first_purpose.starts_with("provider-oauth-device-session:sha256:"));
assert!(!first_purpose.contains("session-123"));
assert_ne!(first_purpose, second_purpose);
let state_key = provider_oauth_state_storage_key("nonce-123");
assert!(state_key.starts_with("provider_oauth_state:sha256:"));
assert!(!state_key.contains("nonce-123"));
assert_eq!(
provider_oauth_batch_task_storage_key("task-123"),
"provider_oauth_batch_task:task-123"
);
let first_task_purpose = provider_oauth_batch_task_secret_purpose("task-123");
let second_task_purpose = provider_oauth_batch_task_secret_purpose("task-456");
assert!(first_task_purpose.starts_with("provider-oauth-batch-task:sha256:"));
assert!(!first_task_purpose.contains("task-123"));
assert_ne!(first_task_purpose, second_task_purpose);
}
#[test]
@@ -10,13 +10,14 @@ use super::log_reported_tunnel_error_event;
use crate::DataLayerError;
use aether_data_contracts::repository::proxy_nodes::{
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
normalize_proxy_metadata, preserve_proxy_metadata_tunnel_security,
reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary,
ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation,
ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation,
ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent,
StoredProxyNodeMetricsBucket, TunnelMetricsSample, PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata,
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery,
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket,
StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, TunnelMetricsSample,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
};
#[derive(Debug, Default)]
@@ -357,6 +358,42 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
Ok(updated)
}
async fn compare_and_set_proxy_password(
&self,
node_id: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let Some(node) = nodes.get_mut(node_id) else {
return Ok(false);
};
if node.proxy_password.as_deref() != Some(expected) {
return Ok(false);
}
node.proxy_password = Some(replacement.to_string());
node.updated_at_unix_secs = Self::now_unix_secs();
Ok(true)
}
async fn compare_and_set_proxy_metadata(
&self,
node_id: &str,
expected: &serde_json::Value,
replacement: &serde_json::Value,
) -> Result<bool, DataLayerError> {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let Some(node) = nodes.get_mut(node_id) else {
return Ok(false);
};
if node.proxy_metadata.as_ref() != Some(expected) {
return Ok(false);
}
node.proxy_metadata = Some(replacement.clone());
node.updated_at_unix_secs = Self::now_unix_secs();
Ok(true)
}
async fn create_manual_node(
&self,
mutation: &ProxyNodeManualCreateMutation,
@@ -369,9 +406,14 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
return Err(Self::duplicate_proxy_node_error(existing));
}
let node_id = requested_proxy_node_id(mutation.node_id.as_deref())?
.unwrap_or_else(|| Uuid::new_v4().to_string());
if let Some(existing) = nodes.get(&node_id) {
return Err(proxy_node_id_in_use_error(existing));
}
let now = Self::now_unix_secs();
let node = StoredProxyNode::new(
Uuid::new_v4().to_string(),
node_id,
mutation.name.clone(),
mutation.ip.clone(),
mutation.port,
@@ -466,18 +508,35 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
) -> Result<StoredProxyNode, DataLayerError> {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let now = Self::now_unix_secs();
let normalized_proxy_metadata = normalize_proxy_metadata(
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
let normalized_proxy_metadata = merge_proxy_metadata_for_registration(
None,
normalize_proxy_metadata(
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
),
);
if let Some(existing_id) = nodes
.iter()
.find(|(_, node)| {
.filter(|(_, node)| {
!node.is_manual && node.ip == mutation.ip && node.port == mutation.port
})
.min_by(|(_, left), (_, right)| {
left.created_at_unix_ms
.unwrap_or(u64::MAX)
.cmp(&right.created_at_unix_ms.unwrap_or(u64::MAX))
.then(left.id.cmp(&right.id))
})
.map(|(node_id, _)| node_id.clone())
{
if let Some(requested_id) = requested_proxy_node_id(mutation.node_id.as_deref())? {
if requested_id != existing_id {
return Err(proxy_node_registration_identity_error(
&requested_id,
&existing_id,
));
}
}
let node = nodes
.get_mut(&existing_id)
.expect("existing proxy node should be present");
@@ -505,9 +564,10 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
if let Some(estimated_max_concurrency) = mutation.estimated_max_concurrency {
node.estimated_max_concurrency = Some(estimated_max_concurrency);
}
if let Some(proxy_metadata) = normalized_proxy_metadata {
node.proxy_metadata = Some(proxy_metadata);
}
node.proxy_metadata = merge_proxy_metadata_for_registration(
node.proxy_metadata.as_ref(),
normalized_proxy_metadata,
);
if node.created_at_unix_ms.is_none() {
node.created_at_unix_ms = now;
}
@@ -515,8 +575,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
return Ok(node.clone());
}
let node_id = requested_proxy_node_id(mutation.node_id.as_deref())?
.unwrap_or_else(|| Uuid::new_v4().to_string());
if let Some(existing) = nodes.get(&node_id) {
return Err(proxy_node_id_in_use_error(existing));
}
let mut node = StoredProxyNode::new(
Uuid::new_v4().to_string(),
node_id,
mutation.name.clone(),
mutation.ip.clone(),
mutation.port,
@@ -555,11 +620,18 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
&self,
mutation: &ProxyNodeHeartbeatMutation,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let (node, sample, now_unix_secs) = {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let Some(node) = nodes.get_mut(&mutation.node_id) else {
return Ok(None);
};
if mutation
.expected_tunnel_generation
.as_deref()
.is_some_and(|expected| expected != node.tunnel_generation)
{
return Ok(None);
}
if !node.tunnel_mode {
return Err(DataLayerError::InvalidInput(
"non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode"
@@ -587,14 +659,11 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
if let Some(value) = mutation.avg_latency_ms {
node.avg_latency_ms = Some(value);
}
let normalized_proxy_metadata = normalize_proxy_metadata(
let normalized_proxy_metadata = normalize_heartbeat_proxy_metadata(
previous_proxy_metadata.as_ref(),
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
);
let normalized_proxy_metadata = preserve_proxy_metadata_tunnel_security(
previous_proxy_metadata.as_ref(),
normalized_proxy_metadata,
);
if let Some(value) = normalized_proxy_metadata {
node.proxy_metadata = Some(value);
}
@@ -671,6 +740,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
});
}
}
drop(nodes);
Ok(Some(node))
}
@@ -686,6 +756,12 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
if !node.is_manual {
return Ok(false);
}
let Some(expected_generation) = mutation.expected_tunnel_generation.as_deref() else {
return Ok(false);
};
if expected_generation != node.tunnel_generation {
return Ok(false);
}
node.total_requests += mutation.total_requests_delta.max(0);
node.failed_requests += mutation.failed_requests_delta.max(0);
@@ -703,6 +779,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
let Some(node) = nodes.get_mut(&mutation.node_id) else {
return Ok(None);
};
if mutation
.expected_tunnel_generation
.as_deref()
.is_some_and(|expected| expected != node.tunnel_generation)
{
return Ok(None);
}
let event_time = mutation
.observed_at_unix_secs
@@ -776,11 +859,11 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
}
async fn delete_node(&self, node_id: &str) -> Result<Option<StoredProxyNode>, DataLayerError> {
let removed = self
.nodes
.write()
.expect("proxy node repository lock")
.remove(node_id);
// Keep the parent lock until all child state is removed. Registration
// also takes this lock, so the same id cannot be recreated between the
// parent delete and cleanup of its events or metrics.
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let removed = nodes.remove(node_id);
if removed.is_some() {
self.events
.write()
@@ -795,6 +878,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
.expect("proxy node repository lock")
.retain(|(metric_node_id, _), _| metric_node_id != node_id);
}
drop(nodes);
Ok(removed)
}
@@ -806,6 +890,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
let Some(node) = nodes.get_mut(&mutation.node_id) else {
return Ok(None);
};
if mutation
.expected_tunnel_generation
.as_deref()
.is_some_and(|expected| expected != node.tunnel_generation)
{
return Ok(None);
}
if node.is_manual {
return Err(DataLayerError::InvalidInput(
"手动节点不支持远程配置下发".to_string(),
@@ -886,6 +977,31 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
}
}
fn requested_proxy_node_id(value: Option<&str>) -> Result<Option<String>, DataLayerError> {
let Some(value) = value else {
return Ok(None);
};
if value.is_empty() || value.trim() != value {
return Err(DataLayerError::InvalidInput(
"proxy node id must be non-empty and unpadded".to_string(),
));
}
Ok(Some(value.to_string()))
}
fn proxy_node_registration_identity_error(requested_id: &str, existing_id: &str) -> DataLayerError {
DataLayerError::InvalidInput(format!(
"proxy node registration identity changed: requested {requested_id}, existing {existing_id}"
))
}
fn proxy_node_id_in_use_error(node: &StoredProxyNode) -> DataLayerError {
DataLayerError::InvalidInput(format!(
"proxy node id is already in use: {} ({}:{})",
node.id, node.ip, node.port
))
}
#[cfg(test)]
mod tests {
use super::InMemoryProxyNodeRepository;
@@ -894,7 +1010,7 @@ mod tests {
ProxyNodeRemoteConfigMutation, ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository,
StoredProxyNode, StoredProxyNodeEvent,
};
use serde_json::json;
use serde_json::{json, Value};
fn sample_node() -> StoredProxyNode {
StoredProxyNode::new(
@@ -937,6 +1053,7 @@ mod tests {
let heartbeat = repository
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
node_id: "node-1".to_string(),
expected_tunnel_generation: None,
heartbeat_interval: Some(45),
active_connections: Some(5),
total_requests_delta: Some(8),
@@ -944,7 +1061,13 @@ mod tests {
failed_requests_delta: Some(2),
dns_failures_delta: Some(1),
stream_errors_delta: Some(3),
proxy_metadata: Some(json!({"arch": "arm64"})),
proxy_metadata: Some(json!({
"arch": "arm64",
"tunnel_security": {
"mode": "disabled",
"encryption_key": "attacker-controlled"
}
})),
proxy_version: Some("1.2.3".to_string()),
})
.await
@@ -966,10 +1089,16 @@ mod tests {
.and_then(|value| value.as_str()),
Some("1.2.3")
);
assert!(heartbeat
.proxy_metadata
.as_ref()
.and_then(|value| value.get("tunnel_security"))
.is_none());
let stale = repository
.update_tunnel_status(&ProxyNodeTunnelStatusMutation {
node_id: "node-1".to_string(),
expected_tunnel_generation: None,
connected: false,
conn_count: 0,
detail: None,
@@ -994,6 +1123,7 @@ mod tests {
let updated = repository
.update_tunnel_status(&ProxyNodeTunnelStatusMutation {
node_id: "node-1".to_string(),
expected_tunnel_generation: None,
connected: false,
conn_count: 0,
detail: None,
@@ -1060,6 +1190,96 @@ mod tests {
assert_eq!(events[0].detail.as_deref(), Some("newer"));
}
#[tokio::test]
async fn delete_cleans_child_state_before_same_id_can_be_reused() {
let old_node = sample_node();
let old_generation = old_node.tunnel_generation.clone();
let repository = InMemoryProxyNodeRepository::seed_with_events(
vec![old_node],
vec![StoredProxyNodeEvent {
id: 1,
node_id: "node-1".to_string(),
event_type: "connected".to_string(),
detail: Some("old incarnation".to_string()),
event_metadata: None,
created_at_unix_ms: Some(1_710_000_000),
}],
);
repository
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
node_id: "node-1".to_string(),
expected_tunnel_generation: Some(old_generation.clone()),
heartbeat_interval: None,
active_connections: Some(1),
total_requests_delta: None,
avg_latency_ms: None,
failed_requests_delta: None,
dns_failures_delta: None,
stream_errors_delta: None,
proxy_metadata: Some(json!({
"tunnel_metrics": {
"connect_errors": 0,
"disconnects": 0,
"error_events_total": 0,
"ws_in_bytes": 0,
"ws_out_bytes": 0,
"ws_in_frames": 0,
"ws_out_frames": 0,
"heartbeat_rtt_last_ms": 1
}
})),
proxy_version: None,
})
.await
.expect("heartbeat should create metric buckets")
.expect("old node should exist");
repository
.delete_node("node-1")
.await
.expect("delete should succeed")
.expect("old node should be removed");
let replacement = repository
.register_node(&ProxyNodeRegistrationMutation {
node_id: Some("node-1".to_string()),
name: "replacement".to_string(),
ip: "127.0.0.2".to_string(),
port: 7002,
region: None,
heartbeat_interval: 30,
active_connections: None,
total_requests: None,
avg_latency_ms: None,
hardware_info: None,
estimated_max_concurrency: None,
proxy_metadata: None,
proxy_version: None,
registered_by: None,
tunnel_mode: true,
})
.await
.expect("same id should be reusable after delete");
assert_ne!(replacement.tunnel_generation, old_generation);
assert!(repository
.list_proxy_node_events("node-1", 10)
.await
.expect("events should read")
.is_empty());
for step in [
crate::repository::proxy_nodes::ProxyNodeMetricsStep::OneMinute,
crate::repository::proxy_nodes::ProxyNodeMetricsStep::OneHour,
] {
assert!(repository
.list_proxy_node_metrics("node-1", step, 0, u64::MAX, 10)
.await
.expect("metrics should read")
.is_empty());
}
}
#[tokio::test]
async fn resets_stale_tunnel_statuses_without_touching_manual_nodes() {
let mut stale_tunnel = sample_node();
@@ -1101,12 +1321,127 @@ mod tests {
assert_eq!(manual.active_connections, 4);
}
#[tokio::test]
async fn registration_rejects_rebinding_existing_endpoint_to_different_node_id() {
let repository = InMemoryProxyNodeRepository::default();
let mutation = ProxyNodeRegistrationMutation {
node_id: Some("stable-node-id".to_string()),
name: "stable-node".to_string(),
ip: "127.0.0.9".to_string(),
port: 7009,
region: None,
heartbeat_interval: 30,
active_connections: None,
total_requests: None,
avg_latency_ms: None,
hardware_info: None,
estimated_max_concurrency: None,
proxy_metadata: Some(json!({"secret_marker": "first"})),
proxy_version: None,
registered_by: None,
tunnel_mode: true,
};
let registered = repository
.register_node(&mutation)
.await
.expect("initial registration should succeed");
assert_eq!(registered.id, "stable-node-id");
let mut conflicting = mutation;
conflicting.node_id = Some("replacement-node-id".to_string());
conflicting.proxy_metadata = Some(json!({"secret_marker": "replacement"}));
assert!(repository.register_node(&conflicting).await.is_err());
let persisted = repository
.find_proxy_node("stable-node-id")
.await
.expect("stable node should read")
.expect("stable node should remain");
assert_eq!(
persisted
.proxy_metadata
.as_ref()
.and_then(|value| value.get("secret_marker")),
Some(&json!("first"))
);
}
#[tokio::test]
async fn registration_preserves_omitted_security_and_allows_rotation() {
let repository = InMemoryProxyNodeRepository::default();
let first_mutation = ProxyNodeRegistrationMutation {
node_id: Some("registration-security-node".to_string()),
name: "registration-security-node".to_string(),
ip: "127.0.0.70".to_string(),
port: 7070,
region: None,
heartbeat_interval: 30,
active_connections: None,
total_requests: None,
avg_latency_ms: None,
hardware_info: None,
estimated_max_concurrency: None,
proxy_metadata: Some(json!({
"version": "1.0.0",
"tunnel_security": {
"mode": "non_tls_required",
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
}
})),
proxy_version: None,
registered_by: None,
tunnel_mode: true,
};
let first = repository
.register_node(&first_mutation)
.await
.expect("first registration should succeed");
let mut refreshed_mutation = first_mutation.clone();
refreshed_mutation.name = "registration-security-node-refreshed".to_string();
refreshed_mutation.proxy_metadata = Some(json!({"runtime": "refreshed"}));
refreshed_mutation.proxy_version = Some("2.0.0".to_string());
let refreshed = repository
.register_node(&refreshed_mutation)
.await
.expect("metadata-only re-registration should succeed");
assert_eq!(refreshed.id, first.id);
assert_eq!(
refreshed
.proxy_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted"))
.and_then(Value::as_str),
Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old")
);
let mut rotated_mutation = refreshed_mutation;
rotated_mutation.proxy_metadata = Some(json!({
"tunnel_security": {
"mode": "non_tls_required",
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new"
}
}));
let rotated = repository
.register_node(&rotated_mutation)
.await
.expect("explicit security rotation should succeed");
assert_eq!(
rotated
.proxy_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted"))
.and_then(Value::as_str),
Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new")
);
}
#[tokio::test]
async fn registers_updates_config_and_unregisters_nodes() {
let repository = InMemoryProxyNodeRepository::default();
let registered = repository
.register_node(&ProxyNodeRegistrationMutation {
node_id: None,
name: "proxy-01".to_string(),
ip: "127.0.0.1".to_string(),
port: 0,
@@ -1131,6 +1466,7 @@ mod tests {
let updated = repository
.update_remote_config(&ProxyNodeRemoteConfigMutation {
node_id: registered.id.clone(),
expected_tunnel_generation: None,
node_name: Some("proxy-02".to_string()),
allowed_ports: Some(vec![443, 8443]),
log_level: Some("info".to_string()),
@@ -1160,6 +1496,7 @@ mod tests {
let after_upgrade = repository
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
node_id: registered.id.clone(),
expected_tunnel_generation: None,
heartbeat_interval: None,
active_connections: Some(2),
total_requests_delta: Some(1),
File diff suppressed because it is too large Load Diff
@@ -26,6 +26,7 @@ pub enum AdminSystemUsageAggregateImportMode {
Skip,
Overwrite,
Error,
ValidateError,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
@@ -4,7 +4,8 @@ use std::sync::RwLock;
use aether_ai_formats::UPSTREAM_IS_STREAM_KEY;
use aether_data_contracts::repository::usage::{
parse_usage_body_ref, usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary,
canonical_usage_body_ref_for, parse_usage_body_ref, sanitize_usage_request_metadata,
usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary,
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
@@ -31,7 +32,8 @@ use serde_json::Value;
use super::{
api_key_usage_contribution, provider_api_key_usage_contribution,
strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure,
sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence,
usage_can_recover_terminal_failure, usage_lifecycle_update_allowed,
usage_request_metadata_client_family, ApiKeyUsageContribution, ApiKeyUsageDelta,
ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest,
StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary,
@@ -159,10 +161,6 @@ fn usage_status_is_finalized(status: &str) -> bool {
matches!(status, "completed" | "failed" | "cancelled")
}
fn usage_status_is_lifecycle(status: &str) -> bool {
matches!(status, "pending" | "streaming")
}
fn merge_usage_timing(existing: Option<u64>, incoming: Option<u64>) -> Option<u64> {
match incoming {
Some(0) | None => existing.or(incoming),
@@ -1176,18 +1174,19 @@ impl UsageReadRepository for InMemoryUsageReadRepository {
}
async fn resolve_body_ref(&self, body_ref: &str) -> Result<Option<Value>, DataLayerError> {
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
return Ok(None);
};
let canonical_ref = usage_body_ref(&request_id, field);
if let Some(value) = self
.detached_bodies
.read()
.expect("usage repository lock")
.get(body_ref)
.get(&canonical_ref)
.cloned()
{
return Ok(Some(value));
}
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
return Ok(None);
};
let usage = self
.by_request_id
.read()
@@ -2676,44 +2675,74 @@ fn usage_body_ref_from_metadata(
.and_then(Value::as_object)
.and_then(|object| object.get(field.as_ref_key()))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(parse_usage_body_ref)
.filter(|(parsed_request_id, parsed_field)| {
parsed_request_id == request_id && *parsed_field == field
})
.map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field))
.and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field))
}
fn sanitize_memory_request_metadata(metadata: Option<Value>) -> Option<Value> {
sanitize_usage_request_metadata(metadata)
}
fn hydrate_legacy_body_refs(item: &mut StoredRequestUsageAudit) {
if item.request_body_ref.is_none() {
item.request_body_ref = usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::RequestBody,
);
}
if item.provider_request_body_ref.is_none() {
item.provider_request_body_ref = usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::ProviderRequestBody,
);
}
if item.response_body_ref.is_none() {
item.response_body_ref = usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::ResponseBody,
);
}
if item.client_response_body_ref.is_none() {
item.client_response_body_ref = usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::ClientResponseBody,
);
}
item.request_body_ref = item
.request_body_ref
.as_deref()
.and_then(|body_ref| {
canonical_usage_body_ref_for(body_ref, &item.request_id, UsageBodyField::RequestBody)
})
.or_else(|| {
usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::RequestBody,
)
});
item.provider_request_body_ref = item
.provider_request_body_ref
.as_deref()
.and_then(|body_ref| {
canonical_usage_body_ref_for(
body_ref,
&item.request_id,
UsageBodyField::ProviderRequestBody,
)
})
.or_else(|| {
usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::ProviderRequestBody,
)
});
item.response_body_ref = item
.response_body_ref
.as_deref()
.and_then(|body_ref| {
canonical_usage_body_ref_for(body_ref, &item.request_id, UsageBodyField::ResponseBody)
})
.or_else(|| {
usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::ResponseBody,
)
});
item.client_response_body_ref = item
.client_response_body_ref
.as_deref()
.and_then(|body_ref| {
canonical_usage_body_ref_for(
body_ref,
&item.request_id,
UsageBodyField::ClientResponseBody,
)
})
.or_else(|| {
usage_body_ref_from_metadata(
item.request_metadata.as_ref(),
&item.request_id,
UsageBodyField::ClientResponseBody,
)
});
}
fn hydrate_client_family(item: &mut StoredRequestUsageAudit) {
@@ -2723,34 +2752,6 @@ fn hydrate_client_family(item: &mut StoredRequestUsageAudit) {
}
}
fn persisted_usage_body_ref(
incoming_ref: Option<&str>,
incoming_body: Option<&Value>,
incoming_state: Option<UsageBodyCaptureState>,
_metadata: Option<&Value>,
existing: Option<&StoredRequestUsageAudit>,
field: UsageBodyField,
) -> Option<String> {
if incoming_state == Some(UsageBodyCaptureState::None) {
return None;
}
if incoming_body.is_some() {
return None;
}
incoming_ref
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| {
existing.and_then(|existing| match field {
UsageBodyField::RequestBody => existing.request_body_ref.clone(),
UsageBodyField::ProviderRequestBody => existing.provider_request_body_ref.clone(),
UsageBodyField::ResponseBody => existing.response_body_ref.clone(),
UsageBodyField::ClientResponseBody => existing.client_response_body_ref.clone(),
})
})
}
fn request_body_capture_replaces_derived_facts(
request_body: Option<&Value>,
request_body_state: Option<UsageBodyCaptureState>,
@@ -2852,9 +2853,68 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
usage: UpsertUsageRecord,
) -> Result<StoredRequestUsageAudit, DataLayerError> {
usage.validate()?;
let usage = strip_deprecated_usage_display_fields(usage);
let capture_usage = usage.clone();
let usage = sanitize_usage_for_persistence(usage);
let mut by_request_id = self.by_request_id.write().expect("usage repository lock");
let existing = by_request_id.get(&usage.request_id).cloned();
if let Some(existing) = existing.as_ref() {
if !usage_lifecycle_update_allowed(
&existing.status,
&existing.billing_status,
existing.updated_at_unix_secs,
existing.finalized_at_unix_secs,
&usage.status,
&usage.billing_status,
usage.updated_at_unix_secs,
usage.finalized_at_unix_secs,
) {
return Ok(existing.clone());
}
let can_recover = usage_can_recover_terminal_failure(
existing.status.as_str(),
existing.billing_status.as_str(),
usage.status.as_str(),
usage.billing_status.as_str(),
);
let completed_terminal_failure_recovery = existing.billing_status == "void"
&& matches!(existing.status.as_str(), "failed" | "cancelled")
&& usage.status == "completed";
if completed_terminal_failure_recovery && !can_recover {
return Ok(existing.clone());
}
}
let capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
if let Some(existing) = by_request_id.get_mut(&usage.request_id) {
existing.request_headers = None;
existing.request_body = None;
existing.request_body_ref = None;
existing.request_body_state = None;
existing.provider_request_headers = None;
existing.provider_request_body = None;
existing.provider_request_body_ref = None;
existing.provider_request_body_state = None;
existing.response_headers = None;
existing.response_body = None;
existing.response_body_ref = None;
existing.response_body_state = None;
existing.client_response_headers = None;
existing.client_response_body = None;
existing.client_response_body_ref = None;
existing.client_response_body_state = None;
existing.request_metadata =
sanitize_usage_request_metadata(existing.request_metadata.take());
}
{
let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock");
for field in [
UsageBodyField::RequestBody,
UsageBodyField::ProviderRequestBody,
UsageBodyField::ResponseBody,
UsageBodyField::ClientResponseBody,
] {
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
}
}
let created_at_unix_ms = by_request_id
.get(&usage.request_id)
@@ -2871,46 +2931,20 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
)
})
.unwrap_or_default();
if existing.as_ref().is_some_and(|existing| {
let can_recover = usage_can_recover_terminal_failure(
existing.status.as_str(),
existing.billing_status.as_str(),
usage.status.as_str(),
usage.billing_status.as_str(),
);
let finalized_lifecycle_regression = usage_status_is_finalized(&existing.status)
&& usage_status_is_lifecycle(&usage.status);
let completed_terminal_failure_recovery = existing.billing_status == "void"
&& matches!(existing.status.as_str(), "failed" | "cancelled")
&& usage.status == "completed";
(finalized_lifecycle_regression || completed_terminal_failure_recovery) && !can_recover
}) {
return Ok(existing.expect("existing usage should be present").clone());
}
if existing.as_ref().is_some_and(|existing| {
existing.billing_status == "pending"
&& existing.status == "streaming"
&& usage.status == "pending"
}) {
return Ok(existing.expect("existing usage should be present").clone());
}
let replace_client_request_body_facts = request_body_capture_replaces_derived_facts(
usage.request_body.as_ref(),
usage.request_body_state,
capture_usage.request_body.as_ref(),
capture_usage.request_body_state,
);
let replace_provider_request_body_facts = request_body_capture_replaces_derived_facts(
usage.provider_request_body.as_ref(),
usage.provider_request_body_state,
capture_usage.provider_request_body.as_ref(),
capture_usage.provider_request_body_state,
);
let clear_request_body = usage.request_body_state == Some(UsageBodyCaptureState::None);
let clear_request_body =
capture_usage.request_body_state == Some(UsageBodyCaptureState::None);
let clear_provider_request_body =
usage.provider_request_body_state == Some(UsageBodyCaptureState::None);
let clear_response_body = usage.response_body_state == Some(UsageBodyCaptureState::None);
let clear_client_response_body =
usage.client_response_body_state == Some(UsageBodyCaptureState::None);
capture_usage.provider_request_body_state == Some(UsageBodyCaptureState::None);
let replace_routing_snapshot = usage_status_is_finalized(&usage.status);
let mut incoming_request_metadata = usage.request_metadata.clone();
let mut incoming_request_metadata = capture_usage.request_metadata.clone();
if incoming_request_metadata.is_some()
&& (clear_request_body || clear_provider_request_body)
{
@@ -2942,61 +2976,11 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
.and_then(|existing| existing.request_metadata.clone())
}
});
let request_body_ref = persisted_usage_body_ref(
usage.request_body_ref.as_deref(),
usage.request_body.as_ref(),
usage.request_body_state,
request_metadata.as_ref(),
existing.as_ref(),
UsageBodyField::RequestBody,
);
let provider_request_body_ref = persisted_usage_body_ref(
usage.provider_request_body_ref.as_deref(),
usage.provider_request_body.as_ref(),
usage.provider_request_body_state,
request_metadata.as_ref(),
existing.as_ref(),
UsageBodyField::ProviderRequestBody,
);
let response_body_ref = persisted_usage_body_ref(
usage.response_body_ref.as_deref(),
usage.response_body.as_ref(),
usage.response_body_state,
request_metadata.as_ref(),
existing.as_ref(),
UsageBodyField::ResponseBody,
);
let client_response_body_ref = persisted_usage_body_ref(
usage.client_response_body_ref.as_deref(),
usage.client_response_body.as_ref(),
usage.client_response_body_state,
request_metadata.as_ref(),
existing.as_ref(),
UsageBodyField::ClientResponseBody,
);
if clear_request_body
|| clear_provider_request_body
|| clear_response_body
|| clear_client_response_body
{
let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock");
for (clear, field) in [
(clear_request_body, UsageBodyField::RequestBody),
(
clear_provider_request_body,
UsageBodyField::ProviderRequestBody,
),
(clear_response_body, UsageBodyField::ResponseBody),
(
clear_client_response_body,
UsageBodyField::ClientResponseBody,
),
] {
if clear {
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
}
}
}
let request_metadata = sanitize_memory_request_metadata(request_metadata);
let request_body_ref = None;
let provider_request_body_ref = None;
let response_body_ref = None;
let client_response_body_ref = None;
let stored = StoredRequestUsageAudit {
id: existing
.as_ref()
@@ -3105,159 +3089,97 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
),
status: usage.status,
billing_status: usage.billing_status,
request_headers: usage.request_headers.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.request_headers.clone())
}),
request_body: if clear_request_body {
None
} else {
usage.request_body.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.request_body.clone())
})
},
request_headers: None,
request_body: None,
request_body_ref,
request_body_state: usage.request_body_state.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.request_body_state)
}),
provider_request_headers: usage.provider_request_headers.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.provider_request_headers.clone())
}),
provider_request_body: if clear_provider_request_body {
None
} else {
usage.provider_request_body.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.provider_request_body.clone())
})
},
request_body_state: capture_usage.request_body_state,
provider_request_headers: None,
provider_request_body: None,
provider_request_body_ref,
provider_request_body_state: usage.provider_request_body_state.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.provider_request_body_state)
}),
response_headers: usage.response_headers.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.response_headers.clone())
}),
response_body: if clear_response_body {
None
} else {
usage.response_body.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.response_body.clone())
})
},
provider_request_body_state: capture_usage.provider_request_body_state,
response_headers: None,
response_body: None,
response_body_ref,
response_body_state: usage.response_body_state.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.response_body_state)
}),
client_response_headers: usage.client_response_headers.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.client_response_headers.clone())
}),
client_response_body: if clear_client_response_body {
None
} else {
usage.client_response_body.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.client_response_body.clone())
})
},
response_body_state: capture_usage.response_body_state,
client_response_headers: None,
client_response_body: None,
client_response_body_ref,
client_response_body_state: usage.client_response_body_state.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.client_response_body_state)
}),
client_response_body_state: capture_usage.client_response_body_state,
candidate_id: if replace_routing_snapshot {
usage.candidate_id
capture_usage.candidate_id
} else {
usage.candidate_id.or_else(|| {
capture_usage.candidate_id.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.routing_candidate_id().map(ToOwned::to_owned))
})
},
candidate_index: if replace_routing_snapshot {
usage.candidate_index
capture_usage.candidate_index
} else {
usage.candidate_index.or_else(|| {
capture_usage.candidate_index.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.routing_candidate_index())
})
},
key_name: if replace_routing_snapshot {
usage.key_name
capture_usage.key_name
} else {
usage.key_name.or_else(|| {
capture_usage.key_name.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.routing_key_name().map(ToOwned::to_owned))
})
},
planner_kind: if replace_routing_snapshot {
usage.planner_kind
capture_usage.planner_kind
} else {
usage.planner_kind.or_else(|| {
capture_usage.planner_kind.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.routing_planner_kind().map(ToOwned::to_owned))
})
},
route_family: if replace_routing_snapshot {
usage.route_family
capture_usage.route_family
} else {
usage.route_family.or_else(|| {
capture_usage.route_family.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.routing_route_family().map(ToOwned::to_owned))
})
},
route_kind: if replace_routing_snapshot {
usage.route_kind
capture_usage.route_kind
} else {
usage.route_kind.or_else(|| {
capture_usage.route_kind.or_else(|| {
existing
.as_ref()
.and_then(|existing| existing.routing_route_kind().map(ToOwned::to_owned))
})
},
execution_path: if replace_routing_snapshot {
usage.execution_path
capture_usage.execution_path
} else {
usage.execution_path.or_else(|| {
capture_usage.execution_path.or_else(|| {
existing.as_ref().and_then(|existing| {
existing.routing_execution_path().map(ToOwned::to_owned)
})
})
},
local_execution_runtime_miss_reason: if replace_routing_snapshot {
usage.local_execution_runtime_miss_reason
capture_usage.local_execution_runtime_miss_reason
} else {
usage.local_execution_runtime_miss_reason.or_else(|| {
existing.as_ref().and_then(|existing| {
existing
.routing_local_execution_runtime_miss_reason()
.map(ToOwned::to_owned)
capture_usage
.local_execution_runtime_miss_reason
.or_else(|| {
existing.as_ref().and_then(|existing| {
existing
.routing_local_execution_runtime_miss_reason()
.map(ToOwned::to_owned)
})
})
})
},
client_family: usage_request_metadata_client_family(request_metadata.as_ref())
.map(ToOwned::to_owned)
@@ -788,6 +788,64 @@ async fn upsert_allows_completed_recovery_after_void_failure() {
assert_eq!(stored.total_tokens, 10);
}
#[tokio::test]
async fn stale_terminal_event_cannot_replace_usage_routing_or_counter_contribution() {
let auth_api_keys = sample_auth_api_key_repository(&["api-key-1"]);
let repository = InMemoryUsageReadRepository::default()
.with_auth_api_key_repository(Arc::clone(&auth_api_keys));
let mut newer = sample_upsert_usage_record("req-stale-terminal");
newer.api_key_id = Some("api-key-1".to_string());
newer.status = "completed".to_string();
newer.status_code = Some(200);
newer.total_tokens = Some(5);
newer.total_cost_usd = Some(0.5);
newer.candidate_id = Some("candidate-new".to_string());
newer.route_kind = Some("route-new".to_string());
newer.updated_at_unix_secs = 200;
newer.finalized_at_unix_secs = Some(200);
repository
.upsert(newer)
.await
.expect("newer terminal usage should upsert");
let mut stale = sample_upsert_usage_record("req-stale-terminal");
stale.api_key_id = Some("api-key-1".to_string());
stale.status = "failed".to_string();
stale.billing_status = "void".to_string();
stale.status_code = Some(503);
stale.total_tokens = Some(999);
stale.total_cost_usd = Some(99.0);
stale.candidate_id = Some("candidate-stale".to_string());
stale.route_kind = Some("route-stale".to_string());
stale.updated_at_unix_secs = 199;
stale.finalized_at_unix_secs = Some(199);
let stored = repository
.upsert(stale)
.await
.expect("stale terminal usage should be ignored");
assert_eq!(stored.status, "completed");
assert_eq!(stored.billing_status, "pending");
assert_eq!(stored.status_code, Some(200));
assert_eq!(stored.total_tokens, 5);
assert_eq!(stored.total_cost_usd, 0.5);
assert_eq!(stored.routing_candidate_id(), Some("candidate-new"));
assert_eq!(stored.routing_route_kind(), Some("route-new"));
assert_eq!(stored.updated_at_unix_secs, 200);
let key = auth_api_keys
.list_export_api_keys_by_ids(&["api-key-1".to_string()])
.await
.expect("api key stats should load")
.into_iter()
.next()
.expect("api key should exist");
assert_eq!(key.total_requests, 1);
assert_eq!(key.total_tokens, 5);
assert_eq!(key.total_cost_usd, 0.5);
}
#[tokio::test]
async fn upsert_rejects_non_authoritative_void_failure_recovery() {
let repository = InMemoryUsageReadRepository::default();
@@ -1189,6 +1247,23 @@ async fn detached_body_seed_moves_large_payloads_behind_usage_refs() {
);
}
#[tokio::test]
async fn seed_discards_cross_request_and_cross_field_body_refs() {
let mut usage = sample_usage("req-ref-target", 100);
usage.request_body_ref = Some("usage://request/req-ref-owner/request_body".to_string());
usage.response_body_ref = Some("usage://request/req-ref-target/request_body".to_string());
let repository = InMemoryUsageReadRepository::seed(vec![usage]);
let stored = repository
.find_by_request_id("req-ref-target")
.await
.expect("find should succeed")
.expect("usage should exist");
assert!(stored.request_body_ref.is_none());
assert!(stored.response_body_ref.is_none());
}
#[tokio::test]
async fn upsert_writes_usage_record() {
let repository = InMemoryUsageReadRepository::default();
@@ -1520,12 +1595,7 @@ async fn upsert_does_not_backfill_typed_body_refs_from_request_metadata() {
.expect("upsert should succeed");
assert_eq!(stored.request_body_ref, None);
assert_eq!(
stored.request_metadata,
Some(json!({
"request_body_ref": "usage://request/req-upsert-body-ref-metadata/request_body"
}))
);
assert_eq!(stored.request_metadata, None);
}
#[tokio::test]
@@ -6,33 +6,35 @@ mod mysql;
pub(crate) use aether_data_contracts::repository::usage::{
api_key_usage_contribution, incoming_usage_can_recover_terminal_failure,
model_usage_contribution, provider_api_key_usage_contribution, provider_api_key_usage_is_error,
provider_api_key_usage_is_success, strip_deprecated_usage_display_fields,
usage_can_recover_terminal_failure, usage_request_metadata_client_family, ApiKeyLastUsedDelta,
ApiKeyUsageContribution, ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution,
ModelUsageDelta, PendingUsageCleanupSummary, ProviderApiKeyUsageContribution,
ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta,
StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary,
StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit,
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow,
StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow,
StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDailySummary,
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow,
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
StoredUsageProviderPerformance, StoredUsageProviderPerformanceProviderRow,
StoredUsageProviderPerformanceSummary, StoredUsageProviderPerformanceTimelineRow,
StoredUsageSettledCostSummary, StoredUsageTimeSeriesBucket, StoredUsageUserTotals,
UpsertUsageRecord, UsageAuditAggregationGroupBy, UsageAuditAggregationQuery,
UsageAuditKeywordSearchQuery, UsageAuditListQuery, UsageAuditSummaryQuery,
UsageBreakdownGroupBy, UsageBreakdownSummaryQuery, UsageCacheAffinityHitSummaryQuery,
UsageCacheAffinityIntervalGroupBy, UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery,
UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupWindow,
UsageCostSavingsSummaryQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot,
UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery,
UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery,
UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery,
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
provider_api_key_usage_is_success, sanitize_usage_capture_controls_for_persistence,
sanitize_usage_for_persistence, strip_deprecated_usage_display_fields,
usage_can_recover_terminal_failure, usage_lifecycle_update_allowed,
usage_request_metadata_client_family, ApiKeyLastUsedDelta, ApiKeyUsageContribution,
ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution, ModelUsageDelta,
PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta,
ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary,
StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow,
StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary,
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
StoredUsageDashboardProviderCount, StoredUsageDashboardStatsSummary,
StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary,
StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance,
StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary,
StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary,
StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UpsertUsageRecord,
UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery,
UsageAuditListQuery, UsageAuditSummaryQuery, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
UsageCacheAffinityHitSummaryQuery, UsageCacheAffinityIntervalGroupBy,
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCleanupPreviewCounts,
UsageCleanupSummary, UsageCleanupWindow, UsageCostSavingsSummaryQuery,
UsageCounterFlushSummary, UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot,
UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository,
UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery, UsageWriteRepository,
};
#[cfg(feature = "postgres")]
File diff suppressed because it is too large Load Diff
@@ -1,11 +1,15 @@
mod memory;
pub use aether_data_contracts::repository::users::{
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
is_last_active_admin_delete_denied, is_last_active_admin_update_denied, is_valid_bcrypt_hash,
last_oauth_unbind_denial, normalize_user_group_name, BindUserOAuthLinkOutcome,
BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome,
LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy,
UserExportSortOrder, UserExportSummary, UserReadRepository,
UserExportSortOrder, UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED,
LAST_ACTIVE_ADMIN_UPDATE_DENIED,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlUserReadRepository;
@@ -14,6 +14,7 @@ use crate::DataLayerError;
struct MemoryVideoTaskIndex {
by_id: BTreeMap<String, StoredVideoTask>,
short_to_id: BTreeMap<String, String>,
request_to_id: BTreeMap<String, String>,
user_external_to_id: BTreeMap<(String, String), String>,
}
@@ -28,6 +29,7 @@ impl InMemoryVideoTaskRepository {
if let Some(short_id) = previous.short_id {
index.short_to_id.remove(&short_id);
}
index.request_to_id.remove(&previous.request_id);
if let (Some(user_id), Some(external_task_id)) =
(previous.user_id, previous.external_task_id)
{
@@ -40,6 +42,9 @@ impl InMemoryVideoTaskRepository {
if let Some(short_id) = &task.short_id {
index.short_to_id.insert(short_id.clone(), task.id.clone());
}
index
.request_to_id
.insert(task.request_id.clone(), task.id.clone());
if let (Some(user_id), Some(external_task_id)) = (&task.user_id, &task.external_task_id) {
index
.user_external_to_id
@@ -49,6 +54,35 @@ impl InMemoryVideoTaskRepository {
task
}
fn ensure_unique_keys_available(
index: &MemoryVideoTaskIndex,
task: &UpsertVideoTask,
) -> Result<(), DataLayerError> {
if let Some(short_id) = task.short_id.as_deref() {
if index
.short_to_id
.get(short_id)
.is_some_and(|existing_id| existing_id != &task.id)
{
return Err(DataLayerError::InvalidInput(format!(
"video task {} conflicts with existing short_id {short_id}",
task.id
)));
}
}
if index
.request_to_id
.get(&task.request_id)
.is_some_and(|existing_id| existing_id != &task.id)
{
return Err(DataLayerError::InvalidInput(format!(
"video task {} conflicts with existing request_id {}",
task.id, task.request_id
)));
}
Ok(())
}
fn matches_filter(task: &StoredVideoTask, filter: &VideoTaskQueryFilter) -> bool {
if let Some(user_id) = filter.user_id.as_deref() {
if task.user_id.as_deref() != Some(user_id) {
@@ -103,6 +137,36 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository {
})
}
async fn find_for_user(
&self,
key: VideoTaskLookupKey<'_>,
user_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let index = self.index.read().expect("video task repository lock");
let task = match key {
VideoTaskLookupKey::Id(id) => index.by_id.get(id),
VideoTaskLookupKey::ShortId(short_id) => index
.short_to_id
.get(short_id)
.and_then(|id| index.by_id.get(id)),
VideoTaskLookupKey::UserExternal {
user_id: lookup_user_id,
external_task_id,
} => {
if lookup_user_id != user_id {
return Ok(None);
}
index
.user_external_to_id
.get(&(lookup_user_id.to_string(), external_task_id.to_string()))
.and_then(|id| index.by_id.get(id))
}
};
Ok(task
.filter(|task| task.user_id.as_deref() == Some(user_id))
.cloned())
}
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
@@ -306,22 +370,34 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository {
#[async_trait]
impl VideoTaskWriteRepository for InMemoryVideoTaskRepository {
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
async fn upsert(&self, mut task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
let mut index = self.index.write().expect("video task repository lock");
Self::ensure_unique_keys_available(&index, &task)?;
if let Some(existing) = index.by_id.get(&task.id) {
existing.ensure_immutable_identity_matches(&task)?;
task.created_at_unix_ms = existing.created_at_unix_ms;
}
Ok(Self::store_locked(&mut index, task.into_stored()))
}
async fn update_if_active(
&self,
task: UpsertVideoTask,
mut task: UpsertVideoTask,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let mut index = self.index.write().expect("video task repository lock");
if Self::ensure_unique_keys_available(&index, &task).is_err() {
return Ok(None);
}
let Some(existing) = index.by_id.get(&task.id) else {
return Ok(None);
};
if !existing.status.is_active() {
return Ok(None);
}
if existing.ensure_immutable_identity_matches(&task).is_err() {
return Ok(None);
}
task.created_at_unix_ms = existing.created_at_unix_ms;
Ok(Some(Self::store_locked(&mut index, task.into_stored())))
}
@@ -455,6 +531,34 @@ mod tests {
.is_some());
}
#[tokio::test]
async fn owner_scoped_lookup_rejects_foreign_user_for_every_identifier() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("upsert should succeed");
for key in [
VideoTaskLookupKey::Id("task-1"),
VideoTaskLookupKey::ShortId("short-task-1"),
VideoTaskLookupKey::UserExternal {
user_id: "user-1",
external_task_id: "ext-task-1",
},
] {
assert!(repo
.find_for_user(key, "user-1")
.await
.expect("owner lookup should succeed")
.is_some());
assert!(repo
.find_for_user(key, "user-2")
.await
.expect("foreign lookup should succeed")
.is_none());
}
}
#[tokio::test]
async fn list_active_only_returns_active_tasks_in_descending_update_order() {
let repo = InMemoryVideoTaskRepository::default();
@@ -478,59 +582,69 @@ mod tests {
}
#[tokio::test]
async fn upsert_replaces_secondary_indexes() {
async fn upsert_rejects_immutable_identity_replacement() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("upsert should succeed");
repo.upsert(UpsertVideoTask {
id: "task-1".to_string(),
short_id: Some("short-task-1b".to_string()),
request_id: "request-task-1b".to_string(),
user_id: Some("user-2".to_string()),
api_key_id: Some("api-key-2".to_string()),
username: Some("user-2".to_string()),
api_key_name: Some("secondary".to_string()),
external_task_id: Some("ext-task-1b".to_string()),
provider_id: Some("provider-2".to_string()),
endpoint_id: Some("endpoint-2".to_string()),
key_id: Some("provider-key-2".to_string()),
client_api_format: Some("gemini:video".to_string()),
provider_api_format: Some("gemini:video".to_string()),
format_converted: false,
model: Some("veo-3".to_string()),
prompt: Some("remix".to_string()),
original_request_body: Some(serde_json::json!({"prompt": "remix"})),
duration_seconds: Some(8),
resolution: Some("1080p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("720p".to_string()),
status: VideoTaskStatus::Processing,
progress_percent: 50,
progress_message: Some("processing".to_string()),
retry_count: 1,
poll_interval_seconds: 10,
next_poll_at_unix_secs: Some(200),
poll_count: 2,
max_poll_count: 360,
created_at_unix_ms: 150,
submitted_at_unix_secs: Some(150),
completed_at_unix_secs: None,
updated_at_unix_secs: 200,
error_code: None,
error_message: None,
video_url: None,
request_metadata: None,
})
.await
.expect("upsert should succeed");
let conflict = repo
.upsert(UpsertVideoTask {
id: "task-1".to_string(),
short_id: Some("short-task-1b".to_string()),
request_id: "request-task-1b".to_string(),
user_id: Some("user-2".to_string()),
api_key_id: Some("api-key-2".to_string()),
username: Some("user-2".to_string()),
api_key_name: Some("secondary".to_string()),
external_task_id: Some("ext-task-1b".to_string()),
provider_id: Some("provider-2".to_string()),
endpoint_id: Some("endpoint-2".to_string()),
key_id: Some("provider-key-2".to_string()),
client_api_format: Some("gemini:video".to_string()),
provider_api_format: Some("gemini:video".to_string()),
format_converted: false,
model: Some("veo-3".to_string()),
prompt: Some("remix".to_string()),
original_request_body: Some(serde_json::json!({"prompt": "remix"})),
duration_seconds: Some(8),
resolution: Some("1080p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("720p".to_string()),
status: VideoTaskStatus::Processing,
progress_percent: 50,
progress_message: Some("processing".to_string()),
retry_count: 1,
poll_interval_seconds: 10,
next_poll_at_unix_secs: Some(200),
poll_count: 2,
max_poll_count: 360,
created_at_unix_ms: 150,
submitted_at_unix_secs: Some(150),
completed_at_unix_secs: None,
updated_at_unix_secs: 200,
error_code: None,
error_message: None,
video_url: None,
request_metadata: None,
})
.await
.expect_err("identity replacement should be rejected");
assert!(conflict.to_string().contains("immutable field short_id"));
let stored = repo
.find(VideoTaskLookupKey::Id("task-1"))
.await
.expect("find should succeed")
.expect("original task should remain");
assert_eq!(stored.request_id, "request-task-1");
assert_eq!(stored.user_id.as_deref(), Some("user-1"));
assert_eq!(stored.status, VideoTaskStatus::Submitted);
assert!(repo
.find(VideoTaskLookupKey::ShortId("short-task-1"))
.await
.expect("find should succeed")
.is_none());
.is_some());
assert!(repo
.find(VideoTaskLookupKey::UserExternal {
user_id: "user-1",
@@ -538,12 +652,117 @@ mod tests {
})
.await
.expect("find should succeed")
.is_none());
.is_some());
assert!(repo
.find(VideoTaskLookupKey::ShortId("short-task-1b"))
.await
.expect("find should succeed")
.is_none());
}
#[tokio::test]
async fn upsert_allows_same_identity_status_update() {
let repo = InMemoryVideoTaskRepository::default();
let task = sample_task("task-1", VideoTaskStatus::Submitted, 100);
repo.upsert(task.clone())
.await
.expect("initial upsert should succeed");
let updated = repo
.upsert(UpsertVideoTask {
status: VideoTaskStatus::Processing,
progress_percent: 50,
poll_count: 2,
created_at_unix_ms: 999,
updated_at_unix_secs: 200,
..task
})
.await
.expect("same identity update should succeed");
assert_eq!(updated.status, VideoTaskStatus::Processing);
assert_eq!(updated.progress_percent, 50);
assert_eq!(updated.poll_count, 2);
assert_eq!(updated.created_at_unix_ms, 90);
assert_eq!(updated.updated_at_unix_secs, 200);
}
#[tokio::test]
async fn update_if_active_rejects_identity_conflict_without_modification() {
let repo = InMemoryVideoTaskRepository::default();
let task = sample_task("task-1", VideoTaskStatus::Submitted, 100);
repo.upsert(task.clone())
.await
.expect("initial upsert should succeed");
let result = repo
.update_if_active(UpsertVideoTask {
user_id: Some("attacker".to_string()),
status: VideoTaskStatus::Completed,
progress_percent: 100,
updated_at_unix_secs: 200,
..task
})
.await
.expect("guarded update should execute");
assert!(result.is_none());
let stored = repo
.find(VideoTaskLookupKey::Id("task-1"))
.await
.expect("find should succeed")
.expect("original task should remain");
assert_eq!(stored.user_id.as_deref(), Some("user-1"));
assert_eq!(stored.status, VideoTaskStatus::Submitted);
assert_eq!(stored.progress_percent, 0);
assert_eq!(stored.updated_at_unix_secs, 100);
}
#[tokio::test]
async fn upsert_rejects_secondary_unique_key_takeover() {
let repo = InMemoryVideoTaskRepository::default();
let original = sample_task("task-1", VideoTaskStatus::Submitted, 100);
repo.upsert(original.clone())
.await
.expect("initial upsert should succeed");
let short_id_conflict = repo
.upsert(UpsertVideoTask {
id: "task-2".to_string(),
request_id: "request-task-2".to_string(),
..original.clone()
})
.await
.expect_err("a short id must not be reassigned to another task");
assert!(short_id_conflict.to_string().contains("existing short_id"));
let request_id_conflict = repo
.upsert(UpsertVideoTask {
id: "task-3".to_string(),
short_id: Some("short-task-3".to_string()),
..original
})
.await
.expect_err("a request id must not be reassigned to another task");
assert!(request_id_conflict
.to_string()
.contains("existing request_id"));
assert!(repo
.find(VideoTaskLookupKey::ShortId("short-task-1"))
.await
.expect("find should succeed")
.is_some());
assert!(repo
.find(VideoTaskLookupKey::Id("task-2"))
.await
.expect("find should succeed")
.is_none());
assert!(repo
.find(VideoTaskLookupKey::Id("task-3"))
.await
.expect("find should succeed")
.is_none());
}
#[tokio::test]
File diff suppressed because it is too large Load Diff
@@ -1,19 +1,34 @@
mod memory;
pub use aether_data_contracts::repository::wallet::{
redeem_code_credits_recharge_balance, redeem_code_payment_method,
redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentCallbackRecord,
AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery,
AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletPaymentOrderRecord,
AdminWalletRefundRecord, AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord,
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput,
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext,
CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput,
canonicalize_payment_method, canonicalize_wallet_refund_fields,
payment_order_is_uncertain_wallet_checkout_placeholder,
payment_order_refund_amounts_are_consistent,
payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response,
project_wallet_recharge_gateway_response, redeem_code_credits_recharge_balance,
redeem_code_payment_method, redeem_code_refundable_amount, stored_timestamp_unix_secs,
validate_admin_redeem_code_batch_input, validate_payment_order_credit_amounts,
validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements,
validate_redeem_wallet_credit, validate_wallet_recharge_order_input,
wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token,
wallet_recharge_checkout_claimed_at, wallet_recharge_checkout_failed_response,
wallet_recharge_checkout_uncertain_response, wallet_recharge_order_created_at_unix_secs,
wallet_recharge_order_is_checkout_placeholder,
wallet_recharge_order_is_reclaimable_placeholder, wallet_recharge_replay_matches,
wallet_recharge_response_is_checkout_placeholder, wallet_refund_proof_is_success,
AdjustWalletBalanceInput, AdminPaymentCallbackRecord, AdminPaymentOrderListQuery,
AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery,
AdminWalletListQuery, AdminWalletPaymentOrderRecord, AdminWalletRefundRecord,
AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord, CanonicalWalletRefundFields,
CompareAndSwapPaymentOrderStripeClientSecretInput, CompleteAdminWalletRefundInput,
CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult,
CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome,
CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome,
CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome,
CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput,
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput,
ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput,
RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback,
StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage,
StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage,
@@ -22,9 +37,10 @@ pub use aether_data_contracts::repository::wallet::{
StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput,
UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome,
WalletReadRepository, WalletReadSeed, WalletReadSnapshot, WalletRepository,
WalletWriteRepository,
WalletWriteRepository, WALLET_RECHARGE_CHECKOUT_CLAIM_LEASE_SECS,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlWalletReadRepository;