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
+555 -109
View File
@@ -3,9 +3,9 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use aether_data_contracts::repository::auth::{
AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository,
AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord,
StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
AuthApiKeyWriteRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord,
CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord,
StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
};
use aether_data_contracts::DataLayerError;
@@ -69,6 +69,21 @@ SELECT
FROM api_keys
"#;
const MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL: &[&str] = &[
"UPDATE request_candidates SET api_key_name = NULL WHERE api_key_id = ?",
"UPDATE video_tasks SET api_key_name = NULL WHERE api_key_id = ?",
"UPDATE `usage` SET api_key_name = NULL WHERE api_key_id = ?",
"UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id = ?",
"UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id = ?",
"UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)",
"UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))",
"UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)",
"UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)",
];
const MYSQL_DELETE_API_KEY_DEPENDENTS_SQL: &[&str] =
&["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"];
#[derive(Debug, Clone)]
pub struct MysqlAuthApiKeyReadRepository {
pool: MysqlPool,
@@ -111,6 +126,17 @@ impl MysqlAuthApiKeyReadRepository {
record: CreateApiKeyInsertRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let now = current_unix_secs();
let mut tx = self.pool.begin().await.map_sql_err()?;
let owner_exists: Option<String> =
sqlx::query_scalar("SELECT id FROM users WHERE id = ? AND is_deleted = 0 FOR UPDATE")
.bind(&record.user_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
if owner_exists.is_none() {
tx.rollback().await.map_sql_err()?;
return Ok(None);
}
sqlx::query(
r#"
INSERT INTO api_keys (
@@ -120,7 +146,7 @@ INSERT INTO api_keys (
total_requests, total_tokens, total_cost_usd, is_standalone,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.api_key_id)
@@ -150,6 +176,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
&record.force_capabilities,
"api_keys.force_capabilities",
)?)
.bind(optional_json_to_string(
&record.feature_settings,
"api_keys.feature_settings",
)?)
.bind(record.is_active)
.bind(optional_i64_from_u64(
record.expires_at_unix_secs,
@@ -165,11 +195,25 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
.bind(record.is_standalone)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
self.reload_export_by_id(&record.api_key_id).await
let reload_sql = format!("{EXPORT_COLUMNS}\nWHERE api_keys.id = ?\nLIMIT 1");
let row = sqlx::query(&reload_sql)
.bind(&record.api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(row) = row else {
tx.rollback().await.map_sql_err()?;
return Err(DataLayerError::UnexpectedValue(format!(
"created api_keys row is missing: {}",
record.api_key_id
)));
};
let created = map_auth_api_key_export_row(&row)?;
tx.commit().await.map_sql_err()?;
Ok(Some(created))
}
}
@@ -186,6 +230,7 @@ struct CreateApiKeyInsertRecord {
rate_limit: Option<i32>,
concurrent_limit: Option<i32>,
force_capabilities: Option<serde_json::Value>,
feature_settings: Option<serde_json::Value>,
is_active: bool,
expires_at_unix_secs: Option<u64>,
auto_delete_on_expiry: bool,
@@ -337,6 +382,7 @@ impl AuthApiKeyReadRepository for MysqlAuthApiKeyReadRepository {
if user_ids.is_empty() {
return Ok(AuthApiKeyExportSummary::default());
}
let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?;
let mut builder = QueryBuilder::<MySql>::new(
r#"
@@ -345,7 +391,7 @@ SELECT
SUM(CASE WHEN is_active = 1 AND (expires_at IS NULL OR expires_at >=
"#,
);
builder.push_bind(now_unix_secs as i64);
builder.push_bind(now_unix_secs);
builder.push(
r#") THEN 1 ELSE 0 END) AS active
FROM api_keys
@@ -429,6 +475,7 @@ WHERE id = ?
rate_limit: Some(record.rate_limit),
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
feature_settings: record.feature_settings,
is_active: record.is_active,
expires_at_unix_secs: record.expires_at_unix_secs,
auto_delete_on_expiry: record.auto_delete_on_expiry,
@@ -457,6 +504,7 @@ WHERE id = ?
rate_limit: record.rate_limit,
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
feature_settings: None,
is_active: record.is_active,
expires_at_unix_secs: record.expires_at_unix_secs,
auto_delete_on_expiry: record.auto_delete_on_expiry,
@@ -472,35 +520,41 @@ WHERE id = ?
&self,
record: UpdateUserApiKeyBasicRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let now = current_unix_secs() as i64;
sqlx::query(
self.update_user_api_key_basic_scoped(record, false).await
}
async fn compare_and_swap_api_key_ciphertext(
&self,
mutation: &CompareAndSwapAuthApiKeyCiphertext,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE api_keys
SET name = COALESCE(?, name),
rate_limit = COALESCE(?, rate_limit),
concurrent_limit = COALESCE(?, concurrent_limit),
ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END,
updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
SET key_encrypted = ?
WHERE BINARY id = BINARY ?
AND BINARY user_id = BINARY ?
AND BINARY key_hash = BINARY ?
AND is_standalone = ?
AND BINARY key_encrypted = BINARY ?
"#,
)
.bind(record.name.as_deref())
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.ip_rules.is_some())
.bind(json_string_from_nested_string_list(
&record.ip_rules,
"api_keys.ip_rules",
)?)
.bind(now)
.bind(&record.api_key_id)
.bind(&record.user_id)
.bind(&mutation.key_encrypted)
.bind(&mutation.api_key_id)
.bind(&mutation.user_id)
.bind(&mutation.key_hash)
.bind(mutation.is_standalone)
.bind(&mutation.expected_key_encrypted)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(&record.api_key_id).await
Ok(result.rows_affected() == 1)
}
async fn update_user_api_key_basic_if_unlocked(
&self,
record: UpdateUserApiKeyBasicRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.update_user_api_key_basic_scoped(record, true).await
}
async fn update_standalone_api_key_basic(
@@ -511,7 +565,9 @@ WHERE id = ?
sqlx::query(
r#"
UPDATE api_keys
SET name = COALESCE(?, name),
SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END,
name = CASE WHEN ? THEN ? ELSE name END,
force_capabilities = CASE WHEN ? THEN ? ELSE force_capabilities END,
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
@@ -525,7 +581,15 @@ WHERE id = ?
AND is_standalone = 1
"#,
)
.bind(record.key_encrypted_present)
.bind(record.key_encrypted.as_deref())
.bind(record.name_present)
.bind(record.name.as_deref())
.bind(record.force_capabilities.is_some())
.bind(optional_json_to_string(
&record.force_capabilities.clone().flatten(),
"api_keys.force_capabilities",
)?)
.bind(record.rate_limit_present)
.bind(record.rate_limit)
.bind(record.concurrent_limit_present)
@@ -565,13 +629,143 @@ WHERE id = ?
self.reload_export_by_id(&record.api_key_id).await
}
async fn restore_api_key_if_matches(
&self,
expected: &StoredAuthApiKeyExportRecord,
restored: &StoredAuthApiKeyExportRecord,
) -> Result<bool, DataLayerError> {
if restored.api_key_id != expected.api_key_id
|| restored.user_id != expected.user_id
|| restored.key_hash != expected.key_hash
|| restored.is_standalone != expected.is_standalone
{
return Ok(false);
}
let mut tx = self.pool.begin().await.map_sql_err()?;
let select_sql = format!("{EXPORT_COLUMNS} WHERE api_keys.id = ? LIMIT 1 FOR UPDATE");
let row = sqlx::query(&select_sql)
.bind(&expected.api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(row) = row else {
tx.rollback().await.map_sql_err()?;
return Ok(false);
};
let current = map_auth_api_key_export_row(&row)?;
if current != *expected {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
let result = sqlx::query(
r#"
UPDATE api_keys
SET key_encrypted = ?,
name = ?,
allowed_providers = ?,
allowed_api_formats = ?,
allowed_models = ?,
ip_rules = ?,
rate_limit = ?,
concurrent_limit = ?,
force_capabilities = ?,
feature_settings = ?,
is_active = ?,
expires_at = ?,
auto_delete_on_expiry = ?,
total_requests = ?,
total_tokens = ?,
total_cost_usd = ?,
last_used_at = ?,
updated_at = ?
WHERE id = ?
AND user_id = ?
AND key_hash = ?
AND is_standalone = ?
"#,
)
.bind(restored.key_encrypted.as_deref())
.bind(restored.name.as_deref())
.bind(json_string_from_string_list(
restored.allowed_providers.as_ref(),
"api_keys.allowed_providers",
)?)
.bind(json_string_from_string_list(
restored.allowed_api_formats.as_ref(),
"api_keys.allowed_api_formats",
)?)
.bind(json_string_from_string_list(
restored.allowed_models.as_ref(),
"api_keys.allowed_models",
)?)
.bind(json_string_from_string_list(
restored.ip_rules.as_ref(),
"api_keys.ip_rules",
)?)
.bind(restored.rate_limit)
.bind(restored.concurrent_limit)
.bind(optional_json_to_string(
&restored.force_capabilities,
"api_keys.force_capabilities",
)?)
.bind(optional_json_to_string(
&restored.feature_settings,
"api_keys.feature_settings",
)?)
.bind(restored.is_active)
.bind(optional_i64_from_u64(
restored.expires_at_unix_secs,
"api_keys.expires_at",
)?)
.bind(restored.auto_delete_on_expiry)
.bind(i64_from_u64(
restored.total_requests,
"api_keys.total_requests",
)?)
.bind(i64_from_u64(
restored.total_tokens,
"api_keys.total_tokens",
)?)
.bind(restored.total_cost_usd)
.bind(optional_i64_from_u64(
restored.last_used_at_unix_secs,
"api_keys.last_used_at",
)?)
.bind(current_unix_secs() as i64)
.bind(&restored.api_key_id)
.bind(&restored.user_id)
.bind(&restored.key_hash)
.bind(restored.is_standalone)
.execute(&mut *tx)
.await
.map_sql_err()?;
if result.rows_affected() != 1 {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
tx.commit().await.map_sql_err()?;
Ok(true)
}
async fn set_user_api_key_active(
&self,
user_id: &str,
api_key_id: &str,
is_active: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_active(api_key_id, Some(user_id), is_active, false)
self.set_active(api_key_id, Some(user_id), is_active, false, false)
.await
}
async fn set_user_api_key_active_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
is_active: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_active(api_key_id, Some(user_id), is_active, false, true)
.await
}
@@ -580,7 +774,8 @@ WHERE id = ?
api_key_id: &str,
is_active: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_active(api_key_id, None, is_active, true).await
self.set_active(api_key_id, None, is_active, true, false)
.await
}
async fn set_user_api_key_locked(
@@ -615,26 +810,23 @@ WHERE id = ?
api_key_id: &str,
allowed_providers: Option<Vec<String>>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
sqlx::query(
r#"
UPDATE api_keys
SET allowed_providers = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
self.set_user_api_key_allowed_providers_scoped(
user_id,
api_key_id,
allowed_providers,
false,
)
.bind(json_string_from_string_list(
allowed_providers.as_ref(),
"api_keys.allowed_providers",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(api_key_id).await
}
async fn set_user_api_key_allowed_providers_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
allowed_providers: Option<Vec<String>>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_user_api_key_allowed_providers_scoped(user_id, api_key_id, allowed_providers, true)
.await
}
async fn set_user_api_key_force_capabilities(
@@ -643,26 +835,28 @@ WHERE id = ?
api_key_id: &str,
force_capabilities: Option<serde_json::Value>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
sqlx::query(
r#"
UPDATE api_keys
SET force_capabilities = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
self.set_user_api_key_force_capabilities_scoped(
user_id,
api_key_id,
force_capabilities,
false,
)
.await
}
async fn set_user_api_key_force_capabilities_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
force_capabilities: Option<serde_json::Value>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_user_api_key_force_capabilities_scoped(
user_id,
api_key_id,
force_capabilities,
true,
)
.bind(optional_json_to_string(
&force_capabilities,
"api_keys.force_capabilities",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(api_key_id).await
}
async fn set_user_api_key_feature_settings(
@@ -671,26 +865,18 @@ WHERE id = ?
api_key_id: &str,
feature_settings: Option<serde_json::Value>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
sqlx::query(
r#"
UPDATE api_keys
SET feature_settings = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
)
.bind(optional_json_to_string(
&feature_settings,
"api_keys.feature_settings",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(api_key_id).await
self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, false)
.await
}
async fn set_user_api_key_feature_settings_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
feature_settings: Option<serde_json::Value>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, true)
.await
}
async fn set_api_key_usage_totals(
@@ -700,6 +886,11 @@ WHERE id = ?
total_tokens: u64,
total_cost_usd: f64,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
if !total_cost_usd.is_finite() {
return Err(DataLayerError::InvalidInput(
"api_keys.total_cost_usd is not finite".to_string(),
));
}
sqlx::query(
r#"
UPDATE api_keys
@@ -710,8 +901,8 @@ SET total_requests = ?,
WHERE id = ?
"#,
)
.bind(total_requests as i64)
.bind(total_tokens as i64)
.bind(i64_from_u64(total_requests, "api_keys.total_requests")?)
.bind(i64_from_u64(total_tokens, "api_keys.total_tokens")?)
.bind(total_cost_usd)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
@@ -726,11 +917,21 @@ WHERE id = ?
user_id: &str,
api_key_id: &str,
) -> Result<bool, DataLayerError> {
self.delete_api_key(api_key_id, Some(user_id), false).await
self.delete_api_key(api_key_id, Some(user_id), false, false)
.await
}
async fn delete_user_api_key_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
) -> Result<bool, DataLayerError> {
self.delete_api_key(api_key_id, Some(user_id), false, true)
.await
}
async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result<bool, DataLayerError> {
self.delete_api_key(api_key_id, None, true).await
self.delete_api_key(api_key_id, None, true, false).await
}
async fn set_standalone_api_key_feature_settings(
@@ -760,12 +961,65 @@ WHERE id = ?
}
impl MysqlAuthApiKeyReadRepository {
async fn update_user_api_key_basic_scoped(
&self,
record: UpdateUserApiKeyBasicRecord,
require_unlocked: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE api_keys
SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END,
name = CASE WHEN ? THEN ? ELSE name END,
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END,
ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END,
feature_settings = CASE WHEN ? THEN ? ELSE feature_settings END,
updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
AND (? = 0 OR is_locked = 0)
"#,
)
.bind(record.key_encrypted_present)
.bind(record.key_encrypted.as_deref())
.bind(record.name_present)
.bind(record.name.as_deref())
.bind(record.rate_limit_present)
.bind(record.rate_limit)
.bind(record.concurrent_limit_present)
.bind(record.concurrent_limit)
.bind(record.ip_rules.is_some())
.bind(json_string_from_nested_string_list(
&record.ip_rules,
"api_keys.ip_rules",
)?)
.bind(record.feature_settings.is_some())
.bind(optional_json_to_string(
&record.feature_settings.clone().flatten(),
"api_keys.feature_settings",
)?)
.bind(current_unix_secs() as i64)
.bind(&record.api_key_id)
.bind(&record.user_id)
.bind(require_unlocked)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.reload_export_by_id(&record.api_key_id).await
}
async fn set_active(
&self,
api_key_id: &str,
user_id: Option<&str>,
is_active: bool,
is_standalone: bool,
require_unlocked: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new("UPDATE api_keys SET is_active = ");
builder
@@ -779,7 +1033,115 @@ impl MysqlAuthApiKeyReadRepository {
if let Some(user_id) = user_id {
builder.push(" AND user_id = ").push_bind(user_id);
}
builder.build().execute(&self.pool).await.map_sql_err()?;
if require_unlocked {
builder.push(" AND is_locked = ").push_bind(false);
}
let result = builder.build().execute(&self.pool).await.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.reload_export_by_id(api_key_id).await
}
async fn set_user_api_key_allowed_providers_scoped(
&self,
user_id: &str,
api_key_id: &str,
allowed_providers: Option<Vec<String>>,
require_unlocked: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE api_keys
SET allowed_providers = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
AND (? = 0 OR is_locked = 0)
"#,
)
.bind(json_string_from_string_list(
allowed_providers.as_ref(),
"api_keys.allowed_providers",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.bind(require_unlocked)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.reload_export_by_id(api_key_id).await
}
async fn set_user_api_key_force_capabilities_scoped(
&self,
user_id: &str,
api_key_id: &str,
force_capabilities: Option<serde_json::Value>,
require_unlocked: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE api_keys
SET force_capabilities = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
AND (? = 0 OR is_locked = 0)
"#,
)
.bind(optional_json_to_string(
&force_capabilities,
"api_keys.force_capabilities",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.bind(require_unlocked)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.reload_export_by_id(api_key_id).await
}
async fn set_user_api_key_feature_settings_scoped(
&self,
user_id: &str,
api_key_id: &str,
feature_settings: Option<serde_json::Value>,
require_unlocked: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE api_keys
SET feature_settings = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
AND (? = 0 OR is_locked = 0)
"#,
)
.bind(optional_json_to_string(
&feature_settings,
"api_keys.feature_settings",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.bind(require_unlocked)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.reload_export_by_id(api_key_id).await
}
@@ -788,22 +1150,76 @@ impl MysqlAuthApiKeyReadRepository {
api_key_id: &str,
user_id: Option<&str>,
is_standalone: bool,
require_unlocked: bool,
) -> Result<bool, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new("DELETE FROM api_keys WHERE id = ");
builder
.push_bind(api_key_id)
.push(" AND is_standalone = ")
.push_bind(is_standalone);
if let Some(user_id) = user_id {
builder.push(" AND user_id = ").push_bind(user_id);
}
let rows_affected = builder
.build()
.execute(&self.pool)
let mut tx = self.pool.begin().await.map_sql_err()?;
let matching_api_key = if let Some(user_id) = user_id {
if require_unlocked {
sqlx::query_scalar::<_, String>(
"SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 AND is_locked = 0 FOR UPDATE",
)
.bind(api_key_id)
.bind(user_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
} else {
sqlx::query_scalar::<_, String>(
"SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 FOR UPDATE",
)
.bind(api_key_id)
.bind(user_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
}
} else {
sqlx::query_scalar::<_, String>(
"SELECT id FROM api_keys WHERE id = ? AND is_standalone = 1 FOR UPDATE",
)
.bind(api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
};
if matching_api_key.is_none() {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
sqlx::query(
"UPDATE wallets SET status = 'disabled', updated_at = UNIX_TIMESTAMP() WHERE api_key_id = ? AND status <> 'disabled'",
)
.bind(api_key_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
for sql in MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL {
sqlx::query(sql)
.bind(api_key_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
for sql in MYSQL_DELETE_API_KEY_DEPENDENTS_SQL {
sqlx::query(sql)
.bind(api_key_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
let result = sqlx::query("DELETE FROM api_keys WHERE id = ? AND is_standalone = ?")
.bind(api_key_id)
.bind(is_standalone)
.execute(&mut *tx)
.await
.map_sql_err()?;
if result.rows_affected() != 1 {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
tx.commit().await.map_sql_err()?;
Ok(true)
}
}
@@ -827,6 +1243,7 @@ async fn summarize_api_keys(
is_standalone: bool,
now_unix_secs: u64,
) -> Result<AuthApiKeyExportSummary, DataLayerError> {
let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?;
let row = sqlx::query(
r#"
SELECT
@@ -836,7 +1253,7 @@ FROM api_keys
WHERE is_standalone = ?
"#,
)
.bind(now_unix_secs as i64)
.bind(now_unix_secs)
.bind(is_standalone)
.fetch_one(pool)
.await
@@ -1037,7 +1454,36 @@ fn map_auth_api_key_export_row(
#[cfg(test)]
mod tests {
use super::MysqlAuthApiKeyReadRepository;
use super::{
MysqlAuthApiKeyReadRepository, MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL,
MYSQL_DELETE_API_KEY_DEPENDENTS_SQL,
};
#[test]
fn api_key_delete_sql_preserves_ids_and_removes_private_snapshots() {
for table in [
"request_candidates",
"video_tasks",
"`usage`",
"stats_daily_api_key",
] {
assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| {
sql.starts_with(&format!("UPDATE {table} "))
&& sql.contains("SET api_key_name = NULL")
&& sql.ends_with("WHERE api_key_id = ?")
}));
}
assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL
.iter()
.any(|sql| sql
.starts_with("UPDATE audit_logs SET description = 'deleted API key event'")));
assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| sql
.starts_with("UPDATE payment_callbacks SET payload = NULL, error_message = NULL")));
assert_eq!(
MYSQL_DELETE_API_KEY_DEPENDENTS_SQL,
&["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"]
);
}
#[tokio::test]
async fn repository_builds_from_lazy_pool() {
@@ -3,7 +3,7 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use aether_data_contracts::repository::auth_modules::*;
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_eq, push_limit, WhereClause};
use aether_data_query::{push_eq, WhereClause};
use crate::error::SqlResultExt;
use crate::MysqlPool;
@@ -35,6 +35,87 @@ SELECT
FROM ldap_configs
"#;
const UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL: &str = r#"
UPDATE ldap_configs
SET
server_url = ?,
bind_dn = ?,
base_dn = ?,
user_search_filter = ?,
username_attr = ?,
email_attr = ?,
display_name_attr = ?,
is_enabled = ?,
is_exclusive = ?,
use_starttls = ?,
connect_timeout = ?,
updated_at = GREATEST(updated_at + 1, ?)
WHERE singleton_key = 1
AND server_url <=> ?
AND bind_dn <=> ?
AND BINARY bind_password_encrypted <=> BINARY ?
AND base_dn <=> ?
AND user_search_filter <=> ?
AND username_attr <=> ?
AND email_attr <=> ?
AND display_name_attr <=> ?
AND is_enabled <=> ?
AND is_exclusive <=> ?
AND use_starttls <=> ?
AND connect_timeout <=> ?
"#;
const UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL: &str = r#"
UPDATE ldap_configs
SET
server_url = ?,
bind_dn = ?,
bind_password_encrypted = ?,
base_dn = ?,
user_search_filter = ?,
username_attr = ?,
email_attr = ?,
display_name_attr = ?,
is_enabled = ?,
is_exclusive = ?,
use_starttls = ?,
connect_timeout = ?,
updated_at = GREATEST(updated_at + 1, ?)
WHERE singleton_key = 1
AND server_url <=> ?
AND bind_dn <=> ?
AND BINARY bind_password_encrypted <=> BINARY ?
AND base_dn <=> ?
AND user_search_filter <=> ?
AND username_attr <=> ?
AND email_attr <=> ?
AND display_name_attr <=> ?
AND is_enabled <=> ?
AND is_exclusive <=> ?
AND use_starttls <=> ?
AND connect_timeout <=> ?
"#;
const INSERT_LDAP_CONFIG_SQL: &str = r#"
INSERT INTO ldap_configs (
singleton_key,
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter,
username_attr,
email_attr,
display_name_attr,
is_enabled,
is_exclusive,
use_starttls,
connect_timeout,
created_at,
updated_at
) VALUES (1, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#;
#[derive(Debug, Clone)]
pub struct MysqlAuthModuleReadRepository {
pool: MysqlPool,
@@ -72,8 +153,7 @@ async fn get_ldap_config(
pool: &MysqlPool,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(LDAP_CONFIG_COLUMNS);
builder.push(" ORDER BY id ASC");
push_limit(&mut builder, 1);
builder.push(" WHERE singleton_key = 1");
let row = builder.build().fetch_optional(pool).await.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
}
@@ -106,97 +186,208 @@ impl AuthModuleReadRepository for MysqlAuthModuleRepository {
#[async_trait]
impl AuthModuleWriteRepository for MysqlAuthModuleRepository {
async fn upsert_ldap_config(
async fn compare_and_swap_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
expected: Option<&StoredLdapModuleConfig>,
replacement: &StoredLdapModuleConfig,
bind_password_update: &LdapBindPasswordUpdate,
) -> Result<CompareAndSwapLdapConfigResult, DataLayerError> {
let persisted =
ldap_config_after_password_update(expected, replacement, bind_password_update)?;
let now = now_unix_secs();
let updated = sqlx::query(
let Some(expected) = expected else {
let insert = sqlx::query(INSERT_LDAP_CONFIG_SQL)
.bind(&persisted.server_url)
.bind(&persisted.bind_dn)
.bind(persisted.bind_password_encrypted.as_deref())
.bind(&persisted.base_dn)
.bind(persisted.user_search_filter.as_deref())
.bind(persisted.username_attr.as_deref())
.bind(persisted.email_attr.as_deref())
.bind(persisted.display_name_attr.as_deref())
.bind(persisted.is_enabled)
.bind(persisted.is_exclusive)
.bind(persisted.use_starttls)
.bind(persisted.connect_timeout)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await;
return match insert {
Ok(result) if result.rows_affected() == 1 => {
Ok(CompareAndSwapLdapConfigResult::Applied(persisted))
}
Ok(_) => Ok(CompareAndSwapLdapConfigResult::Conflict),
Err(error)
if error
.as_database_error()
.is_some_and(|error| error.is_unique_violation()) =>
{
Ok(CompareAndSwapLdapConfigResult::Conflict)
}
Err(error) => Err(DataLayerError::sql(error)),
};
};
let rows_affected = match bind_password_update {
LdapBindPasswordUpdate::Preserve => {
sqlx::query(UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL)
.bind(&replacement.server_url)
.bind(&replacement.bind_dn)
.bind(&replacement.base_dn)
.bind(replacement.user_search_filter.as_deref())
.bind(replacement.username_attr.as_deref())
.bind(replacement.email_attr.as_deref())
.bind(replacement.display_name_attr.as_deref())
.bind(replacement.is_enabled)
.bind(replacement.is_exclusive)
.bind(replacement.use_starttls)
.bind(replacement.connect_timeout)
.bind(now as i64)
.bind(&expected.server_url)
.bind(&expected.bind_dn)
.bind(expected.bind_password_encrypted.as_deref())
.bind(&expected.base_dn)
.bind(expected.user_search_filter.as_deref())
.bind(expected.username_attr.as_deref())
.bind(expected.email_attr.as_deref())
.bind(expected.display_name_attr.as_deref())
.bind(expected.is_enabled)
.bind(expected.is_exclusive)
.bind(expected.use_starttls)
.bind(expected.connect_timeout)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected()
}
LdapBindPasswordUpdate::Set(_) | LdapBindPasswordUpdate::Clear => {
sqlx::query(UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL)
.bind(&replacement.server_url)
.bind(&replacement.bind_dn)
.bind(persisted.bind_password_encrypted.as_deref())
.bind(&replacement.base_dn)
.bind(replacement.user_search_filter.as_deref())
.bind(replacement.username_attr.as_deref())
.bind(replacement.email_attr.as_deref())
.bind(replacement.display_name_attr.as_deref())
.bind(replacement.is_enabled)
.bind(replacement.is_exclusive)
.bind(replacement.use_starttls)
.bind(replacement.connect_timeout)
.bind(now as i64)
.bind(&expected.server_url)
.bind(&expected.bind_dn)
.bind(expected.bind_password_encrypted.as_deref())
.bind(&expected.base_dn)
.bind(expected.user_search_filter.as_deref())
.bind(expected.username_attr.as_deref())
.bind(expected.email_attr.as_deref())
.bind(expected.display_name_attr.as_deref())
.bind(expected.is_enabled)
.bind(expected.is_exclusive)
.bind(expected.use_starttls)
.bind(expected.connect_timeout)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected()
}
};
if rows_affected == 1 {
Ok(CompareAndSwapLdapConfigResult::Applied(persisted))
} else {
Ok(CompareAndSwapLdapConfigResult::Conflict)
}
}
async fn delete_ldap_config_if_matches(
&self,
expected: &StoredLdapModuleConfig,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
r#"
UPDATE ldap_configs
SET
server_url = ?,
bind_dn = ?,
bind_password_encrypted = ?,
base_dn = ?,
user_search_filter = ?,
username_attr = ?,
email_attr = ?,
display_name_attr = ?,
is_enabled = ?,
is_exclusive = ?,
use_starttls = ?,
connect_timeout = ?,
updated_at = ?
WHERE id = (
SELECT id FROM (
SELECT id
FROM ldap_configs
ORDER BY id ASC
LIMIT 1
) selected_ldap_config
)
DELETE FROM ldap_configs
WHERE singleton_key = 1
AND server_url <=> ?
AND bind_dn <=> ?
AND BINARY bind_password_encrypted <=> BINARY ?
AND base_dn <=> ?
AND user_search_filter <=> ?
AND username_attr <=> ?
AND email_attr <=> ?
AND display_name_attr <=> ?
AND is_enabled <=> ?
AND is_exclusive <=> ?
AND use_starttls <=> ?
AND connect_timeout <=> ?
"#,
)
.bind(&config.server_url)
.bind(&config.bind_dn)
.bind(config.bind_password_encrypted.as_deref())
.bind(&config.base_dn)
.bind(config.user_search_filter.as_deref())
.bind(config.username_attr.as_deref())
.bind(config.email_attr.as_deref())
.bind(config.display_name_attr.as_deref())
.bind(config.is_enabled)
.bind(config.is_exclusive)
.bind(config.use_starttls)
.bind(config.connect_timeout)
.bind(now as i64)
.bind(&expected.server_url)
.bind(&expected.bind_dn)
.bind(expected.bind_password_encrypted.as_deref())
.bind(&expected.base_dn)
.bind(expected.user_search_filter.as_deref())
.bind(expected.username_attr.as_deref())
.bind(expected.email_attr.as_deref())
.bind(expected.display_name_attr.as_deref())
.bind(expected.is_enabled)
.bind(expected.is_exclusive)
.bind(expected.use_starttls)
.bind(expected.connect_timeout)
.execute(&self.pool)
.await
.map_sql_err()?;
if updated.rows_affected() == 0 {
sqlx::query(
r#"
INSERT INTO ldap_configs (
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter,
username_attr,
email_attr,
display_name_attr,
is_enabled,
is_exclusive,
use_starttls,
connect_timeout,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&config.server_url)
.bind(&config.bind_dn)
.bind(config.bind_password_encrypted.as_deref())
.bind(&config.base_dn)
.bind(config.user_search_filter.as_deref())
.bind(config.username_attr.as_deref())
.bind(config.email_attr.as_deref())
.bind(config.display_name_attr.as_deref())
.bind(config.is_enabled)
.bind(config.is_exclusive)
.bind(config.use_starttls)
.bind(config.connect_timeout)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
}
self.get_ldap_config().await
.map_sql_err()?
.rows_affected();
Ok(rows_affected == 1)
}
async fn compare_and_swap_ldap_bind_password(
&self,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
r#"
UPDATE ldap_configs
SET bind_password_encrypted = ?, updated_at = GREATEST(updated_at + 1, ?)
WHERE singleton_key = 1
AND BINARY bind_password_encrypted = BINARY ?
"#,
)
.bind(replacement)
.bind(now_unix_secs() as i64)
.bind(expected)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected == 1)
}
}
fn ldap_config_after_password_update(
expected: Option<&StoredLdapModuleConfig>,
replacement: &StoredLdapModuleConfig,
bind_password_update: &LdapBindPasswordUpdate,
) -> Result<StoredLdapModuleConfig, DataLayerError> {
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,
};
Ok(StoredLdapModuleConfig {
bind_password_encrypted,
..replacement.clone()
})
}
fn now_unix_secs() -> u64 {
@@ -209,8 +209,9 @@ impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
impl BackgroundTaskWriteRepository for MysqlBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
mut run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.sanitize_for_persistence();
run.validate()?;
sqlx::query(
r#"
@@ -309,8 +310,9 @@ ON DUPLICATE KEY UPDATE
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
mut event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.sanitize_for_persistence();
event.validate()?;
sqlx::query(
r#"
@@ -358,7 +360,7 @@ fn map_run_row(row: &MySqlRow) -> Result<StoredBackgroundTaskRun, DataLayerError
let finished_at_unix_secs: Option<i64> = row.try_get("finished_at_unix_secs").map_sql_err()?;
let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?;
Ok(StoredBackgroundTaskRun {
let mut run = StoredBackgroundTaskRun {
id: row.try_get("id").map_sql_err()?,
task_key: row.try_get("task_key").map_sql_err()?,
kind: BackgroundTaskKind::from_database(&kind)?,
@@ -381,12 +383,14 @@ fn map_run_row(row: &MySqlRow) -> Result<StoredBackgroundTaskRun, DataLayerError
started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(),
})
};
run.sanitize_persisted_data();
Ok(run)
}
fn map_event_row(row: &MySqlRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
Ok(StoredBackgroundTaskEvent {
let mut event = StoredBackgroundTaskEvent {
id: row.try_get("id").map_sql_err()?,
run_id: row.try_get("run_id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
@@ -396,7 +400,9 @@ fn map_event_row(row: &MySqlRow) -> Result<StoredBackgroundTaskEvent, DataLayerE
"payload_json",
)?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
})
};
event.sanitize_persisted_data();
Ok(event)
}
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
@@ -4,8 +4,9 @@ use sqlx::{mysql::MySqlRow, Row};
use aether_data_contracts::repository::billing::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord,
};
use aether_data_contracts::DataLayerError;
@@ -637,6 +638,157 @@ LIMIT 1
.transpose()
}
async fn compare_and_swap_payment_gateway_secret(
&self,
update: &PaymentGatewaySecretCasUpdate,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE payment_gateway_configs
SET merchant_key_encrypted = ?
WHERE provider = ?
AND BINARY merchant_key_encrypted = BINARY ?
"#,
)
.bind(&update.merchant_key_encrypted)
.bind(update.provider.trim().to_ascii_lowercase())
.bind(&update.expected_merchant_key_encrypted)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() == 1)
}
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_i64();
let mut tx = self.pool.begin().await.map_sql_err()?;
if mutation.expected_existing {
let current = sqlx::query(
r#"
SELECT merchant_key_encrypted
FROM payment_gateway_configs
WHERE provider = ?
LIMIT 1
FOR UPDATE
"#,
)
.bind(&provider)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let current_secret = match current.as_ref() {
Some(row) => row
.try_get::<Option<String>, _>("merchant_key_encrypted")
.map_sql_err()?,
None => {
tx.rollback().await.map_sql_err()?;
return Ok(AdminBillingMutationOutcome::NotFound);
}
};
if current_secret != mutation.expected_merchant_key_encrypted {
tx.rollback().await.map_sql_err()?;
return Ok(AdminBillingMutationOutcome::NotFound);
}
sqlx::query(
r#"
UPDATE payment_gateway_configs
SET
enabled = ?,
endpoint_url = ?,
callback_base_url = ?,
merchant_id = ?,
merchant_key_encrypted = CASE
WHEN ? THEN merchant_key_encrypted
ELSE ?
END,
pay_currency = ?,
usd_exchange_rate = ?,
min_recharge_usd = ?,
channels_json = ?,
updated_at = ?
WHERE provider = ?
"#,
)
.bind(input.enabled)
.bind(&input.endpoint_url)
.bind(input.callback_base_url.as_deref())
.bind(&input.merchant_id)
.bind(input.preserve_existing_secret)
.bind(input.merchant_key_encrypted.as_deref())
.bind(&input.pay_currency)
.bind(input.usd_exchange_rate)
.bind(input.min_recharge_usd)
.bind(json_to_string(&input.channels_json)?)
.bind(now)
.bind(&provider)
.execute(&mut *tx)
.await
.map_sql_err()?;
} else {
let inserted = sqlx::query(
r#"
INSERT INTO payment_gateway_configs (
provider, enabled, endpoint_url, callback_base_url, merchant_id,
merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd,
channels_json, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&provider)
.bind(input.enabled)
.bind(&input.endpoint_url)
.bind(input.callback_base_url.as_deref())
.bind(&input.merchant_id)
.bind(input.merchant_key_encrypted.as_deref())
.bind(&input.pay_currency)
.bind(input.usd_exchange_rate)
.bind(input.min_recharge_usd)
.bind(json_to_string(&input.channels_json)?)
.bind(now)
.bind(now)
.execute(&mut *tx)
.await;
if let Err(err) = inserted {
let unique = matches!(
&err,
sqlx::Error::Database(database_error) if database_error.is_unique_violation()
);
tx.rollback().await.map_sql_err()?;
if unique {
return Ok(AdminBillingMutationOutcome::NotFound);
}
return Err(DataLayerError::sql(err));
}
}
let row = sqlx::query(
r#"
SELECT
provider, enabled, endpoint_url, callback_base_url, merchant_id,
merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd,
channels_json, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs
FROM payment_gateway_configs
WHERE provider = ?
LIMIT 1
"#,
)
.bind(&provider)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let record = map_payment_gateway_config_mysql(&row)?;
tx.commit().await.map_sql_err()?;
Ok(AdminBillingMutationOutcome::Applied(record))
}
async fn upsert_payment_gateway_config(
&self,
input: &PaymentGatewayConfigWriteInput,
@@ -832,12 +832,12 @@ fn map_candidate_selection_row(row: &MySqlRow) -> Result<CandidateSelectionRow,
key_name: row.try_get("key_name").map_sql_err()?,
key_auth_type: row.try_get("key_auth_type").map_sql_err()?,
key_is_active: row.try_get("key_is_active").map_sql_err()?,
key_api_formats: parse_string_list(
parse_json(row.try_get("key_api_formats").ok().flatten())?,
key_api_formats: parse_stored_key_policy_string_list(
row.try_get("key_api_formats").map_sql_err()?,
"provider_api_keys.api_formats",
)?,
key_allowed_models: parse_string_list(
parse_json(row.try_get("key_allowed_models").ok().flatten())?,
key_allowed_models: parse_stored_key_policy_string_list(
row.try_get("key_allowed_models").map_sql_err()?,
"provider_api_keys.allowed_models",
)?,
key_capabilities: parse_json(row.try_get("key_capabilities").ok().flatten())?,
@@ -904,6 +904,82 @@ fn parse_string_list(
parse_string_list_value(&value, field_name)
}
fn parse_stored_key_policy_string_list(
raw: Option<String>,
field_name: &str,
) -> Result<Option<Vec<String>>, DataLayerError> {
let Some(raw) = raw else {
return Ok(None);
};
let value = serde_json::from_str::<serde_json::Value>(&raw).map_err(|err| {
DataLayerError::UnexpectedValue(format!("{field_name} contains invalid JSON: {err}"))
})?;
parse_key_policy_string_list_value(&value, field_name)
}
fn parse_key_policy_string_list_value(
value: &serde_json::Value,
field_name: &str,
) -> Result<Option<Vec<String>>, DataLayerError> {
match value {
serde_json::Value::Null => Err(DataLayerError::UnexpectedValue(format!(
"{field_name} contains JSON null; use SQL NULL for an unset policy"
))),
serde_json::Value::Array(array) => {
parse_key_policy_string_list_array(array, field_name).map(Some)
}
serde_json::Value::String(raw) => parse_embedded_key_policy_string_list(raw, field_name),
_ => Err(DataLayerError::UnexpectedValue(format!(
"{field_name} is not a JSON array"
))),
}
}
fn parse_embedded_key_policy_string_list(
raw: &str,
field_name: &str,
) -> Result<Option<Vec<String>>, DataLayerError> {
let raw = raw.trim();
if raw.is_empty() {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} contains an empty string"
)));
}
if raw.eq_ignore_ascii_case("null") {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} contains stringified JSON null; use SQL NULL for an unset policy"
)));
}
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
return parse_key_policy_string_list_value(&decoded, field_name);
}
Ok(Some(vec![raw.to_string()]))
}
fn parse_key_policy_string_list_array(
array: &[serde_json::Value],
field_name: &str,
) -> Result<Vec<String>, DataLayerError> {
let mut items = Vec::with_capacity(array.len());
for item in array {
let Some(item) = item.as_str() else {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} contains a non-string item"
)));
};
let item = item.trim();
if item.is_empty() {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} contains an empty item"
)));
}
items.push(item.to_string());
}
Ok(items)
}
fn parse_string_list_value(
value: &serde_json::Value,
field_name: &str,
@@ -1108,7 +1184,8 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
#[cfg(test)]
mod tests {
use super::{
api_format_page_query, pool_key_group_by_key_ids_query, pool_key_group_query,
api_format_page_query, parse_stored_key_policy_string_list,
pool_key_group_by_key_ids_query, pool_key_group_query,
provider_model_mapping_api_format_covers, push_key_auth_channel_filter,
requested_model_page_query, vertex_key_auth_channel_matches, ExactPageAccumulator,
MysqlMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT,
@@ -1132,6 +1209,25 @@ mod tests {
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn malformed_key_policy_never_degrades_to_unrestricted() {
for raw in ["null", "\"null\"", "\"\"", "[\"openai:chat\",null]"] {
assert!(parse_stored_key_policy_string_list(
Some(raw.to_string()),
"provider_api_keys.api_formats",
)
.is_err());
}
assert_eq!(
parse_stored_key_policy_string_list(
Some("[\"openai:chat\"]".to_string()),
"provider_api_keys.api_formats",
)
.expect("valid key policy should parse"),
Some(vec!["openai:chat".to_string()])
);
}
#[test]
fn requested_model_page_query_filters_and_pages_before_fetch() {
let query = requested_model_page_query("openai:chat", "gpt-5", 256, 256);
@@ -216,8 +216,9 @@ impl RequestCandidateReadRepository for MysqlRequestCandidateRepository {
impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
mut candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.sanitize_for_persistence();
candidate.validate()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
match upsert_candidate_in_transaction(&mut tx, candidate).await {
@@ -234,12 +235,13 @@ impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository {
async fn upsert_many(
&self,
candidates: Vec<UpsertRequestCandidateRecord>,
mut candidates: Vec<UpsertRequestCandidateRecord>,
) -> Result<usize, DataLayerError> {
if candidates.is_empty() {
return Ok(0);
}
for candidate in &candidates {
for candidate in &mut candidates {
candidate.sanitize_for_persistence();
candidate.validate()?;
}
@@ -423,26 +425,8 @@ ON DUPLICATE KEY UPDATE
THEN status_code
ELSE COALESCE(VALUES(status_code), status_code)
END,
error_type = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN error_type
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN error_type
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN error_type
ELSE COALESCE(VALUES(error_type), error_type)
END,
error_message = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN error_message
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN error_message
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN error_message
ELSE COALESCE(VALUES(error_message), error_message)
END,
error_type = VALUES(error_type),
error_message = NULL,
latency_ms = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
@@ -524,9 +508,10 @@ fn push_endpoint_in_clause<'args>(
}
fn merge_candidate(
candidate: UpsertRequestCandidateRecord,
mut candidate: UpsertRequestCandidateRecord,
existing: Option<StoredRequestCandidate>,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.sanitize_for_persistence();
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|value| {
request_candidate_lifecycle_would_regress(value.status, candidate.status)
});
@@ -538,15 +523,11 @@ fn merge_candidate(
} else {
candidate.status
};
let created_at_unix_ms = candidate
.created_at_unix_ms
let created_at_unix_ms = existing
.as_ref()
.map(|value| value.created_at_unix_ms)
.filter(|value| *value > 1000)
.or_else(|| {
existing
.as_ref()
.map(|value| value.created_at_unix_ms)
.filter(|value| *value > 1000)
})
.or_else(|| candidate.created_at_unix_ms.filter(|value| *value > 1000))
.or(candidate.started_at_unix_ms)
.or(candidate.finished_at_unix_ms)
.unwrap_or_else(current_unix_ms);
@@ -561,35 +542,36 @@ fn merge_candidate(
StoredRequestCandidate::new(
id,
candidate.request_id,
candidate
.user_id
.or_else(|| existing.as_ref().and_then(|value| value.user_id.clone())),
candidate
.api_key_id
.or_else(|| existing.as_ref().and_then(|value| value.api_key_id.clone())),
candidate
.username
.or_else(|| existing.as_ref().and_then(|value| value.username.clone())),
candidate.api_key_name.or_else(|| {
existing
.as_ref()
.and_then(|value| value.api_key_name.clone())
}),
existing
.as_ref()
.and_then(|value| value.user_id.clone())
.or(candidate.user_id),
existing
.as_ref()
.and_then(|value| value.api_key_id.clone())
.or(candidate.api_key_id),
existing
.as_ref()
.and_then(|value| value.username.clone())
.or(candidate.username),
existing
.as_ref()
.and_then(|value| value.api_key_name.clone())
.or(candidate.api_key_name),
to_i32(candidate.candidate_index)?,
to_i32(candidate.retry_index)?,
candidate.provider_id.or_else(|| {
existing
.as_ref()
.and_then(|value| value.provider_id.clone())
}),
candidate.endpoint_id.or_else(|| {
existing
.as_ref()
.and_then(|value| value.endpoint_id.clone())
}),
candidate
.key_id
.or_else(|| existing.as_ref().and_then(|value| value.key_id.clone())),
existing
.as_ref()
.and_then(|value| value.provider_id.clone())
.or(candidate.provider_id),
existing
.as_ref()
.and_then(|value| value.endpoint_id.clone())
.or(candidate.endpoint_id),
existing
.as_ref()
.and_then(|value| value.key_id.clone())
.or(candidate.key_id),
merged_status,
candidate.skip_reason.or_else(|| {
existing
@@ -617,17 +599,7 @@ fn merge_candidate(
.error_type
.or_else(|| existing.as_ref().and_then(|value| value.error_type.clone()))
},
if preserve_existing_lifecycle {
existing
.as_ref()
.and_then(|value| value.error_message.clone())
} else {
candidate.error_message.or_else(|| {
existing
.as_ref()
.and_then(|value| value.error_message.clone())
})
},
None,
if preserve_existing_lifecycle {
match existing.as_ref().and_then(|value| value.latency_ms) {
Some(value) => Some(to_i32_u64(value)?),
@@ -657,9 +629,10 @@ fn merge_candidate(
.and_then(|value| value.required_capabilities.clone())
}),
u64_to_i64(created_at_unix_ms, "request candidate created_at")?,
candidate
.started_at_unix_ms
.or_else(|| existing.as_ref().and_then(|value| value.started_at_unix_ms))
existing
.as_ref()
.and_then(|value| value.started_at_unix_ms)
.or(candidate.started_at_unix_ms)
.map(|value| u64_to_i64(value, "request candidate started_at"))
.transpose()?,
if preserve_existing_lifecycle {
@@ -915,7 +888,7 @@ mod tests {
&request_id,
"initial",
RequestCandidateStatus::Pending,
Some(json!({"initial": true})),
Some(json!({"gateway_execution_runtime": true})),
3_000_000,
);
initial.is_cached = Some(false);
@@ -923,6 +896,18 @@ mod tests {
.upsert(initial)
.await
.expect("initial candidate should insert");
sqlx::query(
"UPDATE request_candidates SET skip_reason = ?, error_type = ?, error_message = ?, extra_data = ?, required_capabilities = ? WHERE request_id = ?",
)
.bind("legacy skip reason with tenant-secret")
.bind("legacy_error_type_with_token")
.bind("Bearer legacy-secret")
.bind(r#"{"gateway_execution_runtime":true,"request_body":{"password":"secret"}}"#)
.bind(r#"{"streaming":true,"internal_capability":"secret"}"#)
.bind(&request_id)
.execute(&pool)
.await
.expect("legacy diagnostics should be injected for the conflict test");
const WRITERS: usize = 8;
let barrier = std::sync::Arc::new(tokio::sync::Barrier::new(WRITERS));
@@ -937,13 +922,22 @@ mod tests {
} else {
RequestCandidateStatus::Streaming
};
let mut extra_data = serde_json::Map::new();
extra_data.insert(format!("writer_{writer}"), json!(writer));
let extra_data = match writer {
0 => json!({"stream_completed": true}),
1 => json!({"cache_1h": true}),
2 => json!({"first_byte_time_ms": 2}),
3 => json!({"pool_key_index": 3}),
4 => json!({"priority_slot": 4}),
5 => json!({"ranking_index": 5}),
6 => json!({"phase": "provider_request"}),
7 => json!({"provider_api_format": "openai:responses"}),
_ => unreachable!("writer index is bounded by WRITERS"),
};
let mut candidate = sample_upsert(
&request_id,
format!("writer-{writer}").as_str(),
status,
Some(serde_json::Value::Object(extra_data)),
Some(extra_data),
3_100_000 + u64::try_from(writer).expect("writer index should fit") * 10,
);
if writer != 0 {
@@ -970,25 +964,60 @@ mod tests {
assert_eq!(candidate.status, RequestCandidateStatus::Success);
assert_eq!(candidate.latency_ms, Some(123));
assert_eq!(candidate.finished_at_unix_ms, Some(3_100_002));
let extra_data = candidate
.extra_data
.as_ref()
.and_then(serde_json::Value::as_object)
.expect("merged extra data should be an object");
assert_eq!(extra_data.get("initial"), Some(&json!(true)));
for writer in 0..WRITERS {
assert_eq!(
extra_data.get(format!("writer_{writer}").as_str()),
Some(&json!(writer))
);
}
assert_eq!(
candidate.extra_data,
Some(json!({
"cache_1h": true,
"first_byte_time_ms": 2,
"gateway_execution_runtime": true,
"phase": "provider_request",
"pool_key_index": 3,
"priority_slot": 4,
"provider_api_format": "openai:responses",
"ranking_index": 5,
"stream_completed": true
}))
);
let raw = sqlx::query(
"SELECT skip_reason, error_type, error_message, extra_data, required_capabilities FROM request_candidates WHERE request_id = ?",
)
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("raw candidate diagnostics should load");
assert!(
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_message")
.expect("error_message should decode")
.is_none()
);
assert_eq!(
sqlx::Row::try_get::<Option<String>, _>(&raw, "skip_reason")
.expect("skip_reason should decode")
.as_deref(),
Some("unclassified_skip")
);
assert_eq!(
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_type")
.expect("error_type should decode")
.as_deref(),
Some("unclassified_error")
);
let raw_extra = sqlx::Row::try_get::<Option<String>, _>(&raw, "extra_data")
.expect("extra_data should decode")
.and_then(|value| serde_json::from_str::<serde_json::Value>(&value).ok());
assert_eq!(raw_extra, candidate.extra_data);
let raw_capabilities =
sqlx::Row::try_get::<Option<String>, _>(&raw, "required_capabilities")
.expect("required_capabilities should decode")
.and_then(|value| serde_json::from_str::<serde_json::Value>(&value).ok());
assert_eq!(raw_capabilities, Some(json!({"streaming": true})));
let batch_request_id = format!("candidate-batch-{}", uuid::Uuid::new_v4());
let mut pending = sample_upsert(
&batch_request_id,
"batch-first",
RequestCandidateStatus::Pending,
Some(json!({"pending": true})),
Some(json!({"gateway_execution_runtime": true})),
4_000_000,
);
pending.is_cached = Some(false);
@@ -996,7 +1025,7 @@ mod tests {
&batch_request_id,
"batch-second",
RequestCandidateStatus::Streaming,
Some(json!({"streaming": true})),
Some(json!({"stream_completed": true})),
4_000_100,
);
streaming.is_cached = None;
@@ -1004,7 +1033,7 @@ mod tests {
&batch_request_id,
"batch-third",
RequestCandidateStatus::Success,
Some(json!({"success": true})),
Some(json!({"cache_1h": true})),
4_000_200,
);
success.is_cached = Some(true);
@@ -1012,7 +1041,7 @@ mod tests {
&batch_request_id,
"batch-fourth",
RequestCandidateStatus::Pending,
Some(json!({"late": true})),
Some(json!({"first_byte_time_ms": 42})),
4_000_300,
);
late_pending.is_cached = None;
@@ -1038,10 +1067,10 @@ mod tests {
assert_eq!(
batch_candidates[0].extra_data,
Some(json!({
"pending": true,
"streaming": true,
"success": true,
"late": true
"cache_1h": true,
"first_byte_time_ms": 42,
"gateway_execution_runtime": true,
"stream_completed": true
}))
);
@@ -1082,7 +1111,7 @@ mod tests {
}
#[test]
fn merge_candidate_keeps_terminal_status_when_streaming_arrives_late() {
fn merge_candidate_preserves_first_identity_and_terminal_fact() {
let existing = StoredRequestCandidate::new(
"candidate-1".to_string(),
"request-1".to_string(),
@@ -1103,7 +1132,7 @@ mod tests {
None,
Some(123),
None,
Some(serde_json::json!({"terminal": true})),
Some(serde_json::json!({"stream_completed": true})),
None,
1_000,
Some(1_001),
@@ -1115,24 +1144,24 @@ mod tests {
UpsertRequestCandidateRecord {
id: "candidate-late".to_string(),
request_id: "request-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
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: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("provider-key-1".to_string()),
status: RequestCandidateStatus::Streaming,
provider_id: Some("attacker-provider".to_string()),
endpoint_id: Some("attacker-endpoint".to_string()),
key_id: Some("attacker-provider-key".to_string()),
status: RequestCandidateStatus::Failed,
skip_reason: None,
is_cached: Some(false),
status_code: Some(200),
error_type: None,
error_message: None,
error_message: Some("Bearer secret-token".to_string()),
latency_ms: Some(9_999),
concurrent_requests: None,
extra_data: Some(serde_json::json!({"late": true})),
extra_data: Some(serde_json::json!({"gateway_execution_runtime": true})),
required_capabilities: None,
created_at_unix_ms: Some(1_050),
started_at_unix_ms: Some(1_051),
@@ -1144,11 +1173,20 @@ mod tests {
assert_eq!(merged.id, "candidate-1");
assert_eq!(merged.status, RequestCandidateStatus::Success);
assert_eq!(merged.user_id.as_deref(), Some("user-1"));
assert_eq!(merged.api_key_id.as_deref(), Some("key-1"));
assert_eq!(merged.provider_id.as_deref(), Some("provider-1"));
assert_eq!(merged.endpoint_id.as_deref(), Some("endpoint-1"));
assert_eq!(merged.key_id.as_deref(), Some("provider-key-1"));
assert!(merged.error_message.is_none());
assert_eq!(merged.latency_ms, Some(123));
assert_eq!(merged.finished_at_unix_ms, Some(1_123));
assert_eq!(
merged.extra_data,
Some(serde_json::json!({"terminal": true, "late": true}))
Some(serde_json::json!({
"gateway_execution_runtime": true,
"stream_completed": true
}))
);
}
@@ -12,6 +12,19 @@ use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, Whe
use crate::error::SqlResultExt;
use crate::MysqlPool;
const OWNER_GUARDED_UPSERT_SQL: &str = r#"
INSERT INTO gemini_file_mappings (
id, file_name, key_id, user_id, display_name, mime_type, source_hash,
created_at, expires_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
display_name = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(display_name), display_name),
mime_type = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(mime_type), mime_type),
source_hash = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(source_hash), source_hash),
expires_at = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(expires_at), expires_at)
"#;
#[derive(Debug, Clone)]
pub struct MysqlGeminiFileMappingRepository {
pool: MysqlPool,
@@ -63,6 +76,85 @@ LIMIT 1
row.as_ref().map(map_row).transpose()
}
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 row = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE BINARY file_name = BINARY ?
AND BINARY user_id = BINARY ?
AND expires_at > ?
LIMIT 1
"#,
)
.bind(file_name)
.bind(user_id)
.bind(i64_from_u64(
now_unix_secs,
"gemini_file_mappings.owner_read_now",
)?)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
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 row = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE BINARY file_name = BINARY ?
AND BINARY key_id = BINARY ?
AND BINARY user_id = BINARY ?
AND expires_at > ?
LIMIT 1
"#,
)
.bind(file_name)
.bind(key_id)
.bind(user_id)
.bind(i64_from_u64(
now_unix_secs,
"gemini_file_mappings.owner_read_now",
)?)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
@@ -185,6 +277,60 @@ ON DUPLICATE KEY UPDATE
self.reload_by_file_name(&record.file_name).await
}
async fn upsert_if_owner_matches(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
record.validate()?;
let mut transaction = self.pool.begin().await.map_sql_err()?;
sqlx::query(OWNER_GUARDED_UPSERT_SQL)
.bind(&record.id)
.bind(&record.file_name)
.bind(&record.key_id)
.bind(&record.user_id)
.bind(&record.display_name)
.bind(&record.mime_type)
.bind(&record.source_hash)
.bind(current_unix_secs() as i64)
.bind(i64_from_u64(
record.expires_at_unix_secs,
"gemini_file_mappings.expires_at",
)?)
.execute(&mut *transaction)
.await
.map_sql_err()?;
let row = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE file_name = ?
LIMIT 1
FOR UPDATE
"#,
)
.bind(&record.file_name)
.fetch_one(&mut *transaction)
.await
.map_sql_err()?;
let stored = map_row(&row)?;
let owner_matches = stored.file_name == record.file_name
&& stored.key_id == record.key_id
&& stored.user_id == record.user_id;
transaction.commit().await.map_sql_err()?;
Ok(owner_matches.then_some(stored))
}
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ?")
.bind(file_name)
@@ -195,6 +341,43 @@ ON DUPLICATE KEY UPDATE
Ok(rows_affected > 0)
}
async fn delete_by_file_name_for_user(
&self,
file_name: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let rows_affected =
sqlx::query(
"DELETE FROM gemini_file_mappings WHERE BINARY file_name = BINARY ? AND BINARY user_id = BINARY ?",
)
.bind(file_name)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn delete_by_file_name_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
"DELETE FROM gemini_file_mappings WHERE BINARY file_name = BINARY ? AND BINARY key_id = BINARY ? AND BINARY user_id = BINARY ?",
)
.bind(file_name)
.bind(key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn delete_by_id(
&self,
mapping_id: &str,
@@ -282,6 +465,11 @@ fn apply_list_filters(
where_clause: &mut WhereClause,
query: &GeminiFileMappingListQuery,
) {
if let Some(user_id) = query.user_id.as_deref() {
where_clause.push_next(builder);
builder.push("BINARY user_id = BINARY ");
builder.push_bind(user_id.to_string());
}
if !query.include_expired {
where_clause.push_next(builder);
builder.push("expires_at > ");
@@ -343,13 +531,17 @@ fn map_row(row: &MySqlRow) -> Result<StoredGeminiFileMapping, DataLayerError> {
#[cfg(test)]
mod tests {
use super::{build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository};
use super::{
build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository,
OWNER_GUARDED_UPSERT_SQL,
};
use aether_data_contracts::repository::gemini_file_mappings::GeminiFileMappingListQuery;
use sqlx::Execute;
#[test]
fn list_query_uses_shared_mysql_filter_and_pagination_rendering() {
let query = GeminiFileMappingListQuery {
user_id: Some("user-1".to_string()),
include_expired: false,
search: Some(" Report ".to_string()),
offset: 5,
@@ -359,7 +551,9 @@ mod tests {
let mut count = build_list_count_query(&query);
let count_sql = count.build().sql().to_string();
assert!(count_sql.contains(" WHERE expires_at > ? AND (LOWER(file_name) LIKE ?"));
assert!(count_sql.contains(
" WHERE BINARY user_id = BINARY ? AND expires_at > ? AND (LOWER(file_name) LIKE ?"
));
assert!(!count_sql.contains("WHERE 1=1"));
let mut rows = build_list_rows_query(&query);
@@ -368,6 +562,12 @@ mod tests {
assert!(rows_sql.contains(" ORDER BY created_at DESC, file_name ASC LIMIT ? OFFSET ?"));
}
#[test]
fn owner_guarded_upsert_uses_exact_file_key_and_user_identity() {
let identity = "BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id)";
assert_eq!(OWNER_GUARDED_UPSERT_SQL.matches(identity).count(), 4);
}
#[tokio::test]
async fn repository_builds_from_lazy_pool() {
let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with(
@@ -2,10 +2,10 @@ use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
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,
};
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause};
@@ -38,8 +38,147 @@ impl MysqlManagementTokenRepository {
.map_sql_err()?;
row.as_ref().map(map_token_row).transpose()
}
async fn get_token_scoped(
&self,
token_id: &str,
expected_user_id: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(TOKEN_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(&mut builder, &mut where_clause, "id", token_id.to_string());
push_optional_eq(
&mut builder,
&mut where_clause,
"user_id",
expected_user_id.map(ToOwned::to_owned),
);
push_limit(&mut builder, 1);
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_row).transpose()
}
async fn update_management_token_scoped(
&self,
record: &UpdateManagementTokenRecord,
expected_user_id: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
let allowed_ips = json_to_string(record.allowed_ips.as_ref())?;
let permissions = json_to_string(record.permissions.as_ref())?;
let now = now_unix_secs();
sqlx::query(UPDATE_MANAGEMENT_TOKEN_SQL)
.bind(record.name.as_deref())
.bind(record.clear_description)
.bind(record.description.as_deref())
.bind(record.clear_allowed_ips)
.bind(allowed_ips)
.bind(permissions)
.bind(record.clear_expires_at)
.bind(
record
.expires_at_unix_secs
.and_then(|value| i64::try_from(value).ok()),
)
.bind(record.is_active)
.bind(now as i64)
.bind(&record.token_id)
.bind(expected_user_id)
.bind(expected_user_id)
.execute(&self.pool)
.await
.map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?;
self.get_token_scoped(&record.token_id, expected_user_id)
.await
}
async fn delete_management_token_scoped(
&self,
token_id: &str,
expected_user_id: Option<&str>,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"DELETE FROM management_tokens WHERE id = ? AND (? IS NULL OR user_id = ?)",
)
.bind(token_id)
.bind(expected_user_id)
.bind(expected_user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn set_management_token_active_scoped(
&self,
token_id: &str,
expected_user_id: Option<&str>,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let result = sqlx::query(
"UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ? AND (? IS NULL OR user_id = ?)",
)
.bind(is_active)
.bind(now_unix_secs() as i64)
.bind(token_id)
.bind(expected_user_id)
.bind(expected_user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token_scoped(token_id, expected_user_id).await
}
async fn regenerate_management_token_secret_scoped(
&self,
mutation: &RegenerateManagementTokenSecret,
expected_user_id: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
let result = sqlx::query(
r#"
UPDATE management_tokens
SET token_hash = ?, token_prefix = ?, updated_at = ?
WHERE id = ? AND (? IS NULL OR user_id = ?)
"#,
)
.bind(&mutation.token_hash)
.bind(mutation.token_prefix.as_deref())
.bind(now_unix_secs() as i64)
.bind(&mutation.token_id)
.bind(expected_user_id)
.bind(expected_user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token_scoped(&mutation.token_id, expected_user_id)
.await
}
}
const UPDATE_MANAGEMENT_TOKEN_SQL: &str = r#"
UPDATE management_tokens
SET name = COALESCE(?, name),
description = CASE WHEN ? THEN NULL ELSE COALESCE(?, description) END,
allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END,
permissions = COALESCE(?, permissions),
expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END,
is_active = COALESCE(?, is_active),
updated_at = ?
WHERE id = ? AND (? IS NULL OR user_id = ?)
"#;
const TOKEN_COLUMNS: &str = r#"
SELECT
id,
@@ -59,6 +198,39 @@ SELECT
FROM management_tokens
"#;
const LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL: &str = r#"
SELECT id
FROM users
WHERE id = ?
AND is_active = TRUE
AND is_deleted = FALSE
AND LOWER(role) = 'admin'
AND security_version = ?
FOR UPDATE
"#;
const LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL: &str = r#"
SELECT
id,
user_id,
token_hash,
name,
description,
token_prefix,
allowed_ips,
permissions,
expires_at AS expires_at_unix_secs,
last_used_at AS last_used_at_unix_secs,
last_used_ip,
COALESCE(usage_count, 0) AS usage_count,
is_active,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM management_tokens
WHERE id = ?
FOR UPDATE
"#;
const TOKEN_WITH_USER_COLUMNS: &str = r#"
SELECT
mt.id,
@@ -220,71 +392,29 @@ INSERT INTO management_tokens (
&self,
record: &UpdateManagementTokenRecord,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
let current = self.get_token(&record.token_id).await?;
let Some(current) = current else {
return Ok(None);
};
let name = record.name.as_deref().unwrap_or(&current.name);
let description = if record.clear_description {
None
} else {
record
.description
.as_deref()
.or(current.description.as_deref())
};
let allowed_ips = if record.clear_allowed_ips {
None
} else {
record.allowed_ips.as_ref().or(current.allowed_ips.as_ref())
};
let permissions = record.permissions.as_ref().or(current.permissions.as_ref());
let expires_at = if record.clear_expires_at {
None
} else {
record.expires_at_unix_secs.or(current.expires_at_unix_secs)
};
let is_active = record.is_active.unwrap_or(current.is_active);
let now = now_unix_secs();
self.update_management_token_scoped(record, None).await
}
let result = sqlx::query(
r#"
UPDATE management_tokens
SET name = ?,
description = ?,
allowed_ips = ?,
permissions = ?,
expires_at = ?,
is_active = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(description)
.bind(json_to_string(allowed_ips)?)
.bind(json_to_string(permissions)?)
.bind(expires_at.and_then(|value| i64::try_from(value).ok()))
.bind(is_active)
.bind(now as i64)
.bind(&record.token_id)
.execute(&self.pool)
.await
.map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&record.token_id).await
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))
.await
}
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM management_tokens WHERE id = ?")
.bind(token_id)
.execute(&self.pool)
self.delete_management_token_scoped(token_id, None).await
}
async fn delete_management_token_for_user(
&self,
token_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
self.delete_management_token_scoped(token_id, Some(user_id))
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn set_management_token_active(
@@ -292,43 +422,135 @@ WHERE id = ?
token_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let result =
sqlx::query("UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ?")
.bind(is_active)
.bind(now_unix_secs() as i64)
.bind(token_id)
.execute(&self.pool)
self.set_management_token_active_scoped(token_id, None, is_active)
.await
}
async fn set_management_token_active_for_user(
&self,
token_id: &str,
user_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
self.set_management_token_active_scoped(token_id, Some(user_id), is_active)
.await
}
async fn activate_management_token_if_matches(
&self,
mutation: &ActivateManagementTokenIfMatches,
) -> Result<bool, DataLayerError> {
mutation.validate()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
let eligible_user =
sqlx::query_scalar::<_, String>(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL)
.bind(&mutation.expected_token.user_id)
.bind(mutation.expected_user_security_version)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
if eligible_user.is_none() {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
self.get_token(token_id).await
let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL)
.bind(&mutation.expected_token.id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let snapshot_matches = match locked.as_ref() {
Some(row) => {
let token_hash: String = row.try_get("token_hash").map_sql_err()?;
let token = map_token_row(row)?;
mutation.matches_locked_token_snapshot(&token, &token_hash)
}
None => false,
};
if !snapshot_matches {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
let result = sqlx::query(
r#"
UPDATE management_tokens
SET is_active = TRUE, updated_at = ?
WHERE id = ?
AND BINARY token_hash = BINARY ?
AND is_active = FALSE
AND (expires_at IS NULL OR expires_at > ?)
"#,
)
.bind(now_unix_secs() as i64)
.bind(&mutation.expected_token.id)
.bind(&mutation.token_hash)
.bind(i64::try_from(mutation.now_unix_secs).unwrap_or(i64::MAX))
.execute(&mut *tx)
.await
.map_sql_err()?;
if result.rows_affected() != 1 {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
tx.commit().await.map_sql_err()?;
Ok(true)
}
async fn delete_inactive_management_token_if_matches(
&self,
mutation: &ActivateManagementTokenIfMatches,
) -> Result<bool, DataLayerError> {
mutation.validate()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL)
.bind(&mutation.expected_token.id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let snapshot_matches = match locked.as_ref() {
Some(row) => {
let token_hash: String = row.try_get("token_hash").map_sql_err()?;
let token = map_token_row(row)?;
mutation.matches_locked_token_snapshot(&token, &token_hash)
}
None => false,
};
if !snapshot_matches {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
let result = sqlx::query(
"DELETE FROM management_tokens WHERE id = ? AND BINARY token_hash = BINARY ? AND is_active = FALSE",
)
.bind(&mutation.expected_token.id)
.bind(&mutation.token_hash)
.execute(&mut *tx)
.await
.map_sql_err()?;
if result.rows_affected() != 1 {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
tx.commit().await.map_sql_err()?;
Ok(true)
}
async fn regenerate_management_token_secret(
&self,
mutation: &RegenerateManagementTokenSecret,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
let result = sqlx::query(
r#"
UPDATE management_tokens
SET token_hash = ?, token_prefix = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&mutation.token_hash)
.bind(mutation.token_prefix.as_deref())
.bind(now_unix_secs() as i64)
.bind(&mutation.token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&mutation.token_id).await
self.regenerate_management_token_secret_scoped(mutation, None)
.await
}
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))
.await
}
async fn record_management_token_usage(
@@ -365,8 +587,18 @@ fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
value.and_then(|value| u64::try_from(value).ok())
fn non_negative_u64(value: i64, field_name: &str) -> Result<u64, DataLayerError> {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"management_tokens.{field_name} must not be negative"
))
})
}
fn optional_unix_secs(value: Option<i64>, field_name: &str) -> Result<Option<u64>, DataLayerError> {
value
.map(|value| non_negative_u64(value, field_name))
.transpose()
}
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
@@ -420,15 +652,30 @@ fn map_token_row(row: &MySqlRow) -> Result<StoredManagementToken, DataLayerError
)
.with_permissions(json_from_string(row.try_get("permissions").map_sql_err()?)?)
.with_runtime_fields(
optional_unix_secs(row.try_get("expires_at_unix_secs").map_sql_err()?),
optional_unix_secs(row.try_get("last_used_at_unix_secs").map_sql_err()?),
optional_unix_secs(
row.try_get("expires_at_unix_secs").map_sql_err()?,
"expires_at",
)?,
optional_unix_secs(
row.try_get("last_used_at_unix_secs").map_sql_err()?,
"last_used_at",
)?,
row.try_get("last_used_ip").map_sql_err()?,
u64::try_from(row.try_get::<i64, _>("usage_count").map_sql_err()?).unwrap_or(0),
non_negative_u64(
row.try_get::<i64, _>("usage_count").map_sql_err()?,
"usage_count",
)?,
row.try_get("is_active").map_sql_err()?,
)
.with_timestamps(
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
optional_unix_secs(
row.try_get("created_at_unix_ms").map_sql_err()?,
"created_at",
)?,
optional_unix_secs(
row.try_get("updated_at_unix_secs").map_sql_err()?,
"updated_at",
)?,
))
}
@@ -454,7 +701,11 @@ fn map_token_with_user_row(
#[cfg(test)]
mod tests {
use super::MysqlManagementTokenRepository;
use super::{
non_negative_u64, optional_unix_secs, MysqlManagementTokenRepository,
LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL, LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL,
UPDATE_MANAGEMENT_TOKEN_SQL,
};
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
#[tokio::test]
@@ -468,6 +719,60 @@ mod tests {
let _repository = MysqlManagementTokenRepository::new(pool);
}
#[test]
fn mysql_install_activation_locks_admin_identity_and_token_snapshot() {
for predicate in [
"is_active = TRUE",
"is_deleted = FALSE",
"LOWER(role) = 'admin'",
"security_version = ?",
"FOR UPDATE",
] {
assert!(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL.contains(predicate));
}
for column in [
"token_hash",
"name",
"description",
"token_prefix",
"allowed_ips",
"permissions",
"expires_at_unix_secs",
"last_used_at_unix_secs",
"last_used_ip",
"usage_count",
"is_active",
"created_at_unix_ms",
"updated_at_unix_secs",
"FOR UPDATE",
] {
assert!(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL.contains(column));
}
}
#[test]
fn mysql_management_token_mapping_rejects_negative_integer_state() {
assert!(optional_unix_secs(Some(-1), "expires_at").is_err());
assert_eq!(
optional_unix_secs(None, "expires_at").expect("SQL NULL should remain optional"),
None
);
assert!(non_negative_u64(-1, "usage_count").is_err());
}
#[test]
fn mysql_management_token_updates_patch_only_explicit_fields() {
for clause in [
"name = COALESCE(?, name)",
"allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END",
"permissions = COALESCE(?, permissions)",
"expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END",
"is_active = COALESCE(?, is_active)",
] {
assert!(UPDATE_MANAGEMENT_TOKEN_SQL.contains(clause));
}
}
#[test]
fn mysql_management_token_pool_config_remains_driver_specific() {
let config = SqlDatabaseConfig {
@@ -3,7 +3,7 @@ use sqlx::{mysql::MySqlRow, Row};
use aether_data_contracts::repository::oauth_providers::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
UpsertOAuthProviderConfigRecord,
UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord,
};
use aether_data_contracts::DataLayerError;
@@ -115,6 +115,13 @@ WHERE users.is_active = 1
)
"#;
const COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL: &str = r#"
UPDATE oauth_providers
SET client_secret_encrypted = ?
WHERE BINARY provider_type = BINARY ?
AND BINARY client_secret_encrypted = BINARY ?
"#;
#[async_trait]
impl OAuthProviderReadRepository for MysqlOAuthProviderRepository {
async fn list_oauth_provider_configs(
@@ -158,12 +165,56 @@ impl OAuthProviderReadRepository for MysqlOAuthProviderRepository {
#[async_trait]
impl OAuthProviderWriteRepository for MysqlOAuthProviderRepository {
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 now = now_unix_secs();
let mut tx = self.pool.begin().await.map_sql_err()?;
let existing_enabled: Option<bool> = if record.is_enabled || force_disable {
None
} else {
sqlx::query_scalar::<_, String>(
"SELECT provider_type FROM oauth_providers ORDER BY provider_type FOR UPDATE",
)
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = ?")
.bind(&record.provider_type)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
};
if existing_enabled == Some(true) {
let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL)
.bind(&record.provider_type)
.bind(&record.provider_type)
.bind(ldap_exclusive)
.bind(&record.provider_type)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let affected_count =
usize::try_from(row.try_get::<i64, _>("locked_count").map_sql_err()?.max(0))
.map_err(|_| {
DataLayerError::UnexpectedValue(
"oauth_providers.locked_user_count overflowed".to_string(),
)
})?;
if affected_count > 0 {
tx.rollback().await.map_sql_err()?;
return Ok(
UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation {
affected_count,
},
);
}
}
sqlx::query(
r#"
INSERT INTO oauth_providers (
@@ -228,27 +279,65 @@ ON DUPLICATE KEY UPDATE
.bind(now as i64)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
self.get_provider(&record.provider_type)
.await?
.ok_or_else(|| {
DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string())
})
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
.bind(&record.provider_type)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let provider = map_oauth_provider_row(&row)?;
tx.commit().await.map_sql_err()?;
Ok(UpsertOAuthProviderConfigOutcome::Upserted(provider))
}
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 result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?")
let result = sqlx::query(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL)
.bind(replacement)
.bind(provider_type)
.bind(expected)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
Ok(result.rows_affected() == 1)
}
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 tx = self.pool.begin().await.map_sql_err()?;
let provider_exists: Option<String> = sqlx::query_scalar(
"SELECT provider_type FROM oauth_providers WHERE provider_type = ? FOR UPDATE",
)
.bind(provider_type)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
if provider_exists.is_none() {
tx.rollback().await.map_sql_err()?;
return Ok(false);
}
let result = sqlx::query(
"DELETE FROM oauth_providers WHERE provider_type = ? AND NOT EXISTS (SELECT 1 FROM user_oauth_links WHERE user_oauth_links.provider_type = oauth_providers.provider_type)",
)
.bind(provider_type)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(result.rows_affected() == 1)
}
}
@@ -378,7 +467,16 @@ fn map_oauth_provider_row(row: &MySqlRow) -> Result<StoredOAuthProviderConfig, D
#[cfg(test)]
mod tests {
use super::MysqlOAuthProviderRepository;
use super::{MysqlOAuthProviderRepository, COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL};
#[test]
fn client_secret_cas_updates_only_the_secret_column() {
assert!(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL
.contains("SET client_secret_encrypted = ?"));
assert!(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL
.contains("BINARY client_secret_encrypted = BINARY ?"));
assert!(!COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL.contains("updated_at"));
}
#[tokio::test]
async fn repository_builds_from_lazy_pool() {
+58 -5
View File
@@ -30,13 +30,20 @@ impl MysqlPoolFactory {
}
pub fn connect_options(&self) -> Result<MySqlConnectOptions, DataLayerError> {
let ssl_mode = if self.config.pool.require_ssl {
MySqlSslMode::Required
} else {
MySqlSslMode::Preferred
};
MySqlConnectOptions::from_str(self.config.url.trim())
.map(|options| {
// Preserve explicit VERIFY_CA/VERIFY_IDENTITY from the URL.
// `require_ssl` is a minimum transport guarantee: upgrade
// weaker modes to Required, never downgrade verification.
let ssl_mode = if self.config.pool.require_ssl
&& !matches!(
options.get_ssl_mode(),
MySqlSslMode::VerifyCa | MySqlSslMode::VerifyIdentity
) {
MySqlSslMode::Required
} else {
options.get_ssl_mode()
};
options
.ssl_mode(ssl_mode)
.statement_cache_capacity(self.config.pool.statement_cache_capacity)
@@ -78,6 +85,52 @@ impl MysqlPoolFactory {
mod tests {
use super::MysqlPoolFactory;
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use sqlx::mysql::MySqlSslMode;
fn ssl_mode(url: &str, require_ssl: bool) -> MySqlSslMode {
MysqlPoolFactory::new(SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: url.to_string(),
pool: SqlPoolConfig {
require_ssl,
..SqlPoolConfig::default()
},
})
.expect("mysql config should build")
.connect_options()
.expect("mysql options should parse")
.get_ssl_mode()
}
#[test]
fn preserves_explicit_mysql_verification_modes() {
assert!(matches!(
ssl_mode(
"mysql://user:pass@localhost/aether?ssl-mode=VERIFY_IDENTITY",
false
),
MySqlSslMode::VerifyIdentity
));
assert!(matches!(
ssl_mode(
"mysql://user:pass@localhost/aether?ssl-mode=VERIFY_CA",
true
),
MySqlSslMode::VerifyCa
));
}
#[test]
fn require_ssl_only_upgrades_weak_mysql_modes() {
for mode in ["DISABLED", "PREFERRED", "REQUIRED"] {
let url = format!("mysql://user:pass@localhost/aether?ssl-mode={mode}");
assert!(matches!(ssl_mode(&url, true), MySqlSslMode::Required));
}
assert!(matches!(
ssl_mode("mysql://user:pass@localhost/aether", false),
MySqlSslMode::Preferred
));
}
#[tokio::test]
async fn factory_builds_lazy_pool_from_valid_config() {
@@ -9,9 +9,11 @@ use sqlx::{
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate,
ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate,
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -550,6 +552,48 @@ WHERE id = ?
self.reload_provider(&provider.id, "updated").await
}
pub async fn compare_and_swap_provider_config(
&self,
update: &ProviderCatalogProviderConfigCasUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.provider_id, "provider catalog provider_id")?;
let expected_config =
optional_json_to_string(&update.expected_config, "providers.expected_config")?;
let config = optional_json_to_string(&update.config, "providers.config")?;
let rows_affected = sqlx::query(
r#"
UPDATE providers
SET config = ?, updated_at = ?
WHERE id = ?
AND config <=> ?
"#,
)
.bind(config)
.bind(current_unix_secs() as i64)
.bind(&update.provider_id)
.bind(expected_config)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected == 1)
}
pub async fn compare_and_swap_provider_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.record_id, "provider catalog provider_id")?;
compare_and_swap_proxy_json(
&self.pool,
"SELECT proxy FROM providers WHERE id = ?",
"UPDATE providers SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?",
update,
"providers.proxy",
)
.await
}
pub async fn delete_provider(&self, provider_id: &str) -> Result<bool, DataLayerError> {
validate_non_empty(provider_id, "provider catalog provider_id")?;
let rows_affected = sqlx::query("DELETE FROM providers WHERE id = ?")
@@ -754,6 +798,21 @@ WHERE id = ?
self.reload_endpoint(&endpoint.id, "updated").await
}
pub async fn compare_and_swap_endpoint_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.record_id, "provider catalog endpoint_id")?;
compare_and_swap_proxy_json(
&self.pool,
"SELECT proxy FROM provider_endpoints WHERE id = ?",
"UPDATE provider_endpoints SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?",
update,
"provider_endpoints.proxy",
)
.await
}
pub async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, DataLayerError> {
validate_non_empty(endpoint_id, "provider catalog endpoint_id")?;
let rows_affected = sqlx::query("DELETE FROM provider_endpoints WHERE id = ?")
@@ -936,6 +995,46 @@ WHERE id = ?
self.reload_key(&key.id, "updated").await
}
pub async fn compare_and_swap_key_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.record_id, "provider catalog key_id")?;
compare_and_swap_proxy_json(
&self.pool,
"SELECT proxy FROM provider_api_keys WHERE id = ?",
"UPDATE provider_api_keys SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?",
update,
"provider_api_keys.proxy",
)
.await
}
pub async fn compare_and_swap_key_credentials(
&self,
update: &ProviderCatalogKeyCredentialsCasUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.key_id, "provider catalog key_id")?;
validate_non_empty(
&update.expected_provider_id,
"provider catalog expected provider_id",
)?;
let rows_affected = sqlx::query(
"UPDATE provider_api_keys SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND BINARY provider_id = BINARY ? AND BINARY COALESCE(api_key, encrypted_key) <=> BINARY ? AND BINARY auth_config <=> BINARY ?",
)
.bind(update.encrypted_api_key.as_deref())
.bind(update.encrypted_auth_config.as_deref())
.bind(&update.key_id)
.bind(&update.expected_provider_id)
.bind(update.expected_encrypted_api_key.as_deref())
.bind(update.expected_encrypted_auth_config.as_deref())
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected == 1)
}
pub async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
@@ -1354,51 +1453,18 @@ WHERE id = ?
Ok(rows_affected > 0)
}
pub 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> {
validate_non_empty(key_id, "provider catalog key_id")?;
validate_non_empty(encrypted_api_key, "provider catalog oauth api_key")?;
let rows_affected = sqlx::query(
r#"
UPDATE provider_api_keys
SET api_key = ?, auth_config = ?, expires_at = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(encrypted_api_key)
.bind(encrypted_auth_config)
.bind(optional_i64_from_u64(
expires_at_unix_secs,
"provider_api_keys.expires_at",
)?)
.bind(current_unix_secs() as i64)
.bind(key_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub 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> {
validate_non_empty(key_id, "provider catalog key_id")?;
let rows_affected = sqlx::query(
r#"
UPDATE provider_api_keys
SET oauth_invalid_at = ?, oauth_invalid_reason = ?,
auth_config = COALESCE(?, auth_config), updated_at = ?
SET oauth_invalid_at = ?, oauth_invalid_reason = ?, updated_at = ?
WHERE id = ?
"#,
)
@@ -1407,7 +1473,6 @@ WHERE id = ?
"provider_api_keys.oauth_invalid_at",
)?)
.bind(oauth_invalid_reason)
.bind(encrypted_auth_config_update)
.bind(updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64)
.bind(key_id)
.execute(&self.pool)
@@ -1755,7 +1820,7 @@ WHERE id = ?
update.expected_encrypted_auth_config.as_deref()
{
builder
.push(" AND auth_config <=> ")
.push(" AND BINARY auth_config <=> BINARY ")
.push_bind(expected_encrypted_auth_config);
}
let rows_affected = builder
@@ -1873,7 +1938,7 @@ SET health_by_format = ?, circuit_breaker_by_format = ?, updated_at = ?
WHERE id = ?
AND JSON_EXTRACT(health_by_format, '$') <=> CAST(? AS JSON)
AND JSON_EXTRACT(circuit_breaker_by_format, '$') <=> CAST(? AS JSON)
AND (? IS NULL OR auth_config <=> ?)
AND (? IS NULL OR BINARY auth_config <=> BINARY ?)
"#,
)
.bind(optional_json_to_string(
@@ -2042,6 +2107,20 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
Self::update_provider(self, provider).await
}
async fn compare_and_swap_provider_config(
&self,
update: &ProviderCatalogProviderConfigCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_swap_provider_config(self, update).await
}
async fn compare_and_swap_provider_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_swap_provider_proxy(self, update).await
}
async fn delete_provider(&self, provider_id: &str) -> Result<bool, DataLayerError> {
Self::delete_provider(self, provider_id).await
}
@@ -2077,6 +2156,13 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
Self::update_endpoint(self, endpoint).await
}
async fn compare_and_swap_endpoint_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_swap_endpoint_proxy(self, update).await
}
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, DataLayerError> {
Self::delete_endpoint(self, endpoint_id).await
}
@@ -2095,6 +2181,20 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
Self::update_key(self, key).await
}
async fn compare_and_swap_key_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_swap_key_proxy(self, update).await
}
async fn compare_and_swap_key_credentials(
&self,
update: &ProviderCatalogKeyCredentialsCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_swap_key_credentials(self, update).await
}
async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
@@ -2189,29 +2289,11 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
Self::clear_key_oauth_invalid_marker(self, key_id).await
}
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> {
Self::update_key_oauth_credentials(
self,
key_id,
encrypted_api_key,
encrypted_auth_config,
expires_at_unix_secs,
)
.await
}
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> {
Self::update_key_oauth_runtime_state(
@@ -2219,7 +2301,6 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
key_id,
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config_update,
updated_at_unix_secs,
)
.await
@@ -2490,6 +2571,44 @@ fn optional_json_to_string(
optional_json_ref_to_string(value.as_ref(), field_name)
}
async fn compare_and_swap_proxy_json(
pool: &MysqlPool,
select_sql: &'static str,
update_sql: &'static str,
update: &ProviderCatalogProxyCasUpdate,
field_name: &'static str,
) -> Result<bool, DataLayerError> {
// Legacy catalog rows may contain semantically identical JSON with Python-style
// whitespace. Comparing a re-serialized serde_json::Value directly to a TEXT column
// would make lazy credential migration conflict forever. Compare the parsed value first,
// then fence the write against the exact raw bytes that were observed.
// Outer None means the row does not exist; inner None is an existing SQL NULL proxy.
let observed_raw: Option<Option<String>> = sqlx::query_scalar::<_, Option<String>>(select_sql)
.bind(&update.record_id)
.fetch_optional(pool)
.await
.map_sql_err()?;
let Some(observed_raw) = observed_raw else {
return Ok(false);
};
let observed = optional_json_from_string(observed_raw.clone(), field_name)?;
if observed != update.expected_proxy {
return Ok(false);
}
let replacement = optional_json_to_string(&update.proxy, field_name)?;
let rows_affected = sqlx::query(update_sql)
.bind(replacement)
.bind(current_unix_secs() as i64)
.bind(&update.record_id)
.bind(observed_raw)
.execute(pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected == 1)
}
fn build_in_query<'a>(
select_sql: &'static str,
column: &'static str,
@@ -3134,11 +3253,11 @@ fn map_key_row(row: &MySqlRow) -> Result<StoredProviderCatalogKey, DataLayerErro
)?,
);
key.note = row.try_get("note").map_sql_err()?;
key.auth_type_by_format = optional_json_from_string(
let auth_type_by_format = optional_json_from_string(
row.try_get("auth_type_by_format").map_sql_err()?,
"provider_api_keys.auth_type_by_format",
)?;
key.allow_auth_channel_mismatch_formats = optional_json_from_string(
let allow_auth_channel_mismatch_formats = optional_json_from_string(
row.try_get("allow_auth_channel_mismatch_formats")
.map_sql_err()?,
"provider_api_keys.allow_auth_channel_mismatch_formats",
@@ -3204,7 +3323,10 @@ fn map_key_row(row: &MySqlRow) -> Result<StoredProviderCatalogKey, DataLayerErro
row.try_get("updated_at_unix_secs").map_sql_err()?,
"provider_api_keys.updated_at",
)?;
Ok::<_, DataLayerError>(key)
key.with_auth_channel_policy_fields(
auth_type_by_format,
allow_auth_channel_mismatch_formats,
)
})?
}
@@ -3238,6 +3360,14 @@ mod tests {
assert!(sql.contains("binary auth_config <=> binary ?"));
}
#[test]
fn credential_cas_migrates_legacy_encrypted_key_with_binary_fence() {
let source = include_str!("provider_catalog.rs");
assert!(source.contains(
"SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND BINARY provider_id = BINARY ? AND BINARY COALESCE(api_key, encrypted_key) <=> BINARY ? AND BINARY auth_config <=> BINARY ?"
));
}
#[test]
fn admin_credential_cas_has_atomic_rotation_guards() {
let source = include_str!("provider_catalog.rs");
File diff suppressed because it is too large Load Diff
@@ -1,10 +1,14 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use sqlx::{mysql::MySqlRow, Acquire, Row};
use aether_data_contracts::repository::settlement::{
finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd,
settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement,
UsageSettlementInput, SETTLEMENT_EPSILON_USD,
settlement_billing_status_for_usage_status, validate_wallet_settlement_values,
ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput,
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput,
ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation,
StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsagePolicyCostReservationState,
UsagePolicyRequestAdmissionState, UsageSettlementInput, SETTLEMENT_EPSILON_USD,
};
use aether_data_contracts::DataLayerError;
@@ -112,6 +116,142 @@ impl MysqlSettlementRepository {
}
}
fn usage_policy_cost_i64(value: u64, field: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds the integer range")))
}
fn usage_policy_cost_u64(value: i64, field: &str) -> Result<u64, DataLayerError> {
u64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("{field} must not be negative")))
}
fn usage_policy_request_admission_from_mysql_row(
row: &MySqlRow,
) -> Result<StoredUsagePolicyRequestAdmission, DataLayerError> {
let state: String = row.try_get("state").map_sql_err()?;
Ok(StoredUsagePolicyRequestAdmission {
request_id: row.try_get("request_id").map_sql_err()?,
subject_id: row.try_get("subject_id").map_sql_err()?,
event_token: row.try_get("event_token").map_sql_err()?,
admitted_at_unix_secs: usage_policy_cost_u64(
row.try_get("admitted_at_unix_secs").map_sql_err()?,
"usage policy request admitted_at",
)?,
retain_until_unix_secs: usage_policy_cost_u64(
row.try_get("retain_until_unix_secs").map_sql_err()?,
"usage policy request retain_until",
)?,
state: UsagePolicyRequestAdmissionState::parse(&state).ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"unknown usage policy request admission state {state}"
))
})?,
released_at_unix_secs: row
.try_get::<Option<i64>, _>("released_at_unix_secs")
.map_sql_err()?
.map(|value| usage_policy_cost_u64(value, "usage policy request released_at"))
.transpose()?,
})
}
const FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL: &str = r#"
SELECT request_id, subject_id, event_token,
admitted_at AS admitted_at_unix_secs,
retain_until AS retain_until_unix_secs,
state, released_at AS released_at_unix_secs
FROM usage_request_admissions
WHERE event_token = ?
FOR UPDATE
"#;
const USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL: &str =
"SET TRANSACTION ISOLATION LEVEL READ COMMITTED";
fn usage_policy_cost_reservation_from_mysql_row(
row: &MySqlRow,
) -> Result<StoredUsagePolicyCostReservation, DataLayerError> {
let state: String = row.try_get("state").map_sql_err()?;
Ok(StoredUsagePolicyCostReservation {
request_id: row.try_get("request_id").map_sql_err()?,
subject_id: row.try_get("subject_id").map_sql_err()?,
reservation_token: row.try_get("reservation_token").map_sql_err()?,
admitted_at_unix_secs: usage_policy_cost_u64(
row.try_get("admitted_at_unix_secs").map_sql_err()?,
"usage policy admitted_at",
)?,
reserved_cost_units: usage_policy_cost_u64(
row.try_get("reserved_cost_units").map_sql_err()?,
"usage policy reserved_cost_units",
)?,
actual_cost_units: row
.try_get::<Option<i64>, _>("actual_cost_units")
.map_sql_err()?
.map(|value| usage_policy_cost_u64(value, "usage policy actual_cost_units"))
.transpose()?,
state: UsagePolicyCostReservationState::parse(&state).ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"unknown usage policy reservation state {state}"
))
})?,
reservation_expires_at_unix_secs: usage_policy_cost_u64(
row.try_get("reservation_expires_at_unix_secs")
.map_sql_err()?,
"usage policy reservation_expires_at",
)?,
retain_until_unix_secs: usage_policy_cost_u64(
row.try_get("retain_until_unix_secs").map_sql_err()?,
"usage policy retain_until",
)?,
finalized_at_unix_secs: row
.try_get::<Option<i64>, _>("finalized_at_unix_secs")
.map_sql_err()?
.map(|value| usage_policy_cost_u64(value, "usage policy finalized_at"))
.transpose()?,
})
}
const FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL: &str = r#"
SELECT
request_id,
subject_id,
reservation_token,
admitted_at AS admitted_at_unix_secs,
reserved_cost_units,
actual_cost_units,
state,
reservation_expires_at AS reservation_expires_at_unix_secs,
retain_until AS retain_until_unix_secs,
finalized_at AS finalized_at_unix_secs
FROM usage_cost_reservations
WHERE reservation_token = ?
FOR UPDATE
"#;
async fn lock_usage_policy_subject_mysql(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
subject_id: &str,
) -> Result<bool, DataLayerError> {
let exists = sqlx::query_scalar::<_, String>(
r#"
SELECT id
FROM users
WHERE id = ?
FOR UPDATE
"#,
)
.bind(subject_id)
.fetch_optional(&mut **tx)
.await
.map_sql_err()?
.is_some();
Ok(exists)
}
fn usage_policy_subject_missing() -> DataLayerError {
DataLayerError::InvalidInput("usage policy subject does not exist".to_string())
}
fn settlement_from_row(row: &MySqlRow) -> Result<StoredUsageSettlement, DataLayerError> {
Ok(StoredUsageSettlement {
request_id: row.try_get("request_id").map_sql_err()?,
@@ -261,7 +401,12 @@ async fn consume_daily_quota_mysql(
wallet_can_overdraft: bool,
now_unix_secs: i64,
) -> Result<DailyQuotaDebitResult, DataLayerError> {
if total_cost_usd <= 0.0 {
if !total_cost_usd.is_finite() || total_cost_usd < 0.0 {
return Err(DataLayerError::InvalidInput(
"daily quota settlement cost must be finite and non-negative".to_string(),
));
}
if total_cost_usd == 0.0 {
return Ok(DailyQuotaDebitResult::default());
}
let rows = sqlx::query(
@@ -335,8 +480,18 @@ WHERE user_entitlement_id = ?
.fetch_one(&mut **tx)
.await
.map_sql_err()?;
if !used.is_finite() || used < 0.0 {
return Err(DataLayerError::UnexpectedValue(
"daily quota usage ledger total is invalid".to_string(),
));
}
let remaining = (grant.daily_quota_usd - used).max(0.0);
total_remaining += remaining;
if !total_remaining.is_finite() {
return Err(DataLayerError::UnexpectedValue(
"daily quota remaining total overflowed".to_string(),
));
}
grants_with_remaining.push((grant, remaining));
}
let insufficient = (!allow_wallet_overage && total_remaining + 0.000_000_01 < total_cost_usd)
@@ -386,6 +541,476 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
#[async_trait]
impl SettlementWriteRepository for MysqlSettlementRepository {
async fn reserve_usage_policy_request(
&self,
input: ReserveUsagePolicyRequestInput,
) -> Result<ReserveUsagePolicyRequestOutcome, DataLayerError> {
input.validate()?;
let admitted_at = usage_policy_cost_i64(
input.admitted_at_unix_secs,
"usage policy request admitted_at",
)?;
let retain_until = usage_policy_cost_i64(
input.retain_until_unix_secs,
"usage policy request retain_until",
)?;
let created_at = now_unix_secs()?;
// Different subjects lock different `users` rows. Under InnoDB's default REPEATABLE READ,
// two missing-token locking reads can retain compatible gap locks and then deadlock when
// both transactions try to insert the same unique event token. READ COMMITTED removes
// that gap-lock cycle while the subject row still serializes each subject's window count.
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query(USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL)
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
return Err(usage_policy_subject_missing());
}
let existing_row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL)
.bind(&input.event_token)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
if let Some(row) = existing_row.as_ref() {
let existing = usage_policy_request_admission_from_mysql_row(row)?;
if existing.request_id != input.request_id || existing.subject_id != input.subject_id {
tx.commit().await.map_sql_err()?;
return Ok(ReserveUsagePolicyRequestOutcome::Conflict);
}
if existing.admitted_at_unix_secs != input.admitted_at_unix_secs {
return Err(DataLayerError::InvalidInput(
"usage policy event_token must keep its original admitted_at".to_string(),
));
}
sqlx::query(
"UPDATE usage_request_admissions SET retain_until = GREATEST(retain_until, ?) WHERE event_token = ?",
)
.bind(retain_until)
.bind(&input.event_token)
.execute(&mut *tx)
.await
.map_sql_err()?;
let outcome = match existing.state {
UsagePolicyRequestAdmissionState::Active => {
ReserveUsagePolicyRequestOutcome::Allowed
}
UsagePolicyRequestAdmissionState::Released => {
ReserveUsagePolicyRequestOutcome::AlreadyReleased
}
};
tx.commit().await.map_sql_err()?;
return Ok(outcome);
}
for (window_index, window) in input.windows.iter().enumerate() {
let used_requests = sqlx::query_scalar::<_, i64>(
r#"
SELECT CAST(COUNT(*) AS SIGNED)
FROM usage_request_admissions
WHERE subject_id = ?
AND state = 'active'
AND admitted_at >= ?
AND admitted_at < ?
"#,
)
.bind(&input.subject_id)
.bind(usage_policy_cost_i64(
window.starts_at_unix_secs,
"usage policy request window start",
)?)
.bind(usage_policy_cost_i64(
window.ends_at_unix_secs,
"usage policy request window end",
)?)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let used_requests =
usage_policy_cost_u64(used_requests, "usage policy request used_requests")?;
if used_requests >= window.limit_requests {
let outcome = ReserveUsagePolicyRequestOutcome::Rejected {
window_index,
limit_requests: window.limit_requests,
used_requests,
};
tx.commit().await.map_sql_err()?;
return Ok(outcome);
}
}
sqlx::query(
r#"
INSERT INTO usage_request_admissions (
request_id, subject_id, event_token, admitted_at, retain_until,
state, released_at, created_at
) VALUES (?, ?, ?, ?, ?, 'active', NULL, ?)
ON DUPLICATE KEY UPDATE event_token = VALUES(event_token)
"#,
)
.bind(&input.request_id)
.bind(&input.subject_id)
.bind(&input.event_token)
.bind(admitted_at)
.bind(retain_until)
.bind(created_at)
.execute(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL)
.bind(&input.event_token)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let existing = usage_policy_request_admission_from_mysql_row(&row)?;
if existing.request_id != input.request_id || existing.subject_id != input.subject_id {
tx.commit().await.map_sql_err()?;
return Ok(ReserveUsagePolicyRequestOutcome::Conflict);
}
if existing.admitted_at_unix_secs != input.admitted_at_unix_secs {
return Err(DataLayerError::InvalidInput(
"usage policy event_token must keep its original admitted_at".to_string(),
));
}
sqlx::query(
"UPDATE usage_request_admissions SET retain_until = GREATEST(retain_until, ?) WHERE event_token = ?",
)
.bind(retain_until)
.bind(&input.event_token)
.execute(&mut *tx)
.await
.map_sql_err()?;
let outcome = match existing.state {
UsagePolicyRequestAdmissionState::Active => ReserveUsagePolicyRequestOutcome::Allowed,
UsagePolicyRequestAdmissionState::Released => {
ReserveUsagePolicyRequestOutcome::AlreadyReleased
}
};
tx.commit().await.map_sql_err()?;
Ok(outcome)
}
async fn release_usage_policy_request_admission(
&self,
input: ReleaseUsagePolicyRequestAdmissionInput,
) -> Result<Option<StoredUsagePolicyRequestAdmission>, DataLayerError> {
input.validate()?;
let released_at = usage_policy_cost_i64(
input.released_at_unix_secs,
"usage policy request released_at",
)?;
let mut tx = self.pool.begin().await.map_sql_err()?;
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
tx.commit().await.map_sql_err()?;
return Ok(None);
}
let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL)
.bind(&input.event_token)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(row) = row else {
tx.commit().await.map_sql_err()?;
return Ok(None);
};
let mut admission = usage_policy_request_admission_from_mysql_row(&row)?;
if admission.request_id != input.request_id || admission.subject_id != input.subject_id {
tx.commit().await.map_sql_err()?;
return Ok(None);
}
if input.released_at_unix_secs < admission.admitted_at_unix_secs {
return Err(DataLayerError::InvalidInput(
"usage policy released_at must not precede admitted_at".to_string(),
));
}
if admission.state == UsagePolicyRequestAdmissionState::Active {
sqlx::query(
"UPDATE usage_request_admissions SET state = 'released', released_at = ? WHERE event_token = ? AND state = 'active'",
)
.bind(released_at)
.bind(&input.event_token)
.execute(&mut *tx)
.await
.map_sql_err()?;
admission.state = UsagePolicyRequestAdmissionState::Released;
admission.released_at_unix_secs = Some(input.released_at_unix_secs);
}
tx.commit().await.map_sql_err()?;
Ok(Some(admission))
}
async fn cleanup_usage_policy_request_admissions(
&self,
now_unix_secs: u64,
batch_size: usize,
) -> Result<usize, DataLayerError> {
if batch_size == 0 {
return Ok(0);
}
let now = usage_policy_cost_i64(now_unix_secs, "usage policy request cleanup timestamp")?;
let limit = i64::try_from(batch_size).unwrap_or(i64::MAX);
let result = sqlx::query(
r#"
DELETE FROM usage_request_admissions
WHERE retain_until <= ?
ORDER BY retain_until, event_token
LIMIT ?
"#,
)
.bind(now)
.bind(limit)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() as usize)
}
async fn reserve_usage_policy_cost(
&self,
input: ReserveUsagePolicyCostInput,
) -> Result<ReserveUsagePolicyCostOutcome, DataLayerError> {
input.validate()?;
let admitted_at =
usage_policy_cost_i64(input.admitted_at_unix_secs, "usage policy admitted_at")?;
let reservation_expires_at = usage_policy_cost_i64(
input.reservation_expires_at_unix_secs,
"usage policy reservation_expires_at",
)?;
let updated_at = now_unix_secs()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
return Err(usage_policy_subject_missing());
}
let existing_row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL)
.bind(&input.reservation_token)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let existing = existing_row
.as_ref()
.map(usage_policy_cost_reservation_from_mysql_row)
.transpose()?;
if let Some(existing) = existing.as_ref() {
if existing.request_id != input.request_id || existing.subject_id != input.subject_id {
tx.commit().await.map_sql_err()?;
return Ok(ReserveUsagePolicyCostOutcome::Conflict);
}
if existing.state != UsagePolicyCostReservationState::Reserved {
let outcome = ReserveUsagePolicyCostOutcome::AlreadyTerminal {
state: existing.state,
};
tx.commit().await.map_sql_err()?;
return Ok(outcome);
}
if existing.admitted_at_unix_secs != input.admitted_at_unix_secs {
return Err(DataLayerError::InvalidInput(
"usage policy reservation_token must keep its original admitted_at".to_string(),
));
}
}
let previous_reserved_cost_units = existing
.as_ref()
.map(|reservation| reservation.reserved_cost_units)
.unwrap_or(0);
let target_reserved_cost_units =
previous_reserved_cost_units.max(input.reserved_cost_units);
for (window_index, window) in input.windows.iter().enumerate() {
let window_start =
usage_policy_cost_i64(window.starts_at_unix_secs, "usage policy window start")?;
let window_end =
usage_policy_cost_i64(window.ends_at_unix_secs, "usage policy window end")?;
let used_cost_units = sqlx::query_scalar::<_, i64>(
r#"
SELECT CAST(COALESCE(SUM(
CASE
WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0)
WHEN state = 'reserved' AND reservation_expires_at > ? THEN reserved_cost_units
ELSE 0
END
), 0) AS SIGNED)
FROM usage_cost_reservations
WHERE subject_id = ?
AND admitted_at >= ?
AND admitted_at < ?
AND reservation_token <> ?
"#,
)
.bind(admitted_at)
.bind(&input.subject_id)
.bind(window_start)
.bind(window_end)
.bind(&input.reservation_token)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let used_cost_units =
usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?;
if used_cost_units
.checked_add(target_reserved_cost_units)
.is_none_or(|total| total > window.limit_cost_units)
{
let outcome = ReserveUsagePolicyCostOutcome::Rejected {
window_index,
limit_cost_units: window.limit_cost_units,
used_cost_units,
};
tx.commit().await.map_sql_err()?;
return Ok(outcome);
}
}
let admitted_at = existing
.as_ref()
.map(|reservation| {
usage_policy_cost_i64(
reservation.admitted_at_unix_secs,
"usage policy admitted_at",
)
})
.transpose()?
.unwrap_or(admitted_at);
sqlx::query(
r#"
INSERT INTO usage_cost_reservations (
request_id, subject_id, reservation_token, admitted_at,
reserved_cost_units, actual_cost_units,
state, reservation_expires_at, retain_until, finalized_at, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, NULL, 'reserved', ?, ?, NULL, ?, ?)
ON DUPLICATE KEY UPDATE
reserved_cost_units = GREATEST(reserved_cost_units, VALUES(reserved_cost_units)),
reservation_expires_at = GREATEST(
reservation_expires_at,
VALUES(reservation_expires_at)
),
retain_until = GREATEST(retain_until, VALUES(retain_until)),
updated_at = VALUES(updated_at)
"#,
)
.bind(&input.request_id)
.bind(&input.subject_id)
.bind(&input.reservation_token)
.bind(admitted_at)
.bind(usage_policy_cost_i64(
target_reserved_cost_units,
"usage policy reserved_cost_units",
)?)
.bind(reservation_expires_at)
.bind(usage_policy_cost_i64(
input.retain_until_unix_secs,
"usage policy retain_until",
)?)
.bind(updated_at)
.bind(updated_at)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(ReserveUsagePolicyCostOutcome::Allowed {
reserved_cost_units: target_reserved_cost_units,
additional_reserved_cost_units: target_reserved_cost_units
.saturating_sub(previous_reserved_cost_units),
})
}
async fn reconcile_usage_policy_cost(
&self,
input: ReconcileUsagePolicyCostInput,
) -> Result<Option<StoredUsagePolicyCostReservation>, DataLayerError> {
input.validate()?;
let actual_cost_units =
usage_policy_cost_i64(input.actual_cost_units, "usage policy actual_cost_units")?;
let finalized_at =
usage_policy_cost_i64(input.finalized_at_unix_secs, "usage policy finalized_at")?;
let updated_at = now_unix_secs()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
tx.commit().await.map_sql_err()?;
return Ok(None);
}
let row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL)
.bind(&input.reservation_token)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(row) = row else {
tx.commit().await.map_sql_err()?;
return Ok(None);
};
let mut reservation = usage_policy_cost_reservation_from_mysql_row(&row)?;
if reservation.request_id != input.request_id || reservation.subject_id != input.subject_id
{
// The token selects the row; audit identity must still match before the reservation
// can be finalized.
tx.commit().await.map_sql_err()?;
return Ok(None);
}
if reservation.state == UsagePolicyCostReservationState::Reserved {
sqlx::query(
r#"
UPDATE usage_cost_reservations
SET state = ?,
actual_cost_units = ?,
finalized_at = ?,
updated_at = ?
WHERE reservation_token = ?
AND request_id = ?
AND subject_id = ?
AND state = 'reserved'
"#,
)
.bind(input.terminal_state.as_str())
.bind(actual_cost_units)
.bind(finalized_at)
.bind(updated_at)
.bind(&input.reservation_token)
.bind(&input.request_id)
.bind(&input.subject_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
reservation.state = input.terminal_state;
reservation.actual_cost_units = Some(input.actual_cost_units);
reservation.finalized_at_unix_secs = Some(input.finalized_at_unix_secs);
}
tx.commit().await.map_sql_err()?;
Ok(Some(reservation))
}
async fn cleanup_usage_policy_cost_reservations(
&self,
now_unix_secs: u64,
batch_size: usize,
) -> Result<usize, DataLayerError> {
if batch_size == 0 {
return Ok(0);
}
let now = usage_policy_cost_i64(now_unix_secs, "usage policy cleanup timestamp")?;
let limit = i64::try_from(batch_size).unwrap_or(i64::MAX);
let result = sqlx::query(
r#"
DELETE FROM usage_cost_reservations
WHERE retain_until <= ?
ORDER BY retain_until, reservation_token
LIMIT ?
"#,
)
.bind(now)
.bind(limit)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() as usize)
}
async fn settle_usage(
&self,
input: UsageSettlementInput,
@@ -465,7 +1090,7 @@ LIMIT 1
let wallet_row = if let Some(api_key_id) = api_key_id {
sqlx::query(
r#"
SELECT id, balance, gift_balance, limit_mode
SELECT id, balance, gift_balance, total_consumed, limit_mode
FROM wallets
WHERE api_key_id = ?
LIMIT 1
@@ -486,7 +1111,7 @@ FOR UPDATE
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
sqlx::query(
r#"
SELECT id, balance, gift_balance, limit_mode
SELECT id, balance, gift_balance, total_consumed, limit_mode
FROM wallets
WHERE user_id = ?
LIMIT 1
@@ -507,14 +1132,20 @@ FOR UPDATE
let wallet_can_overdraft = wallet_row.is_some();
let wallet_available_usd = match wallet_row.as_ref() {
Some(row) => {
let recharge_balance: f64 = row.try_get("balance").map_sql_err()?;
let gift_balance: f64 = row.try_get("gift_balance").map_sql_err()?;
let total_consumed: f64 = row.try_get("total_consumed").map_sql_err()?;
validate_wallet_settlement_values(
recharge_balance,
gift_balance,
total_consumed,
0.0,
)?;
let limit_mode: String = row.try_get("limit_mode").map_sql_err()?;
if limit_mode.eq_ignore_ascii_case("unlimited") {
None
} else {
Some(finite_wallet_available_usd(
row.try_get("balance").map_sql_err()?,
row.try_get("gift_balance").map_sql_err()?,
))
Some(finite_wallet_available_usd(recharge_balance, gift_balance))
}
}
None => Some(0.0),
@@ -593,6 +1224,7 @@ FOR UPDATE
let wallet_id: String = wallet_row.try_get("id").map_sql_err()?;
let before_recharge: f64 = wallet_row.try_get("balance").map_sql_err()?;
let before_gift: f64 = wallet_row.try_get("gift_balance").map_sql_err()?;
let total_consumed: f64 = wallet_row.try_get("total_consumed").map_sql_err()?;
let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?;
let before_total = before_recharge + before_gift;
let mut after_recharge = before_recharge;
@@ -606,6 +1238,13 @@ FOR UPDATE
(after_recharge, after_gift) =
debit_plan.after_balances(before_recharge, before_gift);
}
let total_consumed_after = total_consumed + wallet_debit_cost_usd;
validate_wallet_settlement_values(
after_recharge,
after_gift,
total_consumed_after,
0.0,
)?;
if final_billing_status == "settled" {
sqlx::query(
r#"
@@ -613,14 +1252,14 @@ UPDATE wallets
SET
balance = ?,
gift_balance = ?,
total_consumed = COALESCE(total_consumed, 0) + ?,
total_consumed = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(after_recharge)
.bind(after_gift)
.bind(wallet_debit_cost_usd)
.bind(total_consumed_after)
.bind(updated_at)
.bind(&wallet_id)
.execute(&mut *tx)
@@ -719,10 +1358,11 @@ WHERE id = ?
#[cfg(test)]
mod tests {
use super::MysqlSettlementRepository;
use super::{MysqlSettlementRepository, USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL};
use crate::run_migrations;
use aether_data_contracts::repository::settlement::{
SettlementWriteRepository, UsageSettlementInput,
ReserveUsagePolicyRequestInput, ReserveUsagePolicyRequestOutcome,
SettlementWriteRepository, UsagePolicyRequestWindow, UsageSettlementInput,
};
#[tokio::test]
@@ -736,6 +1376,89 @@ mod tests {
let _repository = MysqlSettlementRepository::new(pool);
}
#[test]
fn request_admission_transactions_use_read_committed() {
assert_eq!(
USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL,
"SET TRANSACTION ISOLATION LEVEL READ COMMITTED"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn cross_subject_same_token_is_allowed_once_without_deadlock_when_url_is_set() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
.ok()
.filter(|value| !value.trim().is_empty())
else {
eprintln!(
"skipping mysql request admission race test because AETHER_TEST_MYSQL_URL is unset"
);
return;
};
let pool = sqlx::mysql::MySqlPoolOptions::new()
.max_connections(2)
.connect(&database_url)
.await
.expect("mysql pool should connect");
run_migrations(&pool)
.await
.expect("mysql migrations should run");
cleanup_request_admission_rows(&pool).await;
sqlx::query(
r#"
INSERT INTO users (id, username, auth_source, created_at, updated_at)
VALUES
('admission-race-user-1', 'admission-race-user-1', 'local', 1, 1),
('admission-race-user-2', 'admission-race-user-2', 'local', 1, 1)
"#,
)
.execute(&pool)
.await
.expect("race users should seed");
let repository = MysqlSettlementRepository::new(pool.clone());
let reserve = |request_id: &str, subject_id: &str| ReserveUsagePolicyRequestInput {
request_id: request_id.to_string(),
subject_id: subject_id.to_string(),
event_token: "admission-race-token".to_string(),
admitted_at_unix_secs: 100,
retain_until_unix_secs: 1_000,
windows: vec![UsagePolicyRequestWindow {
starts_at_unix_secs: 0,
ends_at_unix_secs: 1_000,
limit_requests: 10,
}],
};
let (first, second) = tokio::join!(
repository.reserve_usage_policy_request(reserve(
"admission-race-request-1",
"admission-race-user-1"
)),
repository.reserve_usage_policy_request(reserve(
"admission-race-request-2",
"admission-race-user-2"
))
);
let mut outcomes = vec![
first.expect("first reserve should not deadlock"),
second.expect("second reserve should not deadlock"),
];
outcomes.sort_by_key(|outcome| match outcome {
ReserveUsagePolicyRequestOutcome::Allowed => 0,
ReserveUsagePolicyRequestOutcome::Conflict => 1,
_ => 2,
});
assert_eq!(
outcomes,
vec![
ReserveUsagePolicyRequestOutcome::Allowed,
ReserveUsagePolicyRequestOutcome::Conflict,
]
);
cleanup_request_admission_rows(&pool).await;
}
#[tokio::test]
async fn mysql_repository_settles_once_and_enqueues_provider_delta_when_url_is_set() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
@@ -863,4 +1586,19 @@ WHERE request_id = 'settlement-request-1'
.expect("settlement cleanup should succeed");
}
}
async fn cleanup_request_admission_rows(pool: &sqlx::MySqlPool) {
sqlx::query(
"DELETE FROM usage_request_admissions WHERE event_token = 'admission-race-token'",
)
.execute(pool)
.await
.expect("admission race row cleanup should succeed");
sqlx::query(
"DELETE FROM users WHERE id IN ('admission-race-user-1', 'admission-race-user-2')",
)
.execute(pool)
.await
.expect("admission race user cleanup should succeed");
}
}
+49 -49
View File
@@ -6,7 +6,9 @@ use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use aether_data_contracts::repository::usage::{
strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure,
sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence,
sanitize_usage_request_metadata, usage_can_recover_terminal_failure,
usage_error_category_for_status_code, usage_lifecycle_update_allowed,
usage_request_metadata_client_family, PendingUsageCleanupSummary, StoredRequestUsageAudit,
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardSummary,
StoredUsageUserTotals, UpsertUsageRecord, UsageCleanupExecutionMode, UsageCleanupPreviewCounts,
@@ -437,7 +439,6 @@ ON DUPLICATE KEY UPDATE
const SELECT_STALE_PENDING_USAGE_BATCH_SQL: &str = r#"
SELECT
`usage`.request_id,
`usage`.status,
COALESCE(usage_settlement_snapshots.billing_status, `usage`.billing_status) AS billing_status
FROM `usage`
LEFT JOIN usage_settlement_snapshots
@@ -1046,11 +1047,29 @@ impl UsageWriteRepository for MysqlUsageWriteRepository {
&self,
usage: UpsertUsageRecord,
) -> Result<StoredRequestUsageAudit, DataLayerError> {
let mut usage = strip_deprecated_usage_display_fields(usage);
usage.validate()?;
let prepared_capture = http_capture::prepare_usage_http_capture(&mut usage)?;
// Auxiliary tables may receive only clear tombstones, never request or response content.
let capture_usage = usage.clone();
let mut usage = sanitize_usage_for_persistence(usage);
usage.validate()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
let existing = counters::lock_and_load_usage(&mut tx, &usage.request_id).await?;
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,
) {
let existing = existing.clone();
tx.rollback().await.map_sql_err()?;
return http_capture::hydrate_usage_body_refs(&self.pool, existing).await;
}
}
let recovers_terminal_failure = existing.as_ref().is_some_and(|existing| {
usage_can_recover_terminal_failure(
&existing.status,
@@ -1069,13 +1088,18 @@ impl UsageWriteRepository for MysqlUsageWriteRepository {
}
}
let mut capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
let prepared_capture = http_capture::prepare_usage_http_capture(&mut capture_usage)?;
let capture_update_allowed = recovers_terminal_failure
|| http_capture::capture_update_allowed(existing.as_ref(), &usage.status);
if capture_update_allowed {
http_capture::apply_previous_metadata_tombstones(&mut usage, existing.as_ref());
http_capture::apply_previous_metadata_tombstones(&mut capture_usage, existing.as_ref());
usage.request_metadata =
sanitize_usage_request_metadata(capture_usage.request_metadata.clone());
}
let prepared_snapshots = capture_update_allowed
.then(|| snapshots::from_usage(&usage))
// The control projection preserves safe typed routing and allow-listed billing facts.
.then(|| snapshots::from_usage(&capture_usage))
.transpose()?;
bind_upsert(sqlx::query(UPSERT_USAGE_SQL), &usage)?
.execute(&mut *tx)
@@ -1230,7 +1254,7 @@ SET provider_api_keys.request_count = aggregated.request_count,
&self,
cutoff_unix_secs: u64,
now_unix_secs: u64,
timeout_minutes: u64,
_timeout_minutes: u64,
batch_size: usize,
) -> Result<PendingUsageCleanupSummary, DataLayerError> {
if batch_size == 0 {
@@ -1264,7 +1288,6 @@ SET provider_api_keys.request_count = aggregated.request_count,
.map(|row| {
Ok(StalePendingUsageRow {
request_id: row.try_get("request_id").map_sql_err()?,
status: row.try_get("status").map_sql_err()?,
billing_status: row.try_get("billing_status").map_sql_err()?,
})
})
@@ -1280,7 +1303,8 @@ SET provider_api_keys.request_count = aggregated.request_count,
UPDATE `usage`
SET status = 'completed',
status_code = 200,
error_message = NULL
error_message = NULL,
error_category = NULL
WHERE request_id = ?
"#,
)
@@ -1308,11 +1332,8 @@ WHERE request_id = ?
let candidate_info =
latest_failed_candidate_mysql(&mut tx, &row.request_id).await?;
let (status_code, error_message) = resolve_stale_pending_failure(
candidate_info.as_ref(),
&row.status,
timeout_minutes,
);
let status_code = resolve_stale_pending_status_code(candidate_info.as_ref());
let error_category = usage_error_category_for_status_code(status_code);
let status_code_i64 = i64::from(status_code);
if row.billing_status == "pending" {
sqlx::query(
@@ -1320,7 +1341,8 @@ WHERE request_id = ?
UPDATE `usage`
SET status = 'failed',
status_code = ?,
error_message = ?,
error_message = NULL,
error_category = ?,
billing_status = 'void',
finalized_at = ?,
total_cost_usd = 0,
@@ -1329,7 +1351,7 @@ WHERE request_id = ?
"#,
)
.bind(status_code_i64)
.bind(&error_message)
.bind(error_category)
.bind(to_i64(now_unix_secs, "usage finalized_at")?)
.bind(&row.request_id)
.execute(&mut *tx)
@@ -1347,12 +1369,13 @@ WHERE request_id = ?
UPDATE `usage`
SET status = 'failed',
status_code = ?,
error_message = ?
error_message = NULL,
error_category = ?
WHERE request_id = ?
"#,
)
.bind(status_code_i64)
.bind(&error_message)
.bind(error_category)
.bind(&row.request_id)
.execute(&mut *tx)
.await
@@ -1364,7 +1387,8 @@ WHERE request_id = ?
UPDATE request_candidates
SET status = 'failed',
finished_at = ?,
error_message = '请求超时(服务器可能已重启)'
error_type = 'internal',
error_message = NULL
WHERE request_id = ?
AND status IN ('pending', 'streaming')
"#,
@@ -1451,7 +1475,6 @@ WHERE request_id = ?
struct StalePendingUsageRow {
request_id: String,
status: String,
billing_status: String,
}
@@ -1561,29 +1584,14 @@ ON DUPLICATE KEY UPDATE
Ok(())
}
fn stale_pending_error_message(status: &str, timeout_minutes: u64) -> String {
format!("请求超时: 状态 '{status}' 超过 {timeout_minutes} 分钟未完成")
}
struct FailedCandidateCleanupInfo {
status_code: Option<u16>,
error_message: Option<String>,
}
fn resolve_stale_pending_failure(
candidate: Option<&FailedCandidateCleanupInfo>,
status: &str,
timeout_minutes: u64,
) -> (u16, String) {
match candidate {
Some(info) => (
info.status_code.unwrap_or(502),
info.error_message
.clone()
.unwrap_or_else(|| stale_pending_error_message(status, timeout_minutes)),
),
None => (504, stale_pending_error_message(status, timeout_minutes)),
}
fn resolve_stale_pending_status_code(candidate: Option<&FailedCandidateCleanupInfo>) -> u16 {
candidate
.and_then(|info| info.status_code)
.unwrap_or(if candidate.is_some() { 502 } else { 504 })
}
async fn latest_failed_candidate_mysql(
@@ -1592,7 +1600,7 @@ async fn latest_failed_candidate_mysql(
) -> Result<Option<FailedCandidateCleanupInfo>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT status_code, error_message
SELECT status_code
FROM request_candidates
WHERE request_id = ?
AND status IN ('failed', 'cancelled')
@@ -1615,15 +1623,7 @@ LIMIT 1
.try_get::<Option<i64>, _>("status_code")
.map_sql_err()?
.and_then(|value| u16::try_from(value).ok());
let error_message = row
.try_get::<Option<String>, _>("error_message")
.map_sql_err()?
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
Ok(Some(FailedCandidateCleanupInfo {
status_code,
error_message,
}))
Ok(Some(FailedCandidateCleanupInfo { status_code }))
}
fn bind_upsert<'q>(
@@ -1,11 +1,8 @@
use std::io::Write;
use aether_data_contracts::repository::usage::{
parse_usage_body_ref, usage_body_ref, UsageBodyField, UsageCleanupExecutionMode,
UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow,
UsageCleanupExecutionMode, UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets,
UsageCleanupWindow,
};
use chrono::{DateTime, Utc};
use flate2::{write::GzEncoder, Compression};
use serde_json::Value;
use sqlx::Row;
use tracing::warn;
@@ -66,17 +63,6 @@ OR EXISTS (
)
"#;
const INLINE_OR_COMPRESSED_BODY_PREDICATE: &str = r#"
request_body IS NOT NULL
OR response_body IS NOT NULL
OR provider_request_body IS NOT NULL
OR client_response_body IS NOT NULL
OR request_body_compressed IS NOT NULL
OR response_body_compressed IS NOT NULL
OR provider_request_body_compressed IS NOT NULL
OR client_response_body_compressed IS NOT NULL
"#;
const HEADER_PREDICATE: &str = r#"
request_headers IS NOT NULL
OR response_headers IS NOT NULL
@@ -116,6 +102,20 @@ OR request_body_compressed IS NOT NULL
OR response_body_compressed IS NOT NULL
OR provider_request_body_compressed IS NOT NULL
OR client_response_body_compressed IS NOT NULL
OR EXISTS (
SELECT 1 FROM usage_body_blobs
WHERE usage_body_blobs.request_id = `usage`.request_id
)
OR EXISTS (
SELECT 1 FROM usage_http_audits
WHERE usage_http_audits.request_id = `usage`.request_id
AND (
usage_http_audits.request_body_ref IS NOT NULL
OR usage_http_audits.provider_request_body_ref IS NOT NULL
OR usage_http_audits.response_body_ref IS NOT NULL
OR usage_http_audits.client_response_body_ref IS NOT NULL
)
)
OR (
request_metadata IS NOT NULL
AND JSON_VALID(request_metadata)
@@ -136,43 +136,6 @@ struct CleanupRow {
request_id: String,
}
#[derive(Debug)]
struct BodyRow {
id: String,
request_id: String,
request_body: Option<Value>,
request_body_compressed: Option<Vec<u8>>,
provider_request_body: Option<Value>,
provider_request_body_compressed: Option<Vec<u8>>,
response_body: Option<Value>,
response_body_compressed: Option<Vec<u8>>,
client_response_body: Option<Value>,
client_response_body_compressed: Option<Vec<u8>>,
}
#[derive(Debug, Default)]
struct DetachedRefs {
request_body_ref: Option<String>,
provider_request_body_ref: Option<String>,
response_body_ref: Option<String>,
client_response_body_ref: Option<String>,
}
impl DetachedRefs {
fn any_present(&self) -> bool {
self.request_body_ref.is_some()
|| self.provider_request_body_ref.is_some()
|| self.response_body_ref.is_some()
|| self.client_response_body_ref.is_some()
}
}
struct DetachedBlob {
body_ref: String,
body_field: &'static str,
payload_gzip: Vec<u8>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BodyCleanupKind {
Raw,
@@ -188,10 +151,6 @@ impl BodyCleanupKind {
Self::All => ALL_BODY_PREDICATE,
}
}
fn clears_detached(self) -> bool {
self != Self::Raw
}
}
pub(crate) async fn cleanup_usage(
@@ -263,12 +222,19 @@ pub(crate) async fn cleanup_usage(
};
let detail_newer_than = detail_body_newer_than(window, targets);
let legacy_body_refs_migrated = if targets.detail_body {
migrate_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await?
purge_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await?
} else {
0
};
let body_externalized = if targets.detail_body {
externalize_detail_bodies(pool, window.detail_cutoff, detail_newer_than, batch_size).await?
cleanup_body_fields(
pool,
window.detail_cutoff,
detail_newer_than,
batch_size,
BodyCleanupKind::All,
)
.await?
} else {
0
};
@@ -291,6 +257,8 @@ pub(crate) async fn cleanup_usage(
header_cleaned,
keys_cleaned,
records_deleted,
cost_reservations_deleted: 0,
request_admissions_deleted: 0,
})
}
@@ -611,14 +579,13 @@ WHERE id = ?
.execute(&mut *tx)
.await
.map_sql_err()?;
if kind.clears_detached() {
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
.bind(&row.request_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
sqlx::query(
r#"
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
.bind(&row.request_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
sqlx::query(
r#"
UPDATE usage_http_audits
SET request_body_ref = NULL,
provider_request_body_ref = NULL,
@@ -628,13 +595,12 @@ SET request_body_ref = NULL,
updated_at = UNIX_TIMESTAMP()
WHERE request_id = ?
"#,
)
.bind(&row.request_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
delete_empty_http_audit(&mut tx, &row.request_id).await?;
}
)
.bind(&row.request_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
delete_empty_http_audit(&mut tx, &row.request_id).await?;
}
tx.commit().await.map_sql_err()?;
total = total.saturating_add(row_count);
@@ -670,14 +636,14 @@ WHERE request_id = ?
Ok(())
}
async fn migrate_legacy_body_refs(
async fn purge_legacy_body_refs(
pool: &MysqlPool,
cutoff: DateTime<Utc>,
newer_than: Option<DateTime<Utc>>,
batch_size: usize,
) -> Result<usize, DataLayerError> {
if invalid_window(cutoff, newer_than) {
warn!(%cutoff, ?newer_than, "MySQL usage legacy body-ref migration skipped due to invalid window");
warn!(%cutoff, ?newer_than, "MySQL usage legacy body-ref purge skipped due to invalid window");
return Ok(0);
}
let mut total = 0usize;
@@ -695,7 +661,7 @@ async fn migrate_legacy_body_refs(
}
let row_count = rows.len();
let mut tx = pool.begin().await.map_sql_err()?;
let mut migrated = 0usize;
let mut purged = 0usize;
for row in rows {
let metadata: Option<String> =
sqlx::query_scalar("SELECT request_metadata FROM `usage` WHERE id = ? LIMIT 1")
@@ -704,14 +670,9 @@ async fn migrate_legacy_body_refs(
.await
.map_sql_err()?
.flatten();
let Some((refs, metadata)) =
legacy_body_ref_plan(&row.request_id, metadata.as_deref())?
else {
let Some(metadata) = legacy_body_ref_purge_plan(metadata.as_deref())? else {
continue;
};
if refs.any_present() {
upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?;
}
let updated = sqlx::query(
r#"
UPDATE `usage`
@@ -726,23 +687,23 @@ WHERE id = ?
.await
.map_sql_err()?
.rows_affected();
purge_detached_body_capture(&mut tx, &row.request_id).await?;
if updated > 0 {
migrated += 1;
purged += 1;
}
}
tx.commit().await.map_sql_err()?;
total = total.saturating_add(migrated);
if row_count < batch_size || migrated == 0 {
total = total.saturating_add(purged);
if row_count < batch_size || purged == 0 {
break;
}
}
Ok(total)
}
fn legacy_body_ref_plan(
request_id: &str,
fn legacy_body_ref_purge_plan(
metadata: Option<&str>,
) -> Result<Option<(DetachedRefs, Option<String>)>, DataLayerError> {
) -> Result<Option<Option<String>>, DataLayerError> {
let Some(metadata) = metadata else {
return Ok(None);
};
@@ -752,30 +713,16 @@ fn legacy_body_ref_plan(
let Value::Object(mut object) = value else {
return Ok(None);
};
let mut refs = DetachedRefs::default();
let mut removed = false;
for field in [
UsageBodyField::RequestBody,
UsageBodyField::ProviderRequestBody,
UsageBodyField::ResponseBody,
UsageBodyField::ClientResponseBody,
for key in [
"request_body_ref",
"provider_request_body_ref",
"response_body_ref",
"client_response_body_ref",
] {
let Some(value) = object.remove(field.as_ref_key()) else {
continue;
};
removed = true;
let parsed = 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)
});
set_ref(&mut refs, field, parsed);
if object.remove(key).is_some() {
removed = true;
}
}
if !removed {
return Ok(None);
@@ -791,277 +738,35 @@ fn legacy_body_ref_plan(
})?,
)
};
Ok(Some((refs, metadata)))
Ok(Some(metadata))
}
async fn externalize_detail_bodies(
pool: &MysqlPool,
cutoff: DateTime<Utc>,
newer_than: Option<DateTime<Utc>>,
batch_size: usize,
) -> Result<usize, DataLayerError> {
if invalid_window(cutoff, newer_than) {
warn!(%cutoff, ?newer_than, "MySQL usage body externalization skipped due to invalid window");
return Ok(0);
}
let batch_size = batch_size.clamp(1, 25);
let mut total = 0usize;
loop {
let rows = fetch_body_rows(pool, cutoff, newer_than, batch_size).await?;
if rows.is_empty() {
break;
}
let row_count = rows.len();
let mut externalized = 0usize;
for row in rows {
let (blobs, refs) = build_detached_bodies(&row)?;
let mut tx = pool.begin().await.map_sql_err()?;
for blob in blobs {
sqlx::query(
r#"
INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip)
VALUES (?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
request_id = VALUES(request_id),
body_field = VALUES(body_field),
payload_gzip = VALUES(payload_gzip),
updated_at = UNIX_TIMESTAMP()
"#,
)
.bind(blob.body_ref)
.bind(&row.request_id)
.bind(blob.body_field)
.bind(blob.payload_gzip)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
if refs.any_present() {
upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?;
}
let updated = sqlx::query(
r#"
UPDATE `usage`
SET request_body = NULL,
response_body = NULL,
provider_request_body = NULL,
client_response_body = NULL,
request_body_compressed = NULL,
response_body_compressed = NULL,
provider_request_body_compressed = NULL,
client_response_body_compressed = NULL
WHERE id = ?
"#,
)
.bind(row.id)
.execute(&mut *tx)
.await
.map_sql_err()?
.rows_affected();
tx.commit().await.map_sql_err()?;
if updated > 0 {
externalized += 1;
}
}
total = total.saturating_add(externalized);
if row_count < batch_size || externalized == 0 {
break;
}
}
Ok(total)
}
async fn fetch_body_rows(
pool: &MysqlPool,
cutoff: DateTime<Utc>,
newer_than: Option<DateTime<Utc>>,
batch_size: usize,
) -> Result<Vec<BodyRow>, DataLayerError> {
let newer_than = newer_than.map(|value| value.timestamp());
let sql = format!(
r#"
SELECT id,
request_id,
CAST(request_body AS CHAR) AS request_body,
request_body_compressed,
CAST(provider_request_body AS CHAR) AS provider_request_body,
provider_request_body_compressed,
CAST(response_body AS CHAR) AS response_body,
response_body_compressed,
CAST(client_response_body AS CHAR) AS client_response_body,
client_response_body_compressed
FROM `usage`
WHERE created_at_unix_ms < ?
AND (? IS NULL OR created_at_unix_ms >= ?)
AND ({INLINE_OR_COMPRESSED_BODY_PREDICATE})
ORDER BY created_at_unix_ms ASC, id ASC
LIMIT ?
"#
);
sqlx::query(&sql)
.bind(cutoff.timestamp())
.bind(newer_than)
.bind(newer_than)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.fetch_all(pool)
.await
.map_sql_err()?
.into_iter()
.map(|row| {
Ok(BodyRow {
id: row.try_get("id").map_sql_err()?,
request_id: row.try_get("request_id").map_sql_err()?,
request_body: parse_optional_json(row.try_get("request_body").map_sql_err()?)?,
request_body_compressed: row.try_get("request_body_compressed").map_sql_err()?,
provider_request_body: parse_optional_json(
row.try_get("provider_request_body").map_sql_err()?,
)?,
provider_request_body_compressed: row
.try_get("provider_request_body_compressed")
.map_sql_err()?,
response_body: parse_optional_json(row.try_get("response_body").map_sql_err()?)?,
response_body_compressed: row.try_get("response_body_compressed").map_sql_err()?,
client_response_body: parse_optional_json(
row.try_get("client_response_body").map_sql_err()?,
)?,
client_response_body_compressed: row
.try_get("client_response_body_compressed")
.map_sql_err()?,
})
})
.collect()
}
fn parse_optional_json(raw: Option<String>) -> Result<Option<Value>, DataLayerError> {
raw.map(|raw| {
serde_json::from_str(&raw).map_err(|err| {
DataLayerError::UnexpectedValue(format!("invalid inline usage body JSON: {err}"))
})
})
.transpose()
}
fn build_detached_bodies(
row: &BodyRow,
) -> Result<(Vec<DetachedBlob>, DetachedRefs), DataLayerError> {
let mut blobs = Vec::new();
let mut refs = DetachedRefs::default();
add_detached_body(
&mut blobs,
&mut refs,
&row.request_id,
UsageBodyField::RequestBody,
row.request_body.as_ref(),
row.request_body_compressed.as_deref(),
)?;
add_detached_body(
&mut blobs,
&mut refs,
&row.request_id,
UsageBodyField::ProviderRequestBody,
row.provider_request_body.as_ref(),
row.provider_request_body_compressed.as_deref(),
)?;
add_detached_body(
&mut blobs,
&mut refs,
&row.request_id,
UsageBodyField::ResponseBody,
row.response_body.as_ref(),
row.response_body_compressed.as_deref(),
)?;
add_detached_body(
&mut blobs,
&mut refs,
&row.request_id,
UsageBodyField::ClientResponseBody,
row.client_response_body.as_ref(),
row.client_response_body_compressed.as_deref(),
)?;
Ok((blobs, refs))
}
fn add_detached_body(
blobs: &mut Vec<DetachedBlob>,
refs: &mut DetachedRefs,
request_id: &str,
field: UsageBodyField,
raw: Option<&Value>,
compressed: Option<&[u8]>,
) -> Result<(), DataLayerError> {
let payload_gzip = match raw {
Some(value) => Some(compress_json(value)?),
None => compressed.map(ToOwned::to_owned),
};
let Some(payload_gzip) = payload_gzip else {
return Ok(());
};
let body_ref = usage_body_ref(request_id, field);
blobs.push(DetachedBlob {
body_ref: body_ref.clone(),
body_field: field.as_storage_field(),
payload_gzip,
});
set_ref(refs, field, Some(body_ref));
Ok(())
}
fn compress_json(value: &Value) -> Result<Vec<u8>, DataLayerError> {
let bytes = serde_json::to_vec(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}"))
})?;
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6));
encoder.write_all(&bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}"))
})?;
encoder.finish().map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}"))
})
}
fn set_ref(refs: &mut DetachedRefs, field: UsageBodyField, value: Option<String>) {
match field {
UsageBodyField::RequestBody => refs.request_body_ref = value,
UsageBodyField::ProviderRequestBody => refs.provider_request_body_ref = value,
UsageBodyField::ResponseBody => refs.response_body_ref = value,
UsageBodyField::ClientResponseBody => refs.client_response_body_ref = value,
}
}
async fn upsert_http_audit_refs(
async fn purge_detached_body_capture(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
request_id: &str,
refs: &DetachedRefs,
) -> Result<(), DataLayerError> {
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
.bind(request_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
sqlx::query(
r#"
INSERT INTO usage_http_audits (
request_id,
request_body_ref,
provider_request_body_ref,
response_body_ref,
client_response_body_ref,
body_capture_mode
)
VALUES (?, ?, ?, ?, ?, 'ref_backed')
ON DUPLICATE KEY UPDATE
request_body_ref = COALESCE(VALUES(request_body_ref), request_body_ref),
provider_request_body_ref = COALESCE(VALUES(provider_request_body_ref), provider_request_body_ref),
response_body_ref = COALESCE(VALUES(response_body_ref), response_body_ref),
client_response_body_ref = COALESCE(VALUES(client_response_body_ref), client_response_body_ref),
body_capture_mode = 'ref_backed',
updated_at = UNIX_TIMESTAMP()
UPDATE usage_http_audits
SET request_body_ref = NULL,
provider_request_body_ref = NULL,
response_body_ref = NULL,
client_response_body_ref = NULL,
body_capture_mode = 'none',
updated_at = UNIX_TIMESTAMP()
WHERE request_id = ?
"#,
)
.bind(request_id)
.bind(refs.request_body_ref.as_deref())
.bind(refs.provider_request_body_ref.as_deref())
.bind(refs.response_body_ref.as_deref())
.bind(refs.client_response_body_ref.as_deref())
.execute(&mut **tx)
.await
.map_sql_err()?;
Ok(())
delete_empty_http_audit(tx, request_id).await
}
async fn cleanup_expired_api_keys(
@@ -1127,29 +832,21 @@ ORDER BY expires_at ASC, id ASC
#[cfg(test)]
mod tests {
use std::io::Read;
use flate2::read::GzDecoder;
use serde_json::json;
use super::{compress_json, legacy_body_ref_plan};
use super::{legacy_body_ref_purge_plan, DETAIL_BODY_PREDICATE};
#[test]
fn mysql_cleanup_legacy_body_ref_plan_preserves_unrelated_metadata() {
fn mysql_cleanup_legacy_body_ref_purge_preserves_unrelated_metadata() {
let metadata = json!({
"trace": "kept",
"request_body_ref": "usage://request/request-1/request_body",
"response_body_ref": "usage://request/other/response_body"
})
.to_string();
let (refs, metadata) = legacy_body_ref_plan("request-1", Some(&metadata))
let metadata = legacy_body_ref_purge_plan(Some(&metadata))
.expect("legacy plan should build")
.expect("legacy refs should be present");
assert_eq!(
refs.request_body_ref.as_deref(),
Some("usage://request/request-1/request_body")
);
assert!(refs.response_body_ref.is_none());
assert_eq!(
serde_json::from_str::<serde_json::Value>(
metadata.as_deref().expect("trace metadata should remain")
@@ -1160,18 +857,9 @@ mod tests {
}
#[test]
fn mysql_cleanup_gzip_payload_round_trips() {
let value = json!({"hello": "world"});
let payload = compress_json(&value).expect("body should compress");
let mut decoder = GzDecoder::new(payload.as_slice());
let mut decoded = Vec::new();
decoder
.read_to_end(&mut decoded)
.expect("body should decompress");
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&decoded)
.expect("decoded body should be JSON"),
value
);
fn mysql_detail_cleanup_includes_detached_capture() {
assert!(DETAIL_BODY_PREDICATE.contains("usage_body_blobs"));
assert!(DETAIL_BODY_PREDICATE.contains("usage_http_audits"));
assert!(!DETAIL_BODY_PREDICATE.contains("payload_gzip"));
}
}
@@ -24,6 +24,7 @@ SELECT
id,
kind,
target_id,
target_tunnel_generation,
request_count_delta,
total_requests_delta,
success_count_delta,
@@ -49,6 +50,7 @@ struct DeltaRow {
id: String,
kind: String,
target_id: String,
target_tunnel_generation: Option<String>,
request_count_delta: i64,
total_requests_delta: i64,
success_count_delta: i64,
@@ -71,7 +73,10 @@ struct Aggregates {
provider_api_keys: BTreeMap<String, ProviderApiKeyUsageDelta>,
models: BTreeMap<String, ModelUsageDelta>,
provider_monthly: BTreeMap<String, f64>,
proxy_nodes: BTreeMap<String, ProxyNodeCounterDelta>,
// Keep the node incarnation in the aggregation key. A node id can be
// reused after deletion, so a bare id would route old deltas to the new
// node.
proxy_nodes: BTreeMap<(String, String), ProxyNodeCounterDelta>,
management_tokens: BTreeMap<String, ManagementTokenCounterDelta>,
api_key_last_used: BTreeMap<String, ApiKeyLastUsedDelta>,
}
@@ -142,16 +147,28 @@ impl Aggregates {
.or_default() += row.total_cost_usd_delta;
}
KIND_PROXY_NODE => {
let entry = aggregates
.proxy_nodes
.entry(row.target_id.clone())
.or_insert(ProxyNodeCounterDelta {
let Some(tunnel_generation) = row
.target_tunnel_generation
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
else {
// Legacy rows have no identity fence. Mark them
// processed without applying them to any node.
continue;
};
let aggregate_key = (row.target_id.clone(), tunnel_generation.clone());
let entry = aggregates.proxy_nodes.entry(aggregate_key).or_insert(
ProxyNodeCounterDelta {
node_id: row.target_id.clone(),
expected_tunnel_generation: Some(tunnel_generation),
total_requests_delta: 0,
failed_requests_delta: 0,
dns_failures_delta: 0,
stream_errors_delta: 0,
});
},
);
entry.total_requests_delta += row.total_requests_delta;
entry.failed_requests_delta += row.error_count_delta;
entry.dns_failures_delta += row.dns_failures_delta;
@@ -241,8 +258,8 @@ pub(super) async fn flush(
for (target_id, delta) in &aggregates.provider_monthly {
apply_provider_monthly(&mut tx, target_id, *delta).await?;
}
for (target_id, delta) in &aggregates.proxy_nodes {
apply_proxy_node(&mut tx, target_id, delta).await?;
for ((target_id, tunnel_generation), delta) in &aggregates.proxy_nodes {
apply_proxy_node(&mut tx, target_id, tunnel_generation, delta).await?;
}
for (target_id, delta) in &aggregates.management_tokens {
apply_management_token(&mut tx, target_id, delta).await?;
@@ -283,9 +300,42 @@ pub(super) async fn enqueue_proxy_node(
if delta.is_noop() {
return Ok(false);
}
let Some(expected_tunnel_generation) = delta
.expected_tunnel_generation
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| value.len() <= 64)
.map(ToOwned::to_owned)
else {
// A bare id is not an identity fence. Reject it instead of rebinding
// the delta to whichever incarnation currently owns that id.
return Ok(false);
};
let node_id = delta.node_id.trim().to_string();
let request_id = format!("proxy_node:{node_id}:{}", uuid::Uuid::new_v4());
let mut tx = pool.begin().await.map_sql_err()?;
// Keep the parent lookup lock-free because flush claims outbox rows before
// updating proxy_nodes. The generation is stored in the outbox row and is
// checked again by flush, so a concurrent id reuse can only discard this
// delta, never apply it to the replacement row.
let tunnel_generation: Option<String> = sqlx::query_scalar(
"SELECT tunnel_generation FROM proxy_nodes WHERE id = ? AND BINARY tunnel_generation = BINARY ? LIMIT 1",
)
.bind(&node_id)
.bind(&expected_tunnel_generation)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(_tunnel_generation) = tunnel_generation
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
else {
tx.rollback().await.map_sql_err()?;
return Ok(false);
};
insert_delta(
&mut tx,
DeltaInsert {
@@ -296,6 +346,7 @@ pub(super) async fn enqueue_proxy_node(
error_count_delta: delta.failed_requests_delta,
dns_failures_delta: delta.dns_failures_delta,
stream_errors_delta: delta.stream_errors_delta,
target_tunnel_generation: Some(&expected_tunnel_generation),
..DeltaInsert::default()
},
)
@@ -731,6 +782,7 @@ struct DeltaInsert<'a> {
request_id: &'a str,
kind: &'a str,
target_id: &'a str,
target_tunnel_generation: Option<&'a str>,
request_count_delta: i64,
total_requests_delta: i64,
success_count_delta: i64,
@@ -759,18 +811,20 @@ async fn insert_delta(
sqlx::query(
r#"
INSERT INTO usage_counter_deltas (
id, request_id, kind, target_id, request_count_delta, total_requests_delta,
id, request_id, kind, target_id, target_tunnel_generation,
request_count_delta, total_requests_delta,
success_count_delta, error_count_delta, dns_failures_delta, stream_errors_delta,
total_tokens_delta, total_cost_usd_delta, total_response_time_ms_delta,
last_used_at_unix_secs, last_used_ip, candidate_last_used_at_unix_secs,
removed_last_used_at_unix_secs, usage_created_at_unix_secs, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
.bind(request_id)
.bind(input.kind)
.bind(target_id)
.bind(input.target_tunnel_generation)
.bind(input.request_count_delta)
.bind(input.total_requests_delta)
.bind(input.success_count_delta)
@@ -814,6 +868,7 @@ fn map_row(row: &sqlx::mysql::MySqlRow) -> Result<DeltaRow, DataLayerError> {
id: row.try_get("id").map_sql_err()?,
kind: row.try_get("kind").map_sql_err()?,
target_id: row.try_get("target_id").map_sql_err()?,
target_tunnel_generation: row.try_get("target_tunnel_generation").map_sql_err()?,
request_count_delta: row.try_get("request_count_delta").map_sql_err()?,
total_requests_delta: row.try_get("total_requests_delta").map_sql_err()?,
success_count_delta: row.try_get("success_count_delta").map_sql_err()?,
@@ -997,9 +1052,10 @@ async fn apply_provider_monthly(
async fn apply_proxy_node(
tx: &mut sqlx::Transaction<'_, MySql>,
target_id: &str,
tunnel_generation: &str,
delta: &ProxyNodeCounterDelta,
) -> Result<(), DataLayerError> {
if target_id.trim().is_empty() || delta.is_noop() {
if target_id.trim().is_empty() || tunnel_generation.trim().is_empty() || delta.is_noop() {
return Ok(());
}
sqlx::query(
@@ -1010,7 +1066,7 @@ SET total_requests = total_requests + GREATEST(?, 0),
dns_failures = dns_failures + GREATEST(?, 0),
stream_errors = stream_errors + GREATEST(?, 0),
updated_at = ?
WHERE id = ?
WHERE id = ? AND BINARY tunnel_generation = BINARY ?
"#,
)
.bind(delta.total_requests_delta)
@@ -1019,6 +1075,7 @@ WHERE id = ?
.bind(delta.stream_errors_delta)
.bind(current_unix_secs())
.bind(target_id)
.bind(tunnel_generation)
.execute(&mut **tx)
.await
.map_sql_err()?;
@@ -1,10 +1,9 @@
use std::io::{Read, Write};
use aether_data_contracts::repository::usage::{
parse_usage_body_ref, usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord,
UsageBodyCaptureState, UsageBodyField,
canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json,
usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState,
UsageBodyField,
};
use flate2::{read::GzDecoder, write::GzEncoder, Compression};
use flate2::read::GzDecoder;
use serde_json::{Map, Value};
use sqlx::{mysql::MySqlRow, Row};
@@ -28,9 +27,7 @@ pub(crate) struct PreparedUsageHttpCapture {
#[derive(Debug)]
struct PreparedBody {
field: UsageBodyField,
payload_gzip: Option<Vec<u8>>,
clear_existing: bool,
}
#[derive(Debug, Default)]
@@ -58,15 +55,6 @@ struct HttpAuditStates {
client_response_body_state: Option<UsageBodyCaptureState>,
}
impl HttpAuditStates {
fn any_present(&self) -> bool {
self.request_body_state.is_some()
|| self.provider_request_body_state.is_some()
|| self.response_body_state.is_some()
|| self.client_response_body_state.is_some()
}
}
pub(crate) fn capture_update_allowed(
previous: Option<&StoredRequestUsageAudit>,
incoming_status: &str,
@@ -141,26 +129,10 @@ pub(crate) fn prepare_usage_http_capture(
.then_some(usage.client_response_body.as_ref())
.flatten();
let request_body = prepare_body(
UsageBodyField::RequestBody,
request_body_value,
clear_request,
)?;
let provider_request_body = prepare_body(
UsageBodyField::ProviderRequestBody,
provider_request_body_value,
clear_provider_request,
)?;
let response_body = prepare_body(
UsageBodyField::ResponseBody,
response_body_value,
clear_response,
)?;
let client_response_body = prepare_body(
UsageBodyField::ClientResponseBody,
client_response_body_value,
clear_client_response,
)?;
let request_body = prepare_body(request_body_value)?;
let provider_request_body = prepare_body(provider_request_body_value)?;
let response_body = prepare_body(response_body_value)?;
let client_response_body = prepare_body(client_response_body_value)?;
let refs = HttpAuditRefs {
request_body_ref: resolved_write_ref(
@@ -276,29 +248,13 @@ pub(crate) fn prepare_usage_http_capture(
})
}
fn prepare_body(
field: UsageBodyField,
value: Option<&Value>,
clear_existing: bool,
) -> Result<PreparedBody, DataLayerError> {
Ok(PreparedBody {
field,
payload_gzip: value.map(compress_json).transpose()?,
clear_existing,
})
}
fn compress_json(value: &Value) -> Result<Vec<u8>, DataLayerError> {
let bytes = serde_json::to_vec(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}"))
})?;
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6));
encoder.write_all(&bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}"))
})?;
encoder.finish().map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}"))
})
fn prepare_body(value: Option<&Value>) -> Result<PreparedBody, DataLayerError> {
if value.is_some() {
return Err(DataLayerError::InvalidInput(
"usage body persistence is disabled".to_string(),
));
}
Ok(PreparedBody { payload_gzip: None })
}
fn resolved_write_ref(
@@ -308,9 +264,7 @@ fn resolved_write_ref(
has_blob: bool,
) -> Option<String> {
explicit_ref
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field))
.or_else(|| has_blob.then(|| usage_body_ref(request_id, field)))
}
@@ -378,23 +332,63 @@ pub(crate) async fn sync_usage_http_capture(
request_id: &str,
prepared: &PreparedUsageHttpCapture,
) -> Result<(), DataLayerError> {
for body in [
let bodies = [
&prepared.request_body,
&prepared.provider_request_body,
&prepared.response_body,
&prepared.client_response_body,
] {
sync_body(tx, request_id, body).await?;
];
let contains_capture = prepared.request_headers.is_some()
|| prepared.provider_request_headers.is_some()
|| prepared.response_headers.is_some()
|| prepared.client_response_headers.is_some()
|| prepared.refs.any_present()
|| bodies.iter().any(|body| body.payload_gzip.is_some())
|| prepared.capture_mode != "none";
if contains_capture {
return Err(DataLayerError::InvalidInput(
"usage HTTP capture persistence is disabled".to_string(),
));
}
sqlx::query("DELETE FROM usage_http_audits WHERE request_id = ?")
.bind(request_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
.bind(request_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
sqlx::query(
r#"
UPDATE `usage`
SET request_headers = NULL,
request_body = NULL,
provider_request_headers = NULL,
provider_request_body = NULL,
response_headers = NULL,
response_body = NULL,
client_response_headers = NULL,
client_response_body = NULL,
request_body_compressed = NULL,
provider_request_body_compressed = NULL,
response_body_compressed = NULL,
client_response_body_compressed = NULL
WHERE request_id = ?
"#,
)
.bind(request_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
let headers_present = prepared.request_headers.is_some()
|| prepared.provider_request_headers.is_some()
|| prepared.response_headers.is_some()
|| prepared.client_response_headers.is_some();
if !headers_present
&& !prepared.refs.any_present()
&& !prepared.states.any_present()
&& prepared.capture_mode == "none"
{
if !headers_present && !prepared.refs.any_present() {
return Ok(());
}
@@ -512,67 +506,6 @@ ON DUPLICATE KEY UPDATE
Ok(())
}
async fn sync_body(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
request_id: &str,
body: &PreparedBody,
) -> Result<(), DataLayerError> {
let body_ref = usage_body_ref(request_id, body.field);
if body.clear_existing || body.payload_gzip.is_some() {
sqlx::query(clear_legacy_body_sql(body.field))
.bind(request_id)
.execute(&mut **tx)
.await
.map_sql_err()?;
}
if body.clear_existing {
sqlx::query("DELETE FROM usage_body_blobs WHERE body_ref = ?")
.bind(body_ref)
.execute(&mut **tx)
.await
.map_sql_err()?;
return Ok(());
}
if let Some(payload_gzip) = body.payload_gzip.as_deref() {
sqlx::query(
r#"
INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip)
VALUES (?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
request_id = VALUES(request_id),
body_field = VALUES(body_field),
payload_gzip = VALUES(payload_gzip),
updated_at = UNIX_TIMESTAMP()
"#,
)
.bind(body_ref)
.bind(request_id)
.bind(body.field.as_storage_field())
.bind(payload_gzip)
.execute(&mut **tx)
.await
.map_sql_err()?;
}
Ok(())
}
fn clear_legacy_body_sql(field: UsageBodyField) -> &'static str {
match field {
UsageBodyField::RequestBody => {
"UPDATE `usage` SET request_body = NULL, request_body_compressed = NULL WHERE request_id = ?"
}
UsageBodyField::ProviderRequestBody => {
"UPDATE `usage` SET provider_request_body = NULL, provider_request_body_compressed = NULL WHERE request_id = ?"
}
UsageBodyField::ResponseBody => {
"UPDATE `usage` SET response_body = NULL, response_body_compressed = NULL WHERE request_id = ?"
}
UsageBodyField::ClientResponseBody => {
"UPDATE `usage` SET client_response_body = NULL, client_response_body_compressed = NULL WHERE request_id = ?"
}
}
}
pub(crate) fn hydrate_usage_row(
row: &MySqlRow,
usage: &mut StoredRequestUsageAudit,
@@ -690,8 +623,7 @@ fn resolved_read_ref(
has_compressed: bool,
) -> Option<String> {
audit_ref
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.and_then(|body_ref| canonical_usage_body_ref_for(&body_ref, request_id, field))
.or_else(|| has_compressed.then(|| usage_body_ref(request_id, field)))
.or_else(|| metadata_body_ref(metadata, request_id, field))
}
@@ -704,13 +636,7 @@ fn metadata_body_ref(
metadata
.and_then(|metadata| metadata.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 optional_state(
@@ -752,7 +678,11 @@ pub(crate) async fn hydrate_usage_body_refs(
let Some(body_ref) = usage.body_ref(field) else {
continue;
};
let value = resolve_body_ref(pool, body_ref).await?;
let Some(body_ref) = canonical_usage_body_ref_for(body_ref, &usage.request_id, field)
else {
continue;
};
let value = resolve_body_ref(pool, &body_ref).await?;
match field {
UsageBodyField::RequestBody => usage.request_body = value,
UsageBodyField::ProviderRequestBody => usage.provider_request_body = value,
@@ -767,19 +697,22 @@ pub(crate) async fn resolve_body_ref(
pool: &MysqlPool,
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(payload_gzip) = sqlx::query_scalar::<_, Vec<u8>>(
"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? LIMIT 1",
"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? AND request_id = ? AND body_field = ? LIMIT 1",
)
.bind(body_ref)
.bind(&canonical_ref)
.bind(&request_id)
.bind(field.as_storage_field())
.fetch_optional(pool)
.await
.map_sql_err()?
{
return inflate_json(&payload_gzip).map(Some);
}
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
return Ok(None);
};
let (inline_column, compressed_column) = usage_body_sql_columns(field);
let row = sqlx::query(&format!(
"SELECT CAST({inline_column} AS CHAR) AS inline_body, {compressed_column} AS compressed_body FROM `usage` WHERE request_id = ? LIMIT 1"
@@ -819,11 +752,7 @@ fn usage_body_sql_columns(field: UsageBodyField) -> (&'static str, &'static str)
}
fn inflate_json(bytes: &[u8]) -> Result<Value, DataLayerError> {
let mut decoder = GzDecoder::new(bytes);
let mut decoded = Vec::new();
decoder.read_to_end(&mut decoded).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to decompress usage body: {err}"))
})?;
let decoded = read_decompressed_usage_json(GzDecoder::new(bytes))?;
serde_json::from_slice(&decoded).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to decode usage body JSON: {err}"))
})
@@ -152,6 +152,124 @@ fn mysql_usage_upsert_guards_candidate_identity_metadata_and_routing_from_late_l
.contains("OR (status = 'streaming' AND VALUES(status) = 'pending')"));
}
#[tokio::test]
async fn mysql_stale_terminal_event_is_a_full_transaction_noop_when_url_is_set() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
.ok()
.filter(|value| !value.trim().is_empty())
else {
eprintln!("skipping MySQL stale terminal test because AETHER_TEST_MYSQL_URL is unset");
return;
};
let pool = sqlx::mysql::MySqlPoolOptions::new()
.max_connections(1)
.connect(&database_url)
.await
.expect("mysql test pool should connect");
run_migrations(&pool)
.await
.expect("mysql migrations should run");
let suffix = unique_suffix();
let user_id = format!("stale-user-{suffix}");
let api_key_id = format!("stale-api-key-{suffix}");
let provider_id = format!("stale-provider-{suffix}");
let provider_key_id = format!("stale-provider-key-{suffix}");
let request_id = format!("stale-request-{suffix}");
seed_stats_targets(&pool, &user_id, &api_key_id, &provider_id, &provider_key_id).await;
let repository = MysqlUsageWriteRepository::new(pool.clone());
let mut newer = sample_usage(
&request_id,
&user_id,
&api_key_id,
&provider_id,
&provider_key_id,
"completed",
"pending",
2_000,
);
newer.candidate_id = Some("candidate-new".to_string());
newer.route_kind = Some("route-new".to_string());
repository
.upsert(newer)
.await
.expect("newer terminal usage should upsert");
let counter_rows_before: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = ?")
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("counter rows should count");
let routing_before: (Option<String>, Option<String>) = sqlx::query_as(
"SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?",
)
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("routing snapshot should load");
let settlement_before: (String, Option<f64>) = sqlx::query_as(
"SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?",
)
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("settlement snapshot should load");
let mut stale = sample_usage(
&request_id,
&user_id,
&api_key_id,
&provider_id,
&provider_key_id,
"failed",
"void",
1_999,
);
stale.status_code = Some(503);
stale.total_cost_usd = Some(99.0);
stale.actual_total_cost_usd = Some(98.0);
stale.candidate_id = Some("candidate-stale".to_string());
stale.route_kind = Some("route-stale".to_string());
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_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, 2_000);
let counter_rows_after: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = ?")
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("counter rows should count");
let routing_after: (Option<String>, Option<String>) = sqlx::query_as(
"SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?",
)
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("routing snapshot should load");
let settlement_after: (String, Option<f64>) = sqlx::query_as(
"SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?",
)
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("settlement snapshot should load");
assert_eq!(counter_rows_after, counter_rows_before);
assert_eq!(routing_after, routing_before);
assert_eq!(settlement_after, settlement_before);
}
#[tokio::test]
async fn mysql_usage_write_repository_upserts_and_flushes_counters_when_url_is_set() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
@@ -470,7 +588,7 @@ async fn mysql_concurrent_same_request_upserts_enqueue_counters_once_when_url_is
}
#[tokio::test]
async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
async fn mysql_usage_http_capture_is_not_persisted_when_url_is_set() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
.ok()
.filter(|value| !value.trim().is_empty())
@@ -516,29 +634,17 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
.upsert(rich)
.await
.expect("MySQL canonical capture should upsert");
assert_eq!(
stored.request_headers,
Some(serde_json::json!({"x-client": "one"}))
);
assert_eq!(
stored.request_body,
Some(serde_json::json!({"request": true}))
);
assert_eq!(
stored.request_body_state,
Some(UsageBodyCaptureState::Reference)
);
assert_eq!(
stored.request_body_ref.as_deref(),
Some(format!("usage://request/{request_id}/request_body").as_str())
);
assert!(stored.request_headers.is_none());
assert!(stored.request_body.is_none());
assert!(stored.request_body_state.is_none());
assert!(stored.request_body_ref.is_none());
let blob_count: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = ?")
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("MySQL canonical blobs should count");
assert_eq!(blob_count, 2);
assert_eq!(blob_count, 0);
let legacy_body: Option<String> =
sqlx::query_scalar("SELECT CAST(request_body AS CHAR) FROM `usage` WHERE request_id = ?")
.bind(&request_id)
@@ -561,8 +667,8 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
.upsert(sparse)
.await
.expect("MySQL sparse capture should upsert");
assert_eq!(sparse_stored.request_headers, stored.request_headers);
assert_eq!(sparse_stored.request_body, stored.request_body);
assert!(sparse_stored.request_headers.is_none());
assert!(sparse_stored.request_body.is_none());
let mut clear = sample_usage(
&request_id,
@@ -582,11 +688,8 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
.expect("MySQL explicit none should clear");
assert!(cleared.request_body.is_none());
assert!(cleared.request_body_ref.is_none());
assert_eq!(
cleared.request_body_state,
Some(UsageBodyCaptureState::None)
);
assert_eq!(cleared.provider_request_body, stored.provider_request_body);
assert!(cleared.request_body_state.is_none());
assert!(cleared.provider_request_body.is_none());
}
#[tokio::test]
@@ -992,16 +1095,27 @@ async fn mysql_usage_cleanup_executes_when_url_is_set() {
.await
.expect("MySQL detail cleanup should succeed");
assert!(summary.body_externalized >= 1);
let body_ref: String =
sqlx::query_scalar("SELECT request_body_ref FROM usage_http_audits WHERE request_id = ?")
let stored_body: Option<String> =
sqlx::query_scalar("SELECT CAST(request_body AS CHAR) FROM `usage` WHERE request_id = ?")
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("externalized body ref should load");
assert_eq!(
body_ref,
format!("usage://request/{request_id}/request_body")
);
.expect("purged body should load");
assert!(stored_body.is_none());
let body_blobs: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = ?")
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("purged body blobs should count");
assert_eq!(body_blobs, 0);
let body_audits: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM usage_http_audits WHERE request_id = ?")
.bind(&request_id)
.fetch_one(&pool)
.await
.expect("purged body refs should count");
assert_eq!(body_audits, 0);
let headers_only = UsageCleanupTargets {
detail_body: false,
File diff suppressed because it is too large Load Diff
@@ -72,6 +72,22 @@ impl MysqlVideoTaskRepository {
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_id_for_user(
&self,
id: &str,
user_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE BINARY id = BINARY ? AND BINARY user_id = BINARY ? LIMIT 1"
))
.bind(id)
.bind(user_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_short_id(
&self,
short_id: &str,
@@ -84,13 +100,29 @@ impl MysqlVideoTaskRepository {
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_short_id_for_user(
&self,
short_id: &str,
user_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE BINARY short_id = BINARY ? AND BINARY user_id = BINARY ? LIMIT 1"
))
.bind(short_id)
.bind(user_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_user_external(
&self,
user_id: &str,
external_task_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE user_id = ? AND external_task_id = ? LIMIT 1"
"{VIDEO_TASK_COLUMNS} WHERE BINARY user_id = BINARY ? AND BINARY external_task_id = BINARY ? LIMIT 1"
))
.bind(user_id)
.bind(external_task_id)
@@ -117,6 +149,26 @@ impl VideoTaskReadRepository for MysqlVideoTaskRepository {
}
}
async fn find_for_user(
&self,
key: VideoTaskLookupKey<'_>,
user_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
match key {
VideoTaskLookupKey::Id(id) => self.find_by_id_for_user(id, user_id).await,
VideoTaskLookupKey::ShortId(short_id) => {
self.find_by_short_id_for_user(short_id, user_id).await
}
VideoTaskLookupKey::UserExternal {
user_id: lookup_user_id,
external_task_id,
} if lookup_user_id == user_id => {
self.find_by_user_external(user_id, external_task_id).await
}
VideoTaskLookupKey::UserExternal { .. } => Ok(None),
}
}
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
@@ -258,15 +310,21 @@ impl VideoTaskReadRepository for MysqlVideoTaskRepository {
#[async_trait]
impl VideoTaskWriteRepository for MysqlVideoTaskRepository {
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
async fn upsert(&self, mut task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
task.sanitize_for_persistence();
let id = task.id.clone();
bind_task(sqlx::query(UPSERT_SQL), task, true, false)?
let expected_identity = task.clone();
bind_task(sqlx::query(upsert_sql()), task, true, false)?
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_by_id(&id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("upserted video task missing".into()))
let stored = self.find_by_id(&id).await?.ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"video task {id} conflicts with persisted immutable identity"
))
})?;
stored.ensure_immutable_identity_matches(&expected_identity)?;
Ok(stored)
}
async fn update_if_active(
@@ -368,7 +426,76 @@ FOR UPDATE SKIP LOCKED
}
}
const UPSERT_SQL: &str = r#"
const IMMUTABLE_IDENTITY_MATCH_SQL: &str = r#"BINARY id <=> BINARY VALUES(id)
AND BINARY short_id <=> BINARY VALUES(short_id)
AND BINARY request_id <=> BINARY VALUES(request_id)
AND BINARY user_id <=> BINARY VALUES(user_id)
AND BINARY api_key_id <=> BINARY VALUES(api_key_id)
AND BINARY external_task_id <=> BINARY VALUES(external_task_id)
AND BINARY provider_id <=> BINARY VALUES(provider_id)
AND BINARY endpoint_id <=> BINARY VALUES(endpoint_id)
AND BINARY key_id <=> BINARY VALUES(key_id)
AND BINARY client_api_format <=> BINARY VALUES(client_api_format)
AND BINARY provider_api_format <=> BINARY VALUES(provider_api_format)
AND format_converted <=> VALUES(format_converted)
AND BINARY model <=> BINARY VALUES(model)
AND duration_seconds <=> VALUES(duration_seconds)
AND BINARY resolution <=> BINARY VALUES(resolution)
AND BINARY aspect_ratio <=> BINARY VALUES(aspect_ratio)
AND BINARY size <=> BINARY VALUES(size)"#;
const UPSERT_UPDATE_COLUMNS: &[&str] = &[
"short_id",
"request_id",
"user_id",
"api_key_id",
"username",
"api_key_name",
"external_task_id",
"provider_id",
"endpoint_id",
"key_id",
"client_api_format",
"provider_api_format",
"format_converted",
"model",
"prompt",
"original_request_body",
"duration_seconds",
"resolution",
"aspect_ratio",
"size",
"status",
"progress_percent",
"progress_message",
"retry_count",
"poll_interval_seconds",
"next_poll_at",
"poll_count",
"max_poll_count",
"video_url",
"error_code",
"error_message",
"request_metadata",
"submitted_at",
"completed_at",
"updated_at",
];
fn upsert_sql() -> &'static str {
static SQL: std::sync::OnceLock<String> = std::sync::OnceLock::new();
SQL.get_or_init(|| {
let guarded_updates = UPSERT_UPDATE_COLUMNS
.iter()
.map(|column| {
format!(
" {column} = IF(({IMMUTABLE_IDENTITY_MATCH_SQL}), VALUES({column}), {column})"
)
})
.collect::<Vec<_>>()
.join(",\n");
format!(
r#"
INSERT INTO video_tasks (
id, short_id, request_id, user_id, api_key_id, username, api_key_name,
external_task_id, provider_id, endpoint_id, key_id, client_api_format,
@@ -380,43 +507,11 @@ INSERT INTO video_tasks (
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
short_id = VALUES(short_id),
request_id = VALUES(request_id),
user_id = VALUES(user_id),
api_key_id = VALUES(api_key_id),
username = VALUES(username),
api_key_name = VALUES(api_key_name),
external_task_id = VALUES(external_task_id),
provider_id = VALUES(provider_id),
endpoint_id = VALUES(endpoint_id),
key_id = VALUES(key_id),
client_api_format = VALUES(client_api_format),
provider_api_format = VALUES(provider_api_format),
format_converted = VALUES(format_converted),
model = VALUES(model),
prompt = VALUES(prompt),
original_request_body = VALUES(original_request_body),
duration_seconds = VALUES(duration_seconds),
resolution = VALUES(resolution),
aspect_ratio = VALUES(aspect_ratio),
size = VALUES(size),
status = VALUES(status),
progress_percent = VALUES(progress_percent),
progress_message = VALUES(progress_message),
retry_count = VALUES(retry_count),
poll_interval_seconds = VALUES(poll_interval_seconds),
next_poll_at = VALUES(next_poll_at),
poll_count = VALUES(poll_count),
max_poll_count = VALUES(max_poll_count),
video_url = VALUES(video_url),
error_code = VALUES(error_code),
error_message = VALUES(error_message),
request_metadata = VALUES(request_metadata),
created_at = VALUES(created_at),
submitted_at = VALUES(submitted_at),
completed_at = VALUES(completed_at),
updated_at = VALUES(updated_at)
"#;
{guarded_updates}
"#
)
})
}
const UPDATE_IF_ACTIVE_SQL: &str = r#"
UPDATE video_tasks SET
@@ -452,20 +547,38 @@ UPDATE video_tasks SET
error_code = ?,
error_message = ?,
request_metadata = ?,
created_at = ?,
created_at = COALESCE(created_at, ?),
submitted_at = ?,
completed_at = ?,
updated_at = ?
WHERE id = ?
AND status IN ('pending', 'submitted', 'queued', 'processing')
AND BINARY short_id <=> BINARY ?
AND BINARY request_id <=> BINARY ?
AND BINARY user_id <=> BINARY ?
AND BINARY api_key_id <=> BINARY ?
AND BINARY external_task_id <=> BINARY ?
AND BINARY provider_id <=> BINARY ?
AND BINARY endpoint_id <=> BINARY ?
AND BINARY key_id <=> BINARY ?
AND BINARY client_api_format <=> BINARY ?
AND BINARY provider_api_format <=> BINARY ?
AND format_converted <=> ?
AND BINARY model <=> BINARY ?
AND duration_seconds <=> ?
AND BINARY resolution <=> BINARY ?
AND BINARY aspect_ratio <=> BINARY ?
AND BINARY size <=> BINARY ?
"#;
fn bind_task<'q>(
query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>,
task: UpsertVideoTask,
mut task: UpsertVideoTask,
include_insert_id: bool,
include_update_id: bool,
) -> Result<sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
task.sanitize_for_persistence();
let identity = task.clone();
let original_request_body = json_to_string(&task.original_request_body)?;
let request_metadata = json_to_string(&task.request_metadata)?;
let query = if include_insert_id {
@@ -535,19 +648,45 @@ fn bind_task<'q>(
"video task updated_at",
)?);
if include_update_id {
Ok(bound.bind(task.id))
bind_identity_guard(bound.bind(task.id), identity)
} else {
Ok(bound)
}
}
fn bind_identity_guard<'q>(
query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>,
identity: UpsertVideoTask,
) -> Result<sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
Ok(query
.bind(identity.short_id)
.bind(identity.request_id)
.bind(identity.user_id)
.bind(identity.api_key_id)
.bind(identity.external_task_id)
.bind(identity.provider_id)
.bind(identity.endpoint_id)
.bind(identity.key_id)
.bind(identity.client_api_format)
.bind(identity.provider_api_format)
.bind(identity.format_converted)
.bind(identity.model)
.bind(optional_u32_to_i32(
identity.duration_seconds,
"video task duration_seconds",
)?)
.bind(identity.resolution)
.bind(identity.aspect_ratio)
.bind(identity.size))
}
fn push_filter<'args>(
builder: &mut QueryBuilder<'args, MySql>,
filter: &'args VideoTaskQueryFilter,
created_since_unix_secs: Option<u64>,
) {
if let Some(user_id) = filter.user_id.as_deref() {
push_clause(builder, "user_id = ");
push_clause(builder, "BINARY user_id = BINARY ");
builder.push_bind(user_id);
}
if let Some(status) = filter.status {
@@ -706,13 +845,86 @@ fn optional_u32_to_i32(value: Option<u32>, name: &str) -> Result<Option<i32>, Da
#[cfg(test)]
mod tests {
use super::MysqlVideoTaskRepository;
use super::{
upsert_sql, MysqlVideoTaskRepository, IMMUTABLE_IDENTITY_MATCH_SQL, UPDATE_IF_ACTIVE_SQL,
UPSERT_UPDATE_COLUMNS,
};
use crate::run_migrations;
use aether_data_contracts::repository::video_tasks::{
UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository,
};
use std::sync::Arc;
#[test]
fn mysql_write_sql_atomically_guards_immutable_identity() {
assert!(IMMUTABLE_IDENTITY_MATCH_SQL.contains("BINARY id <=> BINARY VALUES(id)"));
for column in [
"short_id",
"request_id",
"user_id",
"api_key_id",
"external_task_id",
"provider_id",
"endpoint_id",
"key_id",
"client_api_format",
"provider_api_format",
"model",
"resolution",
"aspect_ratio",
"size",
] {
assert!(
IMMUTABLE_IDENTITY_MATCH_SQL
.contains(&format!("BINARY {column} <=> BINARY VALUES({column})")),
"upsert identity predicate should guard {column}"
);
}
for column in ["format_converted", "duration_seconds"] {
assert!(
IMMUTABLE_IDENTITY_MATCH_SQL.contains(&format!("{column} <=> VALUES({column})")),
"upsert identity predicate should guard {column}"
);
}
let upsert = upsert_sql();
assert!(!UPSERT_UPDATE_COLUMNS.contains(&"created_at"));
for column in UPSERT_UPDATE_COLUMNS {
assert!(
upsert.contains(&format!("{column} = IF((BINARY id <=> BINARY VALUES(id)")),
"upsert assignment should be conditional for {column}"
);
}
for column in [
"short_id",
"request_id",
"user_id",
"api_key_id",
"external_task_id",
"provider_id",
"endpoint_id",
"key_id",
"client_api_format",
"provider_api_format",
"model",
"resolution",
"aspect_ratio",
"size",
] {
assert!(
UPDATE_IF_ACTIVE_SQL.contains(&format!("BINARY {column} <=> BINARY ?")),
"active update should guard {column}"
);
}
for column in ["format_converted", "duration_seconds"] {
assert!(
UPDATE_IF_ACTIVE_SQL.contains(&format!("{column} <=> ?")),
"active update should guard {column}"
);
}
assert!(UPDATE_IF_ACTIVE_SQL.contains("created_at = COALESCE(created_at, ?)"));
}
#[tokio::test]
async fn repository_builds_from_lazy_pool() {
let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with(
File diff suppressed because it is too large Load Diff
@@ -75,7 +75,7 @@ fn mysql_wallet_admin_builders_cover_filters_ordering_and_mapping_columns() {
let order_sql = compact_sql(admin_payment_order_list_builder(&order_query, 100, 8, 6).sql());
assert!(order_sql.contains("payment_method = ?"));
assert!(order_sql.contains(
"CASE WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired'"
"CASE WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired'"
));
assert!(order_sql.contains("ORDER BY created_at DESC, id DESC LIMIT ? OFFSET ?"));
@@ -114,6 +114,66 @@ fn compact_sql(sql: &str) -> String {
sql.split_whitespace().collect::<Vec<_>>().join(" ")
}
#[test]
fn mysql_gateway_order_uniqueness_preflights_before_persistent_changes() {
const UNIQUENESS_MIGRATION: &str = include_str!(
"../../migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql"
);
let executable_migration = UNIQUENESS_MIGRATION
.lines()
.filter(|line| !line.trim_start().starts_with("--"))
.collect::<Vec<_>>()
.join("\n");
let migration = compact_sql(&executable_migration);
let initial_cleanup = migration
.find("DROP TEMPORARY TABLE IF EXISTS aether_payment_gateway_order_uniqueness_preflight")
.expect("migration should clean up a same-session failed preflight");
let create_preflight = migration
.find("CREATE TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight")
.expect("migration should create a non-persistent conflict guard");
let seed_preflight = migration
.find("INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) VALUES (1)")
.expect("migration should seed the duplicate-key conflict guard");
let conflict_probe = migration
.find("INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) SELECT 1 FROM payment_orders")
.expect("migration should reject normalized historical conflicts");
let final_cleanup = migration
.find("DROP TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight;")
.expect("migration should remove the successful preflight guard");
let first_persistent_update = migration
.find("UPDATE payment_orders SET payment_method")
.expect("migration should normalize payment order methods");
let alter = migration
.find(
"MODIFY COLUMN gateway_order_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin NULL",
)
.expect("migration should enforce binary gateway identifiers");
let unique = migration
.find("ADD UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id")
.expect("migration should create composite uniqueness");
assert!(migration.contains("GROUP BY LOWER(TRIM(payment_method)), CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin HAVING COUNT(*) > 1 LIMIT 1"));
assert!(migration.contains("WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method))"));
let preflight = &migration[..final_cleanup];
assert!(!preflight.contains("UPDATE payment_orders"));
assert!(!preflight.contains("UPDATE payment_callbacks"));
assert!(!preflight.contains("ALTER TABLE"));
assert!(
initial_cleanup < create_preflight
&& create_preflight < seed_preflight
&& seed_preflight < conflict_probe
&& conflict_probe < final_cleanup
&& final_cleanup < first_persistent_update
&& first_persistent_update < alter,
"the conflict probe must finish before any persistent UPDATE or ALTER"
);
assert!(
alter < unique && !migration.contains("CREATE UNIQUE INDEX"),
"the collation and unique index must be one atomic ALTER TABLE"
);
}
#[tokio::test]
async fn mysql_wallet_read_repository_reads_wallet_contract_views() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")